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