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 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 #[must_use]
280 pub fn is_refine(&self) -> bool {
281 matches!(self, Self::Refine(..))
282 }
283
284 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);