Skip to main content

hir_ty/
variance.rs

1//! Module for inferring the variance of type and lifetime parameters. See the [rustc dev guide]
2//! chapter for more info.
3//!
4//! [rustc dev guide]: https://rustc-dev-guide.rust-lang.org/variance.html
5//!
6//! The implementation here differs from rustc. Rustc does a crate wide fixpoint resolution
7//! as the algorithm for determining variance is a fixpoint computation with potential cycles that
8//! need to be resolved. rust-analyzer does not want a crate-wide analysis though as that would hurt
9//! incrementality too much and as such our query is based on a per item basis.
10//!
11//! This does unfortunately run into the issue that we can run into query cycles which salsa
12//! currently does not allow to be resolved via a fixpoint computation. This will likely be resolved
13//! by the next salsa version. If not, we will likely have to adapt and go with the rustc approach
14//! while installing firewall per item queries to prevent invalidation issues.
15
16use hir_def::{
17    AdtId, GenericDefId, GenericParamId, VariantId,
18    signatures::{StructFlags, StructSignature},
19};
20use rustc_ast_ir::Mutability;
21use rustc_type_ir::{Variance, inherent::IntoKind};
22use stdx::never;
23
24use crate::{
25    db::HirDatabase,
26    generics::{Generics, generics},
27    next_solver::{
28        Const, ConstKind, DbInterner, ExistentialPredicate, GenericArgKind, GenericArgs, Pattern,
29        PatternKind, Region, RegionKind, StoredVariancesOf, TermKind, Ty, TyKind, VariancesOf,
30    },
31};
32
33pub(crate) fn variances_of(db: &dyn HirDatabase, def: GenericDefId) -> VariancesOf<'_> {
34    variances_of_query(db, def).as_ref()
35}
36
37#[salsa::tracked(
38    returns(ref),
39    cycle_fn = crate::variance::variances_of_cycle_fn,
40    cycle_initial = crate::variance::variances_of_cycle_initial,
41)]
42fn variances_of_query(db: &dyn HirDatabase, def: GenericDefId) -> StoredVariancesOf {
43    tracing::debug!("variances_of(def={:?})", def);
44    match def {
45        GenericDefId::FunctionId(_) => (),
46        GenericDefId::AdtId(adt) => {
47            if let AdtId::StructId(id) = adt {
48                let flags = &StructSignature::of(db, id).flags;
49                let types = || crate::next_solver::default_types(db);
50                if flags.contains(StructFlags::IS_UNSAFE_CELL) {
51                    return types().one_invariant.store();
52                } else if flags.contains(StructFlags::IS_PHANTOM_DATA) {
53                    return types().one_covariant.store();
54                }
55            }
56        }
57        _ => return VariancesOf::empty(DbInterner::new_no_crate(db)).store(),
58    }
59
60    let generics = generics(db, def);
61    let count = generics.len(true);
62    if count == 0 {
63        return VariancesOf::empty(DbInterner::new_no_crate(db)).store();
64    }
65    let variances =
66        Context { generics, variances: vec![Variance::Bivariant; count].into_boxed_slice(), db }
67            .solve();
68
69    VariancesOf::new_from_slice(&variances).store()
70}
71
72pub(crate) fn variances_of_cycle_fn(
73    _db: &dyn HirDatabase,
74    _: &salsa::Cycle<'_>,
75    _last_provisional_value: &StoredVariancesOf,
76    value: StoredVariancesOf,
77    _def: GenericDefId,
78) -> StoredVariancesOf {
79    value
80}
81
82fn glb(v1: Variance, v2: Variance) -> Variance {
83    // Greatest lower bound of the variance lattice as defined in The Paper:
84    //
85    //       *
86    //    -     +
87    //       o
88    match (v1, v2) {
89        (Variance::Invariant, _) | (_, Variance::Invariant) => Variance::Invariant,
90
91        (Variance::Covariant, Variance::Contravariant) => Variance::Invariant,
92        (Variance::Contravariant, Variance::Covariant) => Variance::Invariant,
93
94        (Variance::Covariant, Variance::Covariant) => Variance::Covariant,
95
96        (Variance::Contravariant, Variance::Contravariant) => Variance::Contravariant,
97
98        (x, Variance::Bivariant) | (Variance::Bivariant, x) => x,
99    }
100}
101
102pub(crate) fn variances_of_cycle_initial(
103    db: &dyn HirDatabase,
104    _: salsa::Id,
105    def: GenericDefId,
106) -> StoredVariancesOf {
107    let interner = DbInterner::new_no_crate(db);
108    let generics = generics(db, def);
109    let count = generics.len(true);
110
111    VariancesOf::new_from_iter(interner, std::iter::repeat_n(Variance::Bivariant, count)).store()
112}
113
114struct Context<'db> {
115    db: &'db dyn HirDatabase,
116    generics: Generics<'db>,
117    variances: Box<[Variance]>,
118}
119
120impl<'db> Context<'db> {
121    fn solve(mut self) -> Box<[Variance]> {
122        tracing::debug!("solve(generics={:?})", self.generics);
123        match self.generics.def() {
124            GenericDefId::AdtId(adt) => {
125                let db = self.db;
126                let mut add_constraints_from_variant = |variant| {
127                    for (_, field) in db.field_types(variant).iter() {
128                        self.add_constraints_from_ty(
129                            field.ty().instantiate_identity().skip_norm_wip(),
130                            Variance::Covariant,
131                        );
132                    }
133                };
134                match adt {
135                    AdtId::StructId(s) => add_constraints_from_variant(VariantId::StructId(s)),
136                    AdtId::UnionId(u) => add_constraints_from_variant(VariantId::UnionId(u)),
137                    AdtId::EnumId(e) => {
138                        e.enum_variants(db).variants.values().for_each(|&(variant, _)| {
139                            add_constraints_from_variant(VariantId::EnumVariantId(variant))
140                        });
141                    }
142                }
143            }
144            GenericDefId::FunctionId(f) => {
145                let sig =
146                    self.db.callable_item_signature(f.into()).instantiate_identity().skip_binder();
147                self.add_constraints_from_sig(sig.inputs_and_output.iter(), Variance::Covariant);
148            }
149            _ => {}
150        }
151        let mut variances = self.variances;
152
153        // Const parameters are always invariant.
154        // Make all const parameters invariant.
155        for (idx, param) in self.generics.iter_id(false).enumerate() {
156            if let GenericParamId::ConstParamId(_) = param {
157                variances[idx] = Variance::Invariant;
158            }
159        }
160
161        // Functions are permitted to have unused generic parameters: make those invariant.
162        if let GenericDefId::FunctionId(_) = self.generics.def() {
163            variances
164                .iter_mut()
165                .filter(|&&mut v| v == Variance::Bivariant)
166                .for_each(|v| *v = Variance::Invariant);
167        }
168
169        variances
170    }
171
172    /// Adds constraints appropriate for an instance of `ty` appearing
173    /// in a context with the generics defined in `generics` and
174    /// ambient variance `variance`
175    fn add_constraints_from_ty(&mut self, ty: Ty<'db>, variance: Variance) {
176        tracing::debug!("add_constraints_from_ty(ty={:?}, variance={:?})", ty, variance);
177        match ty.kind() {
178            TyKind::Int(_)
179            | TyKind::Uint(_)
180            | TyKind::Float(_)
181            | TyKind::Char
182            | TyKind::Bool
183            | TyKind::Never
184            | TyKind::Str
185            | TyKind::Foreign(..) => {
186                // leaf type -- noop
187            }
188            TyKind::FnDef(..)
189            | TyKind::Coroutine(..)
190            | TyKind::CoroutineClosure(..)
191            | TyKind::Closure(..) => {
192                never!("Unexpected unnameable type in variance computation: {:?}", ty);
193            }
194            TyKind::Ref(lifetime, ty, mutbl) => {
195                self.add_constraints_from_region(lifetime, variance);
196                self.add_constraints_from_mt(ty, mutbl, variance);
197            }
198            TyKind::Array(typ, len) => {
199                self.add_constraints_from_const(len);
200                self.add_constraints_from_ty(typ, variance);
201            }
202            TyKind::Slice(typ) => {
203                self.add_constraints_from_ty(typ, variance);
204            }
205            TyKind::RawPtr(ty, mutbl) => {
206                self.add_constraints_from_mt(ty, mutbl, variance);
207            }
208            TyKind::Tuple(subtys) => {
209                for subty in subtys {
210                    self.add_constraints_from_ty(subty, variance);
211                }
212            }
213            TyKind::Adt(def, args) => {
214                self.add_constraints_from_args(def.def_id().into(), args, variance);
215            }
216            TyKind::Alias(alias) => {
217                // FIXME: Probably not correct wrt. opaques.
218                self.add_constraints_from_invariant_args(alias.args);
219            }
220            TyKind::Dynamic(bounds, region) => {
221                // The type `dyn Trait<T> +'a` is covariant w/r/t `'a`:
222                self.add_constraints_from_region(region, variance);
223
224                for bound in bounds {
225                    match bound.skip_binder() {
226                        ExistentialPredicate::Trait(trait_ref) => {
227                            self.add_constraints_from_invariant_args(trait_ref.args)
228                        }
229                        ExistentialPredicate::Projection(projection) => {
230                            self.add_constraints_from_invariant_args(projection.args);
231                            match projection.term.kind() {
232                                TermKind::Ty(ty) => {
233                                    self.add_constraints_from_ty(ty, Variance::Invariant)
234                                }
235                                TermKind::Const(konst) => self.add_constraints_from_const(konst),
236                            }
237                        }
238                        ExistentialPredicate::AutoTrait(_) => {}
239                    }
240                }
241            }
242
243            // Chalk has no params, so use placeholders for now?
244            TyKind::Param(param) => self.constrain(param.index as usize, variance),
245            TyKind::FnPtr(sig, _) => {
246                self.add_constraints_from_sig(sig.skip_binder().inputs_and_output.iter(), variance);
247            }
248            TyKind::Error(_) => {
249                // we encounter this when walking the trait references for object
250                // types, where we use Error as the Self type
251            }
252            TyKind::Pat(typ, pat) => {
253                self.add_constraints_from_pat(pat);
254                self.add_constraints_from_ty(typ, variance);
255            }
256            TyKind::Bound(..) => {}
257            TyKind::CoroutineWitness(..)
258            | TyKind::Placeholder(..)
259            | TyKind::Infer(..)
260            | TyKind::UnsafeBinder(..) => {
261                never!("unexpected type encountered in variance inference: {:?}", ty)
262            }
263        }
264    }
265
266    fn add_constraints_from_pat(&mut self, pat: Pattern<'db>) {
267        match pat.kind() {
268            PatternKind::Range { start, end } => {
269                self.add_constraints_from_const(start);
270                self.add_constraints_from_const(end);
271            }
272            PatternKind::NotNull => {}
273            PatternKind::Or(patterns) => {
274                for pat in patterns {
275                    self.add_constraints_from_pat(pat)
276                }
277            }
278        }
279    }
280
281    fn add_constraints_from_invariant_args(&mut self, args: GenericArgs<'db>) {
282        for k in args.iter() {
283            match k.kind() {
284                GenericArgKind::Lifetime(lt) => {
285                    self.add_constraints_from_region(lt, Variance::Invariant)
286                }
287                GenericArgKind::Type(ty) => self.add_constraints_from_ty(ty, Variance::Invariant),
288                GenericArgKind::Const(val) => self.add_constraints_from_const(val),
289            }
290        }
291    }
292
293    /// Adds constraints appropriate for a nominal type (enum, struct,
294    /// object, etc) appearing in a context with ambient variance `variance`
295    fn add_constraints_from_args(
296        &mut self,
297        def_id: GenericDefId,
298        args: GenericArgs<'db>,
299        variance: Variance,
300    ) {
301        if args.is_empty() {
302            return;
303        }
304        let variances = self.db.variances_of(def_id);
305
306        for (k, v) in args.iter().zip(variances) {
307            match k.kind() {
308                GenericArgKind::Lifetime(lt) => {
309                    self.add_constraints_from_region(lt, variance.xform(v))
310                }
311                GenericArgKind::Type(ty) => self.add_constraints_from_ty(ty, variance.xform(v)),
312                GenericArgKind::Const(val) => self.add_constraints_from_const(val),
313            }
314        }
315    }
316
317    /// Adds constraints appropriate for a const expression `val`
318    /// in a context with ambient variance `variance`
319    fn add_constraints_from_const(&mut self, c: Const<'db>) {
320        match c.kind() {
321            ConstKind::Unevaluated(c) => self.add_constraints_from_invariant_args(c.args),
322            _ => {}
323        }
324    }
325
326    /// Adds constraints appropriate for a function with signature
327    /// `sig` appearing in a context with ambient variance `variance`
328    fn add_constraints_from_sig(
329        &mut self,
330        mut sig_tys: impl DoubleEndedIterator<Item = Ty<'db>>,
331        variance: Variance,
332    ) {
333        let contra = variance.xform(Variance::Contravariant);
334        let Some(output) = sig_tys.next_back() else {
335            return never!("function signature has no return type");
336        };
337        self.add_constraints_from_ty(output, variance);
338        for input in sig_tys {
339            self.add_constraints_from_ty(input, contra);
340        }
341    }
342
343    /// Adds constraints appropriate for a region appearing in a
344    /// context with ambient variance `variance`
345    fn add_constraints_from_region(&mut self, region: Region<'db>, variance: Variance) {
346        tracing::debug!(
347            "add_constraints_from_region(region={:?}, variance={:?})",
348            region,
349            variance
350        );
351        match region.kind() {
352            RegionKind::ReEarlyParam(param) => self.constrain(param.index as usize, variance),
353            RegionKind::ReStatic => {}
354            RegionKind::ReBound(..) => {
355                // Either a higher-ranked region inside of a type or a
356                // late-bound function parameter.
357                //
358                // We do not compute constraints for either of these.
359            }
360            RegionKind::ReError(_) => {}
361            RegionKind::ReLateParam(..)
362            | RegionKind::RePlaceholder(..)
363            | RegionKind::ReVar(..)
364            | RegionKind::ReErased => {
365                // We don't expect to see anything but 'static or bound
366                // regions when visiting member types or method types.
367                never!(
368                    "unexpected region encountered in variance \
369                      inference: {:?}",
370                    region
371                );
372            }
373        }
374    }
375
376    /// Adds constraints appropriate for a mutability-type pair
377    /// appearing in a context with ambient variance `variance`
378    fn add_constraints_from_mt(&mut self, ty: Ty<'db>, mt: Mutability, variance: Variance) {
379        self.add_constraints_from_ty(
380            ty,
381            match mt {
382                Mutability::Mut => Variance::Invariant,
383                Mutability::Not => variance,
384            },
385        );
386    }
387
388    fn constrain(&mut self, index: usize, variance: Variance) {
389        tracing::debug!(
390            "constrain(index={:?}, variance={:?}, to={:?})",
391            index,
392            self.variances[index],
393            variance
394        );
395        self.variances[index] = glb(self.variances[index], variance);
396    }
397}
398
399#[cfg(test)]
400mod tests {
401    use expect_test::{Expect, expect};
402    use hir_def::{
403        AdtId, GenericDefId, ModuleDefId, hir::generics::GenericParamDataRef, src::HasSource,
404    };
405    use itertools::Itertools;
406    use rustc_type_ir::Variance;
407    use stdx::format_to;
408    use syntax::{AstNode, ast::HasName};
409    use test_fixture::WithFixture;
410
411    use hir_def::Lookup;
412
413    use crate::{db::HirDatabase, test_db::TestDB, variance::generics};
414
415    #[test]
416    fn phantom_data() {
417        check(
418            r#"
419//- minicore: phantom_data
420
421struct Covariant<A> {
422    t: core::marker::PhantomData<A>
423}
424"#,
425            expect![[r#"
426                Covariant[A: covariant]
427            "#]],
428        );
429    }
430
431    #[test]
432    fn rustc_test_variance_types() {
433        check(
434            r#"
435//- minicore: cell
436
437use core::cell::UnsafeCell;
438
439struct InvariantMut<'a,A:'a,B:'a> { //~ ERROR ['a: +, A: o, B: o]
440    t: &'a mut (A,B)
441}
442
443struct InvariantCell<A> { //~ ERROR [A: o]
444    t: UnsafeCell<A>
445}
446
447struct InvariantIndirect<A> { //~ ERROR [A: o]
448    t: InvariantCell<A>
449}
450
451struct Covariant<A> { //~ ERROR [A: +]
452    t: A, u: fn() -> A
453}
454
455struct Contravariant<A> { //~ ERROR [A: -]
456    t: fn(A)
457}
458
459enum Enum<A,B,C> { //~ ERROR [A: +, B: -, C: o]
460    Foo(Covariant<A>),
461    Bar(Contravariant<B>),`
462    Zed(Covariant<C>,Contravariant<C>)
463}
464"#,
465            expect![[r#"
466                InvariantMut['a: covariant, A: invariant, B: invariant]
467                InvariantCell[A: invariant]
468                InvariantIndirect[A: invariant]
469                Covariant[A: covariant]
470                Contravariant[A: contravariant]
471                Enum[A: covariant, B: contravariant, C: invariant]
472            "#]],
473        );
474    }
475
476    #[test]
477    fn type_resolve_error_two_structs_deep() {
478        check(
479            r#"
480struct Hello<'a> {
481    missing: Missing<'a>,
482}
483
484struct Other<'a> {
485    hello: Hello<'a>,
486}
487"#,
488            expect![[r#"
489                Hello['a: bivariant]
490                Other['a: bivariant]
491            "#]],
492        );
493    }
494
495    #[test]
496    fn rustc_test_variance_associated_consts() {
497        check(
498            r#"
499trait Trait {
500    const Const: usize;
501}
502
503struct Foo<T: Trait> { //~ ERROR [T: o]
504    field: [u8; <T as Trait>::Const]
505}
506"#,
507            expect![[r#"
508                Foo[T: invariant]
509            "#]],
510        );
511    }
512
513    #[test]
514    fn rustc_test_variance_associated_types() {
515        check(
516            r#"
517trait Trait<'a> {
518    type Type;
519
520    fn method(&'a self) { }
521}
522
523struct Foo<'a, T : Trait<'a>> { //~ ERROR ['a: +, T: +]
524    field: (T, &'a ())
525}
526
527struct Bar<'a, T : Trait<'a>> { //~ ERROR ['a: o, T: o]
528    field: <T as Trait<'a>>::Type
529}
530
531"#,
532            expect![[r#"
533                method[Self: contravariant, 'a: contravariant]
534                Foo['a: covariant, T: covariant]
535                Bar['a: invariant, T: invariant]
536            "#]],
537        );
538    }
539
540    #[test]
541    fn rustc_test_variance_associated_types2() {
542        // FIXME: RPITs have variance, but we can't treat them as their own thing right now
543        check(
544            r#"
545trait Foo {
546    type Bar;
547}
548
549fn make() -> *const dyn Foo<Bar = &'static u32> {}
550"#,
551            expect![""],
552        );
553    }
554
555    #[test]
556    fn rustc_test_variance_trait_bounds() {
557        check(
558            r#"
559trait Getter<T> {
560    fn get(&self) -> T;
561}
562
563trait Setter<T> {
564    fn get(&self, _: T);
565}
566
567struct TestStruct<U,T:Setter<U>> { //~ ERROR [U: +, T: +]
568    t: T, u: U
569}
570
571enum TestEnum<U,T:Setter<U>> { //~ ERROR [U: *, T: +]
572    //~^ ERROR: `U` is never used
573    Foo(T)
574}
575
576struct TestContraStruct<U,T:Setter<U>> { //~ ERROR [U: *, T: +]
577    //~^ ERROR: `U` is never used
578    t: T
579}
580
581struct TestBox<U,T:Getter<U>+Setter<U>> { //~ ERROR [U: *, T: +]
582    //~^ ERROR: `U` is never used
583    t: T
584}
585"#,
586            expect![[r#"
587                get[Self: contravariant, T: covariant]
588                get[Self: contravariant, T: contravariant]
589                TestStruct[U: covariant, T: covariant]
590                TestEnum[U: bivariant, T: covariant]
591                TestContraStruct[U: bivariant, T: covariant]
592                TestBox[U: bivariant, T: covariant]
593            "#]],
594        );
595    }
596
597    #[test]
598    fn rustc_test_variance_trait_matching() {
599        check(
600            r#"
601
602trait Get<T> {
603    fn get(&self) -> T;
604}
605
606struct Cloner<T:Clone> {
607    t: T
608}
609
610impl<T:Clone> Get<T> for Cloner<T> {
611    fn get(&self) -> T {}
612}
613
614fn get<'a, G>(get: &G) -> i32
615    where G : Get<&'a i32>
616{}
617
618fn pick<'b, G>(get: &'b G, if_odd: &'b i32) -> i32
619    where G : Get<&'b i32>
620{}
621"#,
622            expect![[r#"
623                get[Self: contravariant, T: covariant]
624                Cloner[T: covariant]
625                get[T: invariant]
626                get['a: invariant, G: contravariant]
627                pick['b: contravariant, G: contravariant]
628            "#]],
629        );
630    }
631
632    #[test]
633    fn rustc_test_variance_trait_object_bound() {
634        check(
635            r#"
636enum Option<T> {
637    Some(T),
638    None
639}
640trait T { fn foo(&self); }
641
642struct TOption<'a> { //~ ERROR ['a: +]
643    v: Option<*const (dyn T + 'a)>,
644}
645"#,
646            expect![[r#"
647                Option[T: covariant]
648                foo[Self: contravariant]
649                TOption['a: covariant]
650            "#]],
651        );
652    }
653
654    #[test]
655    fn rustc_test_variance_types_bounds() {
656        check(
657            r#"
658//- minicore: send
659struct TestImm<A, B> { //~ ERROR [A: +, B: +]
660    x: A,
661    y: B,
662}
663
664struct TestMut<A, B:'static> { //~ ERROR [A: +, B: o]
665    x: A,
666    y: &'static mut B,
667}
668
669struct TestIndirect<A:'static, B:'static> { //~ ERROR [A: +, B: o]
670    m: TestMut<A, B>
671}
672
673struct TestIndirect2<A:'static, B:'static> { //~ ERROR [A: o, B: o]
674    n: TestMut<A, B>,
675    m: TestMut<B, A>
676}
677
678trait Getter<A> {
679    fn get(&self) -> A;
680}
681
682trait Setter<A> {
683    fn set(&mut self, a: A);
684}
685
686struct TestObject<A, R> { //~ ERROR [A: o, R: o]
687    n: *const (dyn Setter<A> + Send),
688    m: *const (dyn Getter<R> + Send),
689}
690"#,
691            expect![[r#"
692                TestImm[A: covariant, B: covariant]
693                TestMut[A: covariant, B: invariant]
694                TestIndirect[A: covariant, B: invariant]
695                TestIndirect2[A: invariant, B: invariant]
696                get[Self: contravariant, A: covariant]
697                set[Self: invariant, A: contravariant]
698                TestObject[A: invariant, R: invariant]
699            "#]],
700        );
701    }
702
703    #[test]
704    fn rustc_test_variance_unused_region_param() {
705        check(
706            r#"
707struct SomeStruct<'a> { x: u32 } //~ ERROR parameter `'a` is never used
708enum SomeEnum<'a> { Nothing } //~ ERROR parameter `'a` is never used
709trait SomeTrait<'a> { fn foo(&self); } // OK on traits.
710"#,
711            expect![[r#"
712                SomeStruct['a: bivariant]
713                SomeEnum['a: bivariant]
714                foo[Self: contravariant, 'a: invariant]
715            "#]],
716        );
717    }
718
719    #[test]
720    fn rustc_test_variance_unused_type_param() {
721        check(
722            r#"
723//- minicore: sized
724struct SomeStruct<A> { x: u32 }
725enum SomeEnum<A> { Nothing }
726enum ListCell<T> {
727    Cons(*const ListCell<T>),
728    Nil
729}
730
731struct SelfTyAlias<T>(*const Self);
732struct WithBounds<T: Sized> {}
733struct WithWhereBounds<T> where T: Sized {}
734struct WithOutlivesBounds<T: 'static> {}
735struct DoubleNothing<T> {
736    s: SomeStruct<T>,
737}
738
739"#,
740            expect![[r#"
741                SomeStruct[A: bivariant]
742                SomeEnum[A: bivariant]
743                ListCell[T: bivariant]
744                SelfTyAlias[T: bivariant]
745                WithBounds[T: bivariant]
746                WithWhereBounds[T: bivariant]
747                WithOutlivesBounds[T: bivariant]
748                DoubleNothing[T: bivariant]
749            "#]],
750        );
751    }
752
753    #[test]
754    fn rustc_test_variance_use_contravariant_struct1() {
755        check(
756            r#"
757struct SomeStruct<T>(fn(T));
758
759fn foo<'min,'max>(v: SomeStruct<&'max ()>)
760                  -> SomeStruct<&'min ()>
761    where 'max : 'min
762{}
763"#,
764            expect![[r#"
765                SomeStruct[T: contravariant]
766                foo['min: contravariant, 'max: covariant]
767            "#]],
768        );
769    }
770
771    #[test]
772    fn rustc_test_variance_use_contravariant_struct2() {
773        check(
774            r#"
775struct SomeStruct<T>(fn(T));
776
777fn bar<'min,'max>(v: SomeStruct<&'min ()>)
778                  -> SomeStruct<&'max ()>
779    where 'max : 'min
780{}
781"#,
782            expect![[r#"
783                SomeStruct[T: contravariant]
784                bar['min: covariant, 'max: contravariant]
785            "#]],
786        );
787    }
788
789    #[test]
790    fn rustc_test_variance_use_covariant_struct1() {
791        check(
792            r#"
793struct SomeStruct<T>(T);
794
795fn foo<'min,'max>(v: SomeStruct<&'min ()>)
796                  -> SomeStruct<&'max ()>
797    where 'max : 'min
798{}
799"#,
800            expect![[r#"
801                SomeStruct[T: covariant]
802                foo['min: contravariant, 'max: covariant]
803            "#]],
804        );
805    }
806
807    #[test]
808    fn rustc_test_variance_use_covariant_struct2() {
809        check(
810            r#"
811struct SomeStruct<T>(T);
812
813fn foo<'min,'max>(v: SomeStruct<&'max ()>)
814                  -> SomeStruct<&'min ()>
815    where 'max : 'min
816{}
817"#,
818            expect![[r#"
819                SomeStruct[T: covariant]
820                foo['min: covariant, 'max: contravariant]
821            "#]],
822        );
823    }
824
825    #[test]
826    fn rustc_test_variance_use_invariant_struct1() {
827        check(
828            r#"
829struct SomeStruct<T>(*mut T);
830
831fn foo<'min,'max>(v: SomeStruct<&'max ()>)
832                  -> SomeStruct<&'min ()>
833    where 'max : 'min
834{}
835
836fn bar<'min,'max>(v: SomeStruct<&'min ()>)
837                  -> SomeStruct<&'max ()>
838    where 'max : 'min
839{}
840"#,
841            expect![[r#"
842                SomeStruct[T: invariant]
843                foo['min: invariant, 'max: invariant]
844                bar['min: invariant, 'max: invariant]
845            "#]],
846        );
847    }
848
849    #[test]
850    fn invalid_arg_counts() {
851        check(
852            r#"
853struct S<T>(T);
854struct S2<T>(S<>);
855struct S3<T>(S<T, T>);
856"#,
857            expect![[r#"
858                S[T: covariant]
859                S2[T: bivariant]
860                S3[T: covariant]
861            "#]],
862        );
863    }
864
865    #[test]
866    fn prove_fixedpoint() {
867        check(
868            r#"
869struct FixedPoint<T, U, V>(&'static FixedPoint<(), T, U>, V);
870"#,
871            expect![[r#"
872                FixedPoint[T: covariant, U: covariant, V: covariant]
873            "#]],
874        );
875    }
876
877    #[track_caller]
878    fn check(#[rust_analyzer::rust_fixture] ra_fixture: &str, expected: Expect) {
879        // use tracing_subscriber::{layer::SubscriberExt, Layer};
880        // let my_layer = tracing_subscriber::fmt::layer();
881        // let _g = tracing::subscriber::set_default(tracing_subscriber::registry().with(
882        //     my_layer.with_filter(tracing_subscriber::filter::filter_fn(|metadata| {
883        //         metadata.target().starts_with("hir_ty::variance")
884        //     })),
885        // ));
886        let (db, file_id) = TestDB::with_single_file(ra_fixture);
887
888        crate::attach_db(&db, || {
889            let mut defs: Vec<GenericDefId> = Vec::new();
890            let module = db.module_for_file_opt(file_id.file_id(&db)).unwrap();
891            let def_map = module.def_map(&db);
892            crate::tests::visit_module(&db, def_map, module, &mut |it| {
893                defs.push(match it {
894                    ModuleDefId::FunctionId(it) => it.into(),
895                    ModuleDefId::AdtId(it) => it.into(),
896                    ModuleDefId::ConstId(it) => it.into(),
897                    ModuleDefId::TraitId(it) => it.into(),
898                    ModuleDefId::TypeAliasId(it) => it.into(),
899                    _ => return,
900                })
901            });
902            let defs = defs
903                .into_iter()
904                .filter_map(|def| {
905                    Some((
906                        def,
907                        match def {
908                            GenericDefId::FunctionId(it) => {
909                                let loc = it.lookup(&db);
910                                loc.source(&db).value.name().unwrap()
911                            }
912                            GenericDefId::AdtId(AdtId::EnumId(it)) => {
913                                let loc = it.lookup(&db);
914                                loc.source(&db).value.name().unwrap()
915                            }
916                            GenericDefId::AdtId(AdtId::StructId(it)) => {
917                                let loc = it.lookup(&db);
918                                loc.source(&db).value.name().unwrap()
919                            }
920                            GenericDefId::AdtId(AdtId::UnionId(it)) => {
921                                let loc = it.lookup(&db);
922                                loc.source(&db).value.name().unwrap()
923                            }
924                            GenericDefId::TraitId(_)
925                            | GenericDefId::TypeAliasId(_)
926                            | GenericDefId::ImplId(_)
927                            | GenericDefId::ConstId(_)
928                            | GenericDefId::StaticId(_) => return None,
929                        },
930                    ))
931                })
932                .sorted_by_key(|(_, n)| n.syntax().text_range().start());
933            let mut res = String::new();
934            for (def, name) in defs {
935                let variances = db.variances_of(def);
936                if variances.is_empty() {
937                    continue;
938                }
939                format_to!(
940                    res,
941                    "{name}[{}]\n",
942                    generics(&db, def)
943                        .iter(false)
944                        .map(|(_, param)| match param {
945                            GenericParamDataRef::TypeParamData(type_param_data) => {
946                                type_param_data.name.as_ref().unwrap()
947                            }
948                            GenericParamDataRef::ConstParamData(const_param_data) =>
949                                &const_param_data.name,
950                            GenericParamDataRef::LifetimeParamData(lifetime_param_data) => {
951                                &lifetime_param_data.name
952                            }
953                        })
954                        .zip_eq(variances)
955                        .format_with(", ", |(name, var), f| f(&format_args!(
956                            "{}: {}",
957                            name.as_str(),
958                            match var {
959                                Variance::Covariant => "covariant",
960                                Variance::Invariant => "invariant",
961                                Variance::Contravariant => "contravariant",
962                                Variance::Bivariant => "bivariant",
963                            },
964                        )))
965                );
966            }
967
968            expected.assert_eq(&res);
969        })
970    }
971}