1use 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 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 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 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 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 }
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 self.add_constraints_from_invariant_args(alias.args);
221 }
222 TyKind::Dynamic(bounds, region) => {
223 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 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 }
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 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 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 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 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 }
362 RegionKind::ReError(_) => {}
363 RegionKind::ReLateParam(..)
364 | RegionKind::RePlaceholder(..)
365 | RegionKind::ReVar(..)
366 | RegionKind::ReErased => {
367 never!(
370 "unexpected region encountered in variance \
371 inference: {:?}",
372 region
373 );
374 }
375 }
376 }
377
378 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 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 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}