Line data Source code
1 : #include "rust-tyty-variance-analysis-private.h"
2 : #include "rust-hir-type-check.h"
3 :
4 : namespace Rust {
5 : namespace TyTy {
6 :
7 : BaseType *
8 7258 : lookup_type (HirId ref)
9 : {
10 7258 : BaseType *ty = nullptr;
11 7258 : bool ok = Resolver::TypeCheckContext::get ()->lookup_type (ref, &ty);
12 7258 : rust_assert (ok);
13 7258 : return ty;
14 : }
15 :
16 : namespace VarianceAnalysis {
17 :
18 4971 : CrateCtx::CrateCtx () : private_ctx (new GenericTyPerCrateCtx ()) {}
19 :
20 : // Must be here because of incomplete type.
21 0 : CrateCtx::~CrateCtx () = default;
22 :
23 : void
24 3365 : CrateCtx::add_type_constraints (ADTType &type)
25 : {
26 3365 : private_ctx->process_type (type);
27 3365 : }
28 :
29 : void
30 4706 : CrateCtx::solve ()
31 : {
32 4706 : private_ctx->solve ();
33 4706 : private_ctx->debug_print_solutions ();
34 4706 : }
35 :
36 : std::vector<Variance>
37 0 : CrateCtx::query_generic_variance (const ADTType &type)
38 : {
39 0 : return private_ctx->query_generic_variance (type);
40 : }
41 :
42 : std::vector<Variance>
43 769 : CrateCtx::query_type_variances (BaseType *type)
44 : {
45 769 : TyVisitorCtx ctx (*private_ctx);
46 769 : return ctx.collect_variances (*type);
47 769 : }
48 :
49 : std::vector<Region>
50 129 : CrateCtx::query_type_regions (BaseType *type)
51 : {
52 129 : return private_ctx->query_type_regions (type);
53 : }
54 :
55 : FreeRegions
56 5 : CrateCtx::query_field_regions (const ADTType *parent, size_t variant_index,
57 : size_t field_index,
58 : const FreeRegions &parent_regions)
59 : {
60 5 : return private_ctx->query_field_regions (parent, variant_index, field_index,
61 5 : parent_regions);
62 : }
63 :
64 : Variance
65 0 : Variance::reverse () const
66 : {
67 0 : switch (kind)
68 : {
69 0 : case BIVARIANT:
70 0 : return bivariant ();
71 0 : case COVARIANT:
72 0 : return contravariant ();
73 0 : case CONTRAVARIANT:
74 0 : return covariant ();
75 0 : case INVARIANT:
76 0 : return invariant ();
77 : }
78 :
79 0 : rust_unreachable ();
80 : }
81 :
82 : Variance
83 1600 : Variance::join (Variance lhs, Variance rhs)
84 : {
85 1600 : return {Kind (lhs.kind | rhs.kind)};
86 : }
87 :
88 : void
89 1405 : Variance::join (Variance rhs)
90 : {
91 1405 : *this = join (*this, rhs);
92 1405 : }
93 :
94 : Variance
95 249 : Variance::transform (Variance lhs, Variance rhs)
96 : {
97 249 : switch (lhs.kind)
98 : {
99 0 : case BIVARIANT:
100 0 : return bivariant ();
101 247 : case COVARIANT:
102 247 : return rhs;
103 0 : case CONTRAVARIANT:
104 0 : return rhs.reverse ();
105 2 : case INVARIANT:
106 2 : return invariant ();
107 : }
108 0 : rust_unreachable ();
109 : }
110 :
111 : std::string
112 1302 : Variance::as_string () const
113 : {
114 1302 : switch (kind)
115 : {
116 247 : case BIVARIANT:
117 247 : return "o";
118 979 : case COVARIANT:
119 979 : return "+";
120 0 : case CONTRAVARIANT:
121 0 : return "-";
122 76 : case INVARIANT:
123 76 : return "*";
124 : }
125 0 : rust_unreachable ();
126 : }
127 :
128 : void
129 3365 : GenericTyPerCrateCtx::process_type (ADTType &type)
130 : {
131 3365 : GenericTyVisitorCtx (*this).process_type (type);
132 3365 : }
133 :
134 : void
135 4706 : GenericTyPerCrateCtx::solve ()
136 : {
137 4706 : rust_debug ("Variance analysis solving started:");
138 :
139 : // Fix point iteration
140 4706 : bool changed = true;
141 9443 : while (changed)
142 : {
143 4737 : changed = false;
144 4932 : for (auto constraint : constraints)
145 : {
146 195 : rust_debug ("\tapplying constraint: %s <= %s",
147 : to_string (constraint.target_index).c_str (),
148 : to_string (*constraint.term).c_str ());
149 :
150 195 : auto old_solution = solutions[constraint.target_index];
151 195 : auto new_solution
152 195 : = Variance::join (old_solution, evaluate (constraint.term));
153 :
154 195 : if (old_solution != new_solution)
155 : {
156 35 : rust_debug ("\t\tsolution changed: %s => %s",
157 : old_solution.as_string ().c_str (),
158 : new_solution.as_string ().c_str ());
159 :
160 35 : changed = true;
161 35 : solutions[constraint.target_index] = new_solution;
162 : }
163 : }
164 : }
165 :
166 4706 : constraints.clear ();
167 4706 : constraints.shrink_to_fit ();
168 4706 : }
169 :
170 : void
171 4706 : GenericTyPerCrateCtx::debug_print_solutions ()
172 : {
173 4706 : rust_debug ("Variance analysis results:");
174 :
175 8075 : for (auto type : map_from_ty_orig_ref)
176 : {
177 3369 : auto solution_index = type.second;
178 3369 : auto ref = type.first;
179 :
180 3369 : BaseType *ty = lookup_type (ref);
181 :
182 3369 : std::string result = "\t";
183 3369 : SubstitutionRef *subst = nullptr;
184 :
185 3369 : switch (ty->get_kind ())
186 : {
187 3365 : case TypeKind::ADT:
188 3365 : subst = static_cast<ADTType *> (ty);
189 : break;
190 0 : case TypeKind::FNDEF:
191 0 : subst = static_cast<FnType *> (ty);
192 : break;
193 0 : case TypeKind::CLOSURE:
194 0 : subst = static_cast<ClosureType *> (ty);
195 : break;
196 4 : case TypeKind::PROJECTION:
197 4 : subst = static_cast<ProjectionType *> (ty);
198 : break;
199 0 : default:
200 0 : rust_unreachable ();
201 : }
202 :
203 6738 : result += ty->get_name ();
204 3369 : result += "<";
205 :
206 3369 : size_t i = solution_index;
207 3391 : for (auto ®ion : subst->get_used_arguments ().get_regions ())
208 : {
209 22 : (void) region;
210 22 : if (i > solution_index)
211 1 : result += ", ";
212 44 : result += solutions[i].as_string ();
213 22 : i++;
214 : }
215 4649 : for (auto ¶m : subst->get_substs ())
216 : {
217 1280 : if (i > solution_index)
218 143 : result += ", ";
219 2560 : result += param.get_type_representation ().as_string ();
220 1280 : result += "=";
221 2560 : result += solutions[i].as_string ();
222 1280 : i++;
223 : }
224 :
225 3369 : result += ">";
226 3369 : rust_debug ("%s", result.c_str ());
227 3369 : }
228 4706 : }
229 :
230 : tl::optional<SolutionIndex>
231 3884 : GenericTyPerCrateCtx::lookup_type_index (HirId orig_ref)
232 : {
233 3884 : auto it = map_from_ty_orig_ref.find (orig_ref);
234 3884 : if (it != map_from_ty_orig_ref.end ())
235 : {
236 451 : return it->second;
237 : }
238 3433 : return tl::nullopt;
239 : }
240 :
241 : void
242 3365 : GenericTyVisitorCtx::process_type (ADTType &ty)
243 : {
244 3365 : rust_debug ("add_type_constraints: %s", ty.as_string ().c_str ());
245 :
246 3365 : first_lifetime = lookup_or_add_type (ty.get_orig_ref ());
247 3365 : first_type = first_lifetime + ty.get_used_arguments ().get_regions ().size ();
248 :
249 4641 : for (auto ¶m : ty.get_substs ())
250 1276 : param_names.push_back (param.get_type_representation ().as_string ());
251 :
252 7459 : for (const auto &variant : ty.get_variants ())
253 : {
254 4094 : if (variant->get_variant_type () != VariantDef::NUM
255 4094 : && variant->get_variant_type () != VariantDef::UNIT)
256 : {
257 7400 : for (const auto &field : variant->get_fields ())
258 4665 : add_constraints_from_ty (field->get_field_type (),
259 4665 : Variance::covariant ());
260 : }
261 : }
262 3365 : }
263 :
264 : std::string
265 0 : GenericTyPerCrateCtx::to_string (const Term &term) const
266 : {
267 0 : switch (term.kind)
268 : {
269 0 : case Term::CONST:
270 0 : return term.const_val.as_string ();
271 0 : case Term::REF:
272 0 : return "v(" + to_string (term.ref) + ")";
273 0 : case Term::TRANSFORM:
274 0 : return "(" + to_string (*term.transform.lhs) + " x "
275 0 : + to_string (*term.transform.rhs) + ")";
276 : }
277 0 : rust_unreachable ();
278 : }
279 :
280 : std::string
281 0 : GenericTyPerCrateCtx::to_string (SolutionIndex index) const
282 : {
283 : // Search all values in def_id_to_solution_index_start and find key for
284 : // largest value smaller than index
285 0 : std::pair<HirId, SolutionIndex> best = {0, 0};
286 :
287 0 : for (const auto &ty_map : map_from_ty_orig_ref)
288 : {
289 0 : if (ty_map.second <= index && ty_map.first > best.first)
290 : best = ty_map;
291 : }
292 0 : rust_assert (best.first != 0);
293 :
294 0 : BaseType *ty = lookup_type (best.first);
295 :
296 0 : std::string result = "";
297 0 : if (auto adt = ty->try_as<ADTType> ())
298 : {
299 0 : result += (adt->get_identifier ());
300 : }
301 : else
302 : {
303 0 : result += ty->as_string ();
304 : }
305 :
306 0 : result += "[" + std::to_string (index - best.first) + "]";
307 0 : return result;
308 : }
309 :
310 : Variance
311 585 : GenericTyPerCrateCtx::evaluate (Term *term)
312 : {
313 585 : switch (term->kind)
314 : {
315 195 : case Term::CONST:
316 195 : return term->const_val;
317 195 : case Term::REF:
318 195 : return solutions[term->ref];
319 195 : case Term::TRANSFORM:
320 195 : return Variance::transform (evaluate (term->transform.lhs),
321 195 : evaluate (term->transform.rhs));
322 : }
323 0 : rust_unreachable ();
324 : }
325 :
326 : std::vector<Variance>
327 141 : GenericTyPerCrateCtx::query_generic_variance (const ADTType &type)
328 : {
329 141 : auto solution_index = lookup_type_index (type.get_orig_ref ());
330 141 : rust_assert (solution_index.has_value ());
331 141 : auto num_lifetimes = type.get_num_lifetime_params ();
332 141 : auto num_types = type.get_num_type_params ();
333 :
334 141 : std::vector<Variance> result;
335 141 : result.reserve (num_lifetimes + num_types);
336 :
337 326 : for (size_t i = 0; i < num_lifetimes + num_types; ++i)
338 : {
339 44 : result.push_back (solutions[solution_index.value () + i]);
340 : }
341 :
342 141 : return result;
343 : }
344 :
345 : FreeRegions
346 5 : GenericTyPerCrateCtx::query_field_regions (const ADTType *parent,
347 : size_t variant_index,
348 : size_t field_index,
349 : const FreeRegions &parent_regions)
350 : {
351 5 : auto orig = lookup_type (parent->get_orig_ref ());
352 5 : FieldVisitorCtx ctx (*this, *parent->as<const SubstitutionRef> (),
353 5 : parent_regions);
354 5 : return ctx.collect_regions (*orig->as<const ADTType> ()
355 5 : ->get_variants ()
356 5 : .at (variant_index)
357 5 : ->get_fields ()
358 5 : .at (field_index)
359 5 : ->get_field_type ());
360 5 : }
361 : std::vector<Region>
362 129 : GenericTyPerCrateCtx::query_type_regions (BaseType *type)
363 : {
364 129 : TyVisitorCtx ctx (*this);
365 129 : return ctx.collect_regions (*type);
366 129 : }
367 :
368 : SolutionIndex
369 3743 : GenericTyVisitorCtx::lookup_or_add_type (HirId hir_id)
370 : {
371 3743 : BaseType *ty = lookup_type (hir_id);
372 3743 : auto index = ctx.lookup_type_index (hir_id);
373 3743 : if (index.has_value ())
374 : {
375 310 : return index.value ();
376 : }
377 :
378 3433 : SubstitutionRef *subst = nullptr;
379 3433 : switch (ty->get_kind ())
380 : {
381 3429 : case TypeKind::ADT:
382 3429 : subst = static_cast<ADTType *> (ty);
383 : break;
384 :
385 0 : case TypeKind::FNDEF:
386 0 : subst = static_cast<FnType *> (ty);
387 : break;
388 :
389 0 : case TypeKind::CLOSURE:
390 0 : subst = static_cast<ClosureType *> (ty);
391 : break;
392 :
393 4 : case TypeKind::PROJECTION:
394 4 : subst = static_cast<ProjectionType *> (ty);
395 : break;
396 :
397 0 : default:
398 0 : rust_sorry_at (
399 : ty->get_locus (),
400 : "This is a compiler bug: Unhandled type in variance analysis");
401 0 : break;
402 : }
403 3433 : rust_assert (subst != nullptr);
404 :
405 3433 : auto solution_index = ctx.solutions.size ();
406 3433 : ctx.map_from_ty_orig_ref.emplace (ty->get_orig_ref (), solution_index);
407 :
408 3433 : auto num_lifetime_param = subst->get_used_arguments ().get_regions ().size ();
409 3433 : auto num_type_param = subst->get_num_substitutions ();
410 :
411 8169 : for (size_t i = 0; i < num_lifetime_param + num_type_param; ++i)
412 1303 : ctx.solutions.emplace_back (Variance::bivariant ());
413 :
414 3433 : return solution_index;
415 : }
416 :
417 : void
418 5709 : GenericTyVisitorCtx::add_constraints_from_ty (BaseType *type, Term variance)
419 : {
420 5709 : rust_debug ("\tadd_constraint_from_ty: %s with v=%s",
421 : type->as_string ().c_str (), ctx.to_string (variance).c_str ());
422 :
423 5709 : Visitor visitor (*this, variance);
424 5709 : type->accept_vis (visitor);
425 5709 : }
426 :
427 : void
428 1561 : GenericTyVisitorCtx::add_constraint (SolutionIndex index, Term term)
429 : {
430 1561 : rust_debug ("\t\tadd_constraint: %s", ctx.to_string (term).c_str ());
431 :
432 1561 : if (term.kind == Term::CONST)
433 : {
434 : // Constant terms do not depend on other solutions, so we can
435 : // immediately apply them.
436 1405 : ctx.solutions[index].join (term.const_val);
437 : }
438 : else
439 : {
440 156 : ctx.constraints.emplace_back (index, new Term (term));
441 : }
442 1561 : }
443 :
444 : void
445 27 : GenericTyVisitorCtx::add_constraints_from_region (const Region ®ion,
446 : Term term)
447 : {
448 27 : if (region.is_early_bound ())
449 : {
450 16 : add_constraint (first_lifetime + region.get_index (), term);
451 : }
452 27 : }
453 :
454 : void
455 378 : GenericTyVisitorCtx::add_constraints_from_generic_args (HirId ref,
456 : SubstitutionRef &subst,
457 : Term variance,
458 : bool invariant_args)
459 : {
460 378 : SolutionIndex solution_index = lookup_or_add_type (ref);
461 :
462 378 : size_t num_lifetimes = subst.get_used_arguments ().get_regions ().size ();
463 378 : size_t num_types = subst.get_substs ().size ();
464 :
465 571 : for (size_t i = 0; i < num_lifetimes + num_types; ++i)
466 : {
467 : // TODO: What about variance from other crates?
468 193 : auto variance_i
469 : = invariant_args
470 193 : ? Term::make_transform (variance, Variance::invariant ())
471 189 : : Term::make_transform (variance,
472 : Term::make_ref (solution_index + i));
473 :
474 193 : if (i < num_lifetimes)
475 : {
476 1 : auto region_i = i;
477 1 : auto ®ion
478 1 : = subst.get_substitution_arguments ().get_mut_regions ()[region_i];
479 1 : add_constraints_from_region (region, variance_i);
480 : }
481 : else
482 : {
483 192 : auto type_i = i - num_lifetimes;
484 192 : auto arg = subst.get_arg_at (type_i);
485 192 : if (arg.has_value ())
486 : {
487 182 : add_constraints_from_ty (arg.value ().get_tyty (), variance_i);
488 : }
489 : }
490 : }
491 378 : }
492 : void
493 1545 : GenericTyVisitorCtx::add_constrints_from_param (ParamType &type, Term variance)
494 : {
495 1545 : auto it
496 1545 : = std::find (param_names.begin (), param_names.end (), type.get_name ());
497 1545 : rust_assert (it != param_names.end ());
498 :
499 1545 : auto index = first_type + std::distance (param_names.begin (), it);
500 :
501 1545 : add_constraint (index, variance);
502 1545 : }
503 :
504 : Term
505 10 : GenericTyVisitorCtx::contra (Term variance)
506 : {
507 10 : return Term::make_transform (variance, Variance::contravariant ());
508 : }
509 :
510 : void
511 1126 : TyVisitorCtx::add_constraints_from_ty (BaseType *ty, Variance variance)
512 : {
513 1126 : Visitor visitor (*this, variance);
514 1126 : ty->accept_vis (visitor);
515 1126 : }
516 :
517 : void
518 284 : TyVisitorCtx::add_constraints_from_region (const Region ®ion,
519 : Variance variance)
520 : {
521 284 : variances.push_back (variance);
522 284 : regions.push_back (region);
523 284 : }
524 :
525 : void
526 141 : TyVisitorCtx::add_constraints_from_generic_args (HirId ref,
527 : SubstitutionRef &subst,
528 : Variance variance,
529 : bool invariant_args)
530 : {
531 : // Handle function
532 141 : auto variances
533 141 : = ctx.query_generic_variance (*lookup_type (ref)->as<ADTType> ());
534 :
535 141 : size_t num_lifetimes = subst.get_used_arguments ().get_regions ().size ();
536 141 : size_t num_types = subst.get_substs ().size ();
537 :
538 185 : for (size_t i = 0; i < num_lifetimes + num_types; ++i)
539 : {
540 : // TODO: What about variance from other crates?
541 44 : auto variance_i
542 : = invariant_args
543 44 : ? Variance::transform (variance, Variance::invariant ())
544 44 : : Variance::transform (variance, variances[i]);
545 :
546 44 : if (i < num_lifetimes)
547 : {
548 44 : auto region_i = i;
549 44 : auto ®ion = subst.get_used_arguments ().get_regions ()[region_i];
550 44 : add_constraints_from_region (region, variance_i);
551 : }
552 : else
553 : {
554 0 : auto type_i = i - num_lifetimes;
555 0 : auto arg = subst.get_arg_at (type_i);
556 0 : if (arg.has_value ())
557 : {
558 0 : add_constraints_from_ty (arg.value ().get_tyty (), variance_i);
559 : }
560 : }
561 : }
562 141 : }
563 :
564 : FreeRegions
565 5 : FieldVisitorCtx::collect_regions (BaseType &ty)
566 : {
567 : // Segment the regions into ranges for each type parameter. Type parameter
568 : // at index i contains regions from type_param_ranges[i] to
569 : // type_param_ranges[i+1] (exclusive).;
570 5 : type_param_ranges.push_back (subst.get_num_lifetime_params ());
571 :
572 5 : for (size_t i = 0; i < subst.get_num_type_params (); i++)
573 : {
574 0 : auto arg = subst.get_arg_at (i);
575 0 : rust_assert (arg.has_value ());
576 0 : type_param_ranges.push_back (
577 0 : ctx.query_type_regions (arg.value ().get_tyty ()).size ());
578 : }
579 :
580 5 : add_constraints_from_ty (&ty, Variance::covariant ());
581 5 : return regions;
582 : }
583 :
584 : void
585 8 : FieldVisitorCtx::add_constraints_from_ty (BaseType *ty, Variance variance)
586 : {
587 8 : Visitor visitor (*this, variance);
588 8 : ty->accept_vis (visitor);
589 8 : }
590 :
591 : void
592 3 : FieldVisitorCtx::add_constraints_from_region (const Region ®ion,
593 : Variance variance)
594 : {
595 3 : if (region.is_early_bound ())
596 : {
597 3 : regions.push_back (parent_regions[region.get_index ()]);
598 : }
599 0 : else if (region.is_late_bound ())
600 : {
601 0 : rust_debug ("Ignoring late bound region");
602 : }
603 3 : }
604 :
605 : void
606 0 : FieldVisitorCtx::add_constrints_from_param (ParamType ¶m, Variance variance)
607 : {
608 0 : size_t param_i = subst.get_used_arguments ().find_symbol (param).value ();
609 0 : for (size_t i = type_param_ranges[param_i];
610 0 : i < type_param_ranges[param_i + 1]; i++)
611 : {
612 0 : regions.push_back (parent_regions[i]);
613 : }
614 0 : }
615 :
616 : Variance
617 0 : TyVisitorCtx::contra (Variance variance)
618 : {
619 0 : return Variance::transform (variance, Variance::contravariant ());
620 : }
621 :
622 : Term
623 189 : Term::make_ref (SolutionIndex index)
624 : {
625 189 : Term term;
626 189 : term.kind = REF;
627 189 : term.ref = index;
628 189 : return term;
629 : }
630 :
631 : Term
632 203 : Term::make_transform (Term lhs, Term rhs)
633 : {
634 203 : if (lhs.is_const () && rhs.is_const ())
635 : {
636 10 : return Variance::transform (lhs.const_val, rhs.const_val);
637 : }
638 :
639 193 : Term term;
640 193 : term.kind = TRANSFORM;
641 193 : term.transform.lhs = new Term (lhs);
642 193 : term.transform.rhs = new Term (rhs);
643 193 : return term;
644 : }
645 :
646 : } // namespace VarianceAnalysis
647 : } // namespace TyTy
648 : } // namespace Rust
|