Skip to main content

rustc_hir_analysis/variance/
solve.rs

1//! Constraint solving
2//!
3//! The final phase iterates over the constraints, refining the variance
4//! for each inferred until a fixed point is reached. This will be the
5//! optimal solution to the constraints. The final variance for each
6//! inferred is then written into the `variance_map` in the tcx.
7
8use rustc_data_structures::fx::FxHashSet;
9use rustc_hir::def_id::DefIdMap;
10use rustc_middle::ty;
11use tracing::debug;
12
13use super::constraints::*;
14use super::terms::VarianceTerm::*;
15use super::terms::*;
16
17fn glb(v1: ty::Variance, v2: ty::Variance) -> ty::Variance {
18    // Greatest lower bound of the variance lattice as defined in The Paper:
19    //
20    //       *
21    //    -     +
22    //       o
23    match (v1, v2) {
24        (ty::Invariant, _) | (_, ty::Invariant) => ty::Invariant,
25
26        (ty::Covariant, ty::Contravariant) => ty::Invariant,
27        (ty::Contravariant, ty::Covariant) => ty::Invariant,
28
29        (ty::Covariant, ty::Covariant) => ty::Covariant,
30
31        (ty::Contravariant, ty::Contravariant) => ty::Contravariant,
32
33        (x, ty::Bivariant) | (ty::Bivariant, x) => x,
34    }
35}
36struct SolveContext<'a, 'tcx> {
37    terms_cx: TermsContext<'a, 'tcx>,
38    constraints: Vec<Constraint<'a>>,
39
40    // Maps from an InferredIndex to the inferred value for that variable.
41    solutions: Vec<ty::Variance>,
42}
43
44pub(crate) fn solve_constraints<'tcx>(
45    constraints_cx: ConstraintContext<'_, 'tcx>,
46) -> ty::CrateVariancesMap<'tcx> {
47    let ConstraintContext { terms_cx, mut constraints, .. } = constraints_cx;
48
49    let mut overridden_inferreds = FxHashSet::default();
50    let mut solutions = ::alloc::vec::from_elem(ty::Bivariant, terms_cx.inferred_terms.len())vec![ty::Bivariant; terms_cx.inferred_terms.len()];
51    // prime the solutions for certain lang items which have hard-coded variance
52    for (id, variances) in &terms_cx.lang_items {
53        let InferredIndex(start) = terms_cx.inferred_starts[id];
54        for (i, &variance) in variances.iter().enumerate() {
55            solutions[start + i] = variance;
56            overridden_inferreds.insert(start + i);
57        }
58    }
59
60    // ensure the solutions for overridden inferreds are never constrained by anything else
61    if !overridden_inferreds.is_empty() {
62        constraints.retain(|Constraint { inferred: InferredIndex(inferred), .. }| {
63            !overridden_inferreds.contains(inferred)
64        });
65    }
66
67    let mut solutions_cx = SolveContext { terms_cx, constraints, solutions };
68    solutions_cx.solve();
69    let variances = solutions_cx.create_map();
70
71    ty::CrateVariancesMap { variances }
72}
73
74impl<'a, 'tcx> SolveContext<'a, 'tcx> {
75    fn solve(&mut self) {
76        // Propagate constraints until a fixed point is reached. Note
77        // that the maximum number of iterations is 2C where C is the
78        // number of constraints (each variable can change values at most
79        // twice). Since number of constraints is linear in size of the
80        // input, so is the inference process.
81        let mut changed = true;
82        while changed {
83            changed = false;
84
85            for constraint in &self.constraints {
86                let Constraint { inferred, variance: term } = *constraint;
87                let InferredIndex(inferred) = inferred;
88                let variance = self.evaluate(term);
89                let old_value = self.solutions[inferred];
90                let new_value = glb(variance, old_value);
91                if old_value != new_value {
92                    {
    use ::tracing::__macro_support::Callsite as _;
    static __CALLSITE: ::tracing::callsite::DefaultCallsite =
        {
            static META: ::tracing::Metadata<'static> =
                {
                    ::tracing_core::metadata::Metadata::new("event /rustc-dev/c36f1457196e315bc204b9564a6a5a7fe7f5a51f/compiler/rustc_hir_analysis/src/variance/solve.rs:92",
                        "rustc_hir_analysis::variance::solve",
                        ::tracing::Level::DEBUG,
                        ::tracing_core::__macro_support::Option::Some("/rustc-dev/c36f1457196e315bc204b9564a6a5a7fe7f5a51f/compiler/rustc_hir_analysis/src/variance/solve.rs"),
                        ::tracing_core::__macro_support::Option::Some(92u32),
                        ::tracing_core::__macro_support::Option::Some("rustc_hir_analysis::variance::solve"),
                        ::tracing_core::field::FieldSet::new(&["message"],
                            ::tracing_core::callsite::Identifier(&__CALLSITE)),
                        ::tracing::metadata::Kind::EVENT)
                };
            ::tracing::callsite::DefaultCallsite::new(&META)
        };
    let enabled =
        ::tracing::Level::DEBUG <= ::tracing::level_filters::STATIC_MAX_LEVEL
                &&
                ::tracing::Level::DEBUG <=
                    ::tracing::level_filters::LevelFilter::current() &&
            {
                let interest = __CALLSITE.interest();
                !interest.is_never() &&
                    ::tracing::__macro_support::__is_enabled(__CALLSITE.metadata(),
                        interest)
            };
    if enabled {
        (|value_set: ::tracing::field::ValueSet|
                    {
                        let meta = __CALLSITE.metadata();
                        ::tracing::Event::dispatch(meta, &value_set);
                        ;
                    })({
                #[allow(unused_imports)]
                use ::tracing::field::{debug, display, Value};
                __CALLSITE.metadata().fields().value_set_all(&[(::tracing::__macro_support::Option::Some(&format_args!("updating inferred {0} from {1:?} to {2:?} due to {3:?}",
                                                    inferred, old_value, new_value, term) as
                                            &dyn ::tracing::field::Value))])
            });
    } else { ; }
};debug!(
93                        "updating inferred {} \
94                            from {:?} to {:?} due to {:?}",
95                        inferred, old_value, new_value, term
96                    );
97
98                    self.solutions[inferred] = new_value;
99                    changed = true;
100                }
101            }
102        }
103    }
104
105    fn enforce_const_invariance(&self, generics: &ty::Generics, variances: &mut [ty::Variance]) {
106        let tcx = self.terms_cx.tcx;
107
108        // Make all const parameters invariant.
109        for param in generics.own_params.iter() {
110            if let ty::GenericParamDefKind::Const { .. } = param.kind {
111                variances[param.index as usize] = ty::Invariant;
112            }
113        }
114
115        // Make all the const parameters in the parent invariant (recursively).
116        if let Some(def_id) = generics.parent {
117            self.enforce_const_invariance(tcx.generics_of(def_id), variances);
118        }
119    }
120
121    fn create_map(&self) -> DefIdMap<&'tcx [ty::Variance]> {
122        let tcx = self.terms_cx.tcx;
123
124        let solutions = &self.solutions;
125        DefIdMap::from(self.terms_cx.inferred_starts.items().map(
126            |(&def_id, &InferredIndex(start))| {
127                let generics = tcx.generics_of(def_id);
128                let count = generics.count();
129
130                let variances = tcx.arena.alloc_slice(&solutions[start..(start + count)]);
131
132                // Const parameters are always invariant.
133                self.enforce_const_invariance(generics, variances);
134
135                // Functions are permitted to have unused generic parameters: make those invariant.
136                if let ty::FnDef(..) =
137                    tcx.type_of(def_id).instantiate_identity().skip_norm_wip().kind()
138                {
139                    for variance in variances.iter_mut() {
140                        if *variance == ty::Bivariant {
141                            *variance = ty::Invariant;
142                        }
143                    }
144                }
145
146                (def_id.to_def_id(), &*variances)
147            },
148        ))
149    }
150
151    fn evaluate(&self, term: VarianceTermPtr<'a>) -> ty::Variance {
152        match *term {
153            ConstantTerm(v) => v,
154
155            TransformTerm(t1, t2) => {
156                let v1 = self.evaluate(t1);
157                let v2 = self.evaluate(t2);
158                v1.xform(v2)
159            }
160
161            InferredTerm(InferredIndex(index)) => self.solutions[index],
162        }
163    }
164}