Skip to main content

flux_middle/rty/
binder.rs

1use std::slice;
2
3pub use flux_arc_interner::{List, impl_slice_internable};
4use flux_common::tracked_span_bug;
5use flux_macros::{TypeFoldable, TypeVisitable};
6use flux_rustc_bridge::{
7    ToRustc,
8    ty::{BoundRegion, BoundRegionKind, Region},
9};
10use itertools::Itertools;
11use rustc_data_structures::unord::UnordMap;
12use rustc_macros::{Decodable, Encodable, TyDecodable, TyEncodable};
13use rustc_middle::ty::TyCtxt;
14use rustc_span::Symbol;
15
16use super::{
17    Expr, GenericArg, InferMode, RefineParam, Sort,
18    fold::TypeFoldable,
19    subst::{self, BoundVarReplacer, FnMutDelegate},
20};
21
22#[derive(Clone, Debug, TyEncodable, TyDecodable)]
23pub struct EarlyBinder<T>(pub T);
24
25impl<T> EarlyBinder<T> {
26    pub fn as_ref(&self) -> EarlyBinder<&T> {
27        EarlyBinder(&self.0)
28    }
29
30    pub fn as_deref(&self) -> EarlyBinder<&T::Target>
31    where
32        T: std::ops::Deref,
33    {
34        EarlyBinder(self.0.deref())
35    }
36
37    pub fn map<U>(self, f: impl FnOnce(T) -> U) -> EarlyBinder<U> {
38        EarlyBinder(f(self.0))
39    }
40
41    pub fn try_map<U, E>(self, f: impl FnOnce(T) -> Result<U, E>) -> Result<EarlyBinder<U>, E> {
42        Ok(EarlyBinder(f(self.0)?))
43    }
44
45    pub fn skip_binder(self) -> T {
46        self.0
47    }
48
49    pub fn skip_binder_ref(&self) -> &T {
50        &self.0
51    }
52
53    pub fn instantiate_identity(self) -> T {
54        self.0
55    }
56}
57
58impl<I: IntoIterator> EarlyBinder<I> {
59    pub fn iter_identity(self) -> impl Iterator<Item = I::Item> {
60        self.0.into_iter()
61    }
62}
63
64impl<T: TypeFoldable> EarlyBinder<T> {
65    pub fn instantiate(self, tcx: TyCtxt, args: &[GenericArg], refine_args: &[Expr]) -> T {
66        self.as_ref().instantiate_ref(tcx, args, refine_args)
67    }
68}
69
70impl<T: TypeFoldable> EarlyBinder<&T> {
71    pub fn instantiate_ref(self, tcx: TyCtxt, args: &[GenericArg], refine_args: &[Expr]) -> T {
72        self.0
73            .try_fold_with(&mut subst::GenericsSubstFolder::new(
74                subst::GenericArgsDelegate(args, tcx),
75                refine_args,
76            ))
77            .into_ok()
78    }
79}
80
81impl EarlyBinder<RefineParam> {
82    pub fn name(&self) -> Symbol {
83        self.skip_binder_ref().name
84    }
85}
86
87#[derive(Clone, Eq, PartialEq, Hash, TyEncodable, TyDecodable)]
88pub struct Binder<T> {
89    vars: List<BoundVariableKind>,
90    value: T,
91}
92
93impl<T> Binder<T> {
94    pub fn bind_with_vars(value: T, vars: BoundVariableKinds) -> Binder<T> {
95        Binder { vars, value }
96    }
97
98    pub fn dummy(value: T) -> Binder<T> {
99        Binder::bind_with_vars(value, List::empty())
100    }
101
102    pub fn bind_with_sorts(value: T, sorts: &[Sort]) -> Binder<T> {
103        Binder::bind_with_vars(value, sorts.iter().cloned().map_into().collect())
104    }
105
106    pub fn bind_with_sort(value: T, sort: Sort) -> Binder<T> {
107        Binder::bind_with_sorts(value, &[sort])
108    }
109
110    pub fn vars(&self) -> &List<BoundVariableKind> {
111        &self.vars
112    }
113
114    pub fn as_ref(&self) -> Binder<&T> {
115        Binder { vars: self.vars.clone(), value: &self.value }
116    }
117
118    pub fn skip_binder(self) -> T {
119        self.value
120    }
121
122    pub fn skip_binder_ref(&self) -> &T {
123        self.as_ref().skip_binder()
124    }
125
126    pub fn rebind<U>(&self, value: U) -> Binder<U> {
127        Binder { vars: self.vars.clone(), value }
128    }
129
130    pub fn map<U>(self, f: impl FnOnce(T) -> U) -> Binder<U> {
131        Binder { vars: self.vars, value: f(self.value) }
132    }
133
134    pub fn map_ref<U>(&self, f: impl FnOnce(&T) -> U) -> Binder<U> {
135        Binder { vars: self.vars.clone(), value: f(&self.value) }
136    }
137
138    pub fn try_map<U, E>(self, f: impl FnOnce(T) -> Result<U, E>) -> Result<Binder<U>, E> {
139        Ok(Binder { vars: self.vars, value: f(self.value)? })
140    }
141
142    #[track_caller]
143    pub fn sort(&self) -> Sort {
144        match &self.vars[..] {
145            [BoundVariableKind::Refine(sort, ..)] => sort.clone(),
146            _ => tracked_span_bug!("expected single-sorted binder"),
147        }
148    }
149
150    pub fn sorts(&self) -> List<Sort> {
151        self.vars
152            .iter()
153            .map(|kind| {
154                let BoundVariableKind::Refine(sort, ..) = kind else {
155                    tracked_span_bug!("unexpected BoundVariable");
156                };
157                sort.clone()
158            })
159            .collect()
160    }
161}
162
163impl<T> Binder<T>
164where
165    T: TypeFoldable,
166{
167    pub fn replace_bound_vars(
168        &self,
169        mut replace_region: impl FnMut(BoundRegion) -> Region,
170        mut replace_expr: impl FnMut(&Sort, InferMode, BoundReftKind) -> Expr,
171    ) -> T {
172        self.try_replace_bound_vars(
173            |br| Ok::<_, !>(replace_region(br)),
174            |sort, mode, kind| Ok(replace_expr(sort, mode, kind)),
175        )
176        .into_ok()
177    }
178
179    /// Like [`Binder::replace_bound_vars`] but the callbacks can fail.
180    pub fn try_replace_bound_vars<E>(
181        &self,
182        mut replace_region: impl FnMut(BoundRegion) -> Result<Region, E>,
183        mut replace_expr: impl FnMut(&Sort, InferMode, BoundReftKind) -> Result<Expr, E>,
184    ) -> Result<T, E> {
185        let mut exprs = UnordMap::default();
186        let mut regions = UnordMap::default();
187        let delegate = FnMutDelegate::new(
188            |breft| {
189                if let Some(expr) = exprs.get(&breft.var) {
190                    return Ok(Expr::clone(expr));
191                }
192                let (sort, mode, kind) = self.vars[breft.var.as_usize()].expect_refine();
193                let expr = replace_expr(sort, mode, kind)?;
194                exprs.insert(breft.var, expr.clone());
195                Ok(expr)
196            },
197            |br| {
198                if let Some(region) = regions.get(&br.var) {
199                    return Ok(*region);
200                }
201                let region = replace_region(br)?;
202                regions.insert(br.var, region);
203                Ok(region)
204            },
205        );
206
207        self.value
208            .try_fold_with(&mut BoundVarReplacer::new(delegate))
209    }
210
211    pub fn replace_bound_refts(&self, exprs: &[Expr]) -> T {
212        let delegate = FnMutDelegate::new(
213            |breft| Ok::<_, !>(exprs[breft.var.as_usize()].clone()),
214            |br| tracked_span_bug!("unexpected escaping region {br:?}"),
215        );
216        self.value
217            .try_fold_with(&mut BoundVarReplacer::new(delegate))
218            .into_ok()
219    }
220
221    pub fn replace_bound_reft(&self, expr: &Expr) -> T {
222        debug_assert!(matches!(&self.vars[..], [BoundVariableKind::Refine(..)]));
223        self.replace_bound_refts(slice::from_ref(expr))
224    }
225
226    pub fn replace_bound_refts_with(
227        &self,
228        mut f: impl FnMut(&Sort, InferMode, BoundReftKind) -> Expr,
229    ) -> T {
230        let exprs = self
231            .vars
232            .iter()
233            .map(|param| {
234                let (sort, mode, kind) = param.expect_refine();
235                f(sort, mode, kind)
236            })
237            .collect_vec();
238        self.replace_bound_refts(&exprs)
239    }
240}
241
242impl<'tcx, V> ToRustc<'tcx> for Binder<V>
243where
244    V: ToRustc<'tcx, T: rustc_middle::ty::TypeVisitable<TyCtxt<'tcx>>>,
245{
246    type T = rustc_middle::ty::Binder<'tcx, V::T>;
247
248    fn to_rustc(&self, tcx: TyCtxt<'tcx>) -> Self::T {
249        let vars = BoundVariableKind::to_rustc(&self.vars, tcx);
250        let value = self.value.to_rustc(tcx);
251        rustc_middle::ty::Binder::bind_with_vars(value, vars)
252    }
253}
254
255#[derive(
256    Clone, PartialEq, Eq, Hash, Debug, TyEncodable, TyDecodable, TypeVisitable, TypeFoldable,
257)]
258pub enum BoundVariableKind {
259    Region(BoundRegionKind),
260    Refine(Sort, InferMode, BoundReftKind),
261}
262
263impl BoundVariableKind {
264    pub fn expect_refine(&self) -> (&Sort, InferMode, BoundReftKind) {
265        if let BoundVariableKind::Refine(sort, mode, kind) = self {
266            (sort, *mode, *kind)
267        } else {
268            tracked_span_bug!("expected `BoundVariableKind::Refine`")
269        }
270    }
271
272    pub fn expect_sort(&self) -> &Sort {
273        self.expect_refine().0
274    }
275
276    /// Returns `true` if the bound variable kind is [`Refine`].
277    ///
278    /// [`Refine`]: BoundVariableKind::Refine
279    #[must_use]
280    pub fn is_refine(&self) -> bool {
281        matches!(self, Self::Refine(..))
282    }
283
284    // We can't implement [`ToRustc`] on [`List<BoundVariableKind>`] because of coherence so we add
285    // it here
286    fn to_rustc<'tcx>(
287        vars: &[Self],
288        tcx: TyCtxt<'tcx>,
289    ) -> &'tcx rustc_middle::ty::List<rustc_middle::ty::BoundVariableKind<'tcx>> {
290        tcx.mk_bound_variable_kinds_from_iter(vars.iter().flat_map(|kind| {
291            match kind {
292                BoundVariableKind::Region(brk) => {
293                    Some(rustc_middle::ty::BoundVariableKind::Region(brk.to_rustc(tcx)))
294                }
295                BoundVariableKind::Refine(..) => None,
296            }
297        }))
298    }
299}
300
301impl From<Sort> for BoundVariableKind {
302    fn from(sort: Sort) -> Self {
303        Self::Refine(sort, InferMode::EVar, BoundReftKind::Anon)
304    }
305}
306
307pub type BoundVariableKinds = List<BoundVariableKind>;
308
309#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq, PartialOrd, Ord, Encodable, Decodable)]
310pub enum BoundReftKind {
311    Anon,
312    Named(Symbol),
313}
314
315impl_slice_internable!(BoundVariableKind);