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.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 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 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 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 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 }
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 self.add_constraints_from_invariant_args(alias.args);
219 }
220 TyKind::Dynamic(bounds, region) => {
221 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 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 }
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 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 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 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 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 }
360 RegionKind::ReError(_) => {}
361 RegionKind::ReLateParam(..)
362 | RegionKind::RePlaceholder(..)
363 | RegionKind::ReVar(..)
364 | RegionKind::ReErased => {
365 never!(
368 "unexpected region encountered in variance \
369 inference: {:?}",
370 region
371 );
372 }
373 }
374 }
375
376 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 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 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}