Skip to main content

cairo_lang_lowering/lower/
refs.rs

1use cairo_lang_defs::ids::LanguageElementId;
2use cairo_lang_proc_macros::DebugWithDb;
3use cairo_lang_semantic::expr::fmt::ExprFormatter;
4use cairo_lang_semantic::expr::inference::InferenceError;
5use cairo_lang_semantic::items::structure::StructSemantic;
6use cairo_lang_semantic::types::{peel_snapshots, wrap_in_snapshots};
7use cairo_lang_semantic::usage::MemberPath;
8use cairo_lang_semantic::{self as semantic, ConcreteTypeId, MemberAccessKind, TypeLongId};
9use cairo_lang_utils::ordered_hash_map::OrderedHashMap;
10use cairo_lang_utils::{Intern, extract_matches, try_extract_matches};
11use itertools::{Itertools, chain};
12
13use super::block_builder::BlockStructRecomposer;
14use super::context::{LoweringContext, VarRequest};
15use crate::VariableId;
16use crate::ids::LocationId;
17
18/// Information about members captured by the closure and their types.
19#[derive(Clone, Debug)]
20pub struct ClosureInfo<'db> {
21    // TODO(TomerStarkware): unite copiable members and snapshots into a single map.
22    /// The members captured by the closure (not as snapshot).
23    pub members: OrderedHashMap<MemberPath<'db>, semantic::TypeId<'db>>,
24    /// The types of the captured snapshot variables.
25    pub snapshots: OrderedHashMap<MemberPath<'db>, semantic::TypeId<'db>>,
26}
27
28/// The result of `SemanticLoweringMapping::assemble_value`.
29pub enum AssembleValueError<'db> {
30    /// The variable was moved before.
31    Moved(MovedVar<'db>),
32    /// The variable is missing from `SemanticLoweringMapping::scattered`.
33    Missing,
34}
35
36#[derive(Clone, Default, Debug)]
37pub struct SemanticLoweringMapping<'db> {
38    /// Maps member paths ([MemberPath]) to lowered variable ids or scattered variable ids.
39    scattered: OrderedHashMap<MemberPath<'db>, Value<'db>>,
40}
41impl<'db> SemanticLoweringMapping<'db> {
42    /// Returns the topmost mapped member path containing the given member path, or None no such
43    /// member path exists in the mapping.
44    pub fn topmost_mapped_containing_member_path(
45        &self,
46        mut member_path: MemberPath<'db>,
47    ) -> Option<MemberPath<'db>> {
48        let mut res = None;
49        loop {
50            if self.scattered.contains_key(&member_path) {
51                res = Some(member_path.clone());
52            }
53            let MemberPath::Member { parent, .. } = member_path else {
54                return res;
55            };
56            member_path = *parent;
57        }
58    }
59
60    pub fn destructure_closure(
61        &mut self,
62        ctx: &mut BlockStructRecomposer<'_, '_, 'db>,
63        closure_var: VariableId,
64        closure_info: &ClosureInfo<'db>,
65    ) -> Vec<VariableId> {
66        ctx.deconstruct_by_types(
67            closure_var,
68            chain!(closure_info.members.values(), closure_info.snapshots.values()).cloned(),
69        )
70    }
71
72    pub fn get(
73        &mut self,
74        mut ctx: BlockStructRecomposer<'_, '_, 'db>,
75        path: &MemberPath<'db>,
76    ) -> Result<VariableId, AssembleValueError<'db>> {
77        let value = self.break_into_value(&mut ctx, path).ok_or(AssembleValueError::Missing)?;
78        let base_var = path.base_var();
79        let location_stable_ptr = match ctx.ctx.semantic_defs.get(&base_var) {
80            Some(binding) => binding.stable_ptr(ctx.ctx.db),
81            None => base_var.untyped_stable_ptr(ctx.ctx.db),
82        };
83        let location = ctx.ctx.get_location(location_stable_ptr);
84        Self::assemble_value(&mut ctx, value, location).map_err(AssembleValueError::Moved)
85    }
86
87    pub fn introduce(&mut self, path: MemberPath<'db>, var: VariableId) {
88        self.scattered.insert(path, Value::Var(var));
89    }
90
91    pub fn update(
92        &mut self,
93        ctx: &mut BlockStructRecomposer<'_, '_, 'db>,
94        path: &MemberPath<'db>,
95        var: VariableId,
96    ) -> Option<()> {
97        // TODO(TomerStarkware): check if path is captured by a closure and invalidate the closure.
98        // Right now this can only happen if we take a snapshot of the variable (as the
99        // snapshot function returns a new var).
100        // we need to ensure the borrow checker invalidates the closure when mutable capture
101        // is supported.
102
103        let value = self.break_into_value(ctx, path)?;
104        *value = Value::Var(var);
105        Some(())
106    }
107
108    /// Marks the variable at the given path as moved.
109    ///
110    /// This function should be called for non-copyable variables.
111    pub fn mark_as_used(
112        &mut self,
113        mut ctx: BlockStructRecomposer<'_, '_, 'db>,
114        path: &MemberPath<'db>,
115        moved: MovedVar<'db>,
116    ) {
117        *self.break_into_value(&mut ctx, path).unwrap() = Value::MovedVar(moved);
118    }
119
120    /// Assembles a [VariableId] from the given [Value] by recursively reconstructing it if it is
121    /// currently deconstructed.
122    ///
123    /// Returns a [MovedVar] if the variable, or any of its members, were moved before.
124    fn assemble_value(
125        ctx: &mut BlockStructRecomposer<'_, '_, 'db>,
126        value: &mut Value<'db>,
127        location: LocationId<'db>,
128    ) -> Result<VariableId, MovedVar<'db>> {
129        match value {
130            Value::Var(var) => Ok(*var),
131            Value::MovedVar(moved) => Err(moved.clone()),
132            Value::Scattered(scattered) => {
133                let ty = scattered.ty;
134                let mut moved_var = None;
135                let members = scattered
136                    .members
137                    .iter_mut()
138                    .map(|(_, value)| match Self::assemble_value(ctx, value, location) {
139                        Ok(var) => var,
140                        Err(moved) => {
141                            let var = moved.var_id;
142                            moved_var.get_or_insert(moved);
143                            var
144                        }
145                    })
146                    .collect_vec();
147                let var = ctx.reconstruct(ty, members, location);
148                *value = Value::Var(var);
149                if let Some(MovedVar { var_id: _, inference_error, last_use_location }) = moved_var
150                {
151                    Err(MovedVar { var_id: var, inference_error, last_use_location })
152                } else {
153                    Ok(var)
154                }
155            }
156        }
157    }
158
159    fn break_into_value(
160        &mut self,
161        ctx: &mut BlockStructRecomposer<'_, '_, 'db>,
162        path: &MemberPath<'db>,
163    ) -> Option<&mut Value<'db>> {
164        if self.scattered.contains_key(path) {
165            return self.scattered.get_mut(path);
166        }
167
168        let MemberPath::Member { parent, kind } = path else {
169            return None;
170        };
171
172        // The type of the parent aggregate, taken from the member's [MemberAccessKind] (the
173        // authoritative semantic type).
174        let ty = aggregate_ty(ctx.ctx.db, kind);
175        let parent_value = self.break_into_value(ctx, parent)?;
176        match parent_value {
177            Value::Var(var) => {
178                let members = ctx.deconstruct(ty, *var);
179                let members = OrderedHashMap::from_iter(
180                    members.into_iter().map(|(kind, var)| (kind, Value::Var(var))),
181                );
182                let scattered = Scattered { ty, members };
183                *parent_value = Value::Scattered(Box::new(scattered));
184            }
185            &mut Value::MovedVar(MovedVar { var_id, ref inference_error, last_use_location }) => {
186                let location = ctx.ctx.variables[var_id].location;
187                let inference_error = inference_error.clone();
188                let members = OrderedHashMap::from_iter(
189                    aggregate_members(ctx.ctx.db, ty).into_iter().map(|(kind, member_ty)| {
190                        (
191                            kind,
192                            Value::MovedVar(MovedVar {
193                                var_id: ctx.ctx.new_var(VarRequest { ty: member_ty, location }),
194                                inference_error: inference_error.clone(),
195                                last_use_location,
196                            }),
197                        )
198                    }),
199                );
200                let scattered = Scattered { ty, members };
201                *parent_value = Value::Scattered(Box::new(scattered));
202            }
203            Value::Scattered(..) => {}
204        };
205        extract_matches!(parent_value, Value::Scattered).members.get_mut(kind)
206    }
207}
208
209impl<'db> cairo_lang_debug::debug::DebugWithDb<'db> for SemanticLoweringMapping<'db> {
210    type Db = ExprFormatter<'db>;
211
212    fn fmt(&self, f: &mut std::fmt::Formatter<'_>, db: &ExprFormatter<'db>) -> std::fmt::Result {
213        for (member_path, value) in self.scattered.iter() {
214            writeln!(f, "{:?}: {value}", member_path.debug(db))?;
215        }
216        Ok(())
217    }
218}
219
220/// Merges [SemanticLoweringMapping] from multiple blocks to a single [SemanticLoweringMapping].
221///
222/// The mapping from semantic variables to lowered variables in the new block follows these rules:
223///
224/// * Variables mapped to the same lowered variable across all input blocks are kept as-is.
225/// * Local variables that appear in only a subset of the blocks are removed.
226/// * Variables with different mappings across blocks are remapped to a new lowered variable, by
227///   invoking the `remapped_callback` function.
228pub fn merge_semantics<'db, 'a>(
229    mappings: impl Iterator<Item = &'a SemanticLoweringMapping<'db>>,
230    remapped_callback: &mut impl FnMut(&MemberPath<'db>) -> VariableId,
231) -> SemanticLoweringMapping<'db>
232where
233    'db: 'a,
234{
235    // A map from [MemberPath] to its [Value] in the `mappings` where it appears.
236    // If the number of [Value]s is not the length of `mappings`, it is later dropped.
237    let mut path_to_values: OrderedHashMap<MemberPath<'_>, Vec<Value<'_>>> = Default::default();
238
239    let mut n_mappings = 0;
240    for map in mappings {
241        for (path, var) in map.scattered.iter() {
242            path_to_values.entry(path.clone()).or_default().push(var.clone());
243        }
244        n_mappings += 1;
245    }
246
247    let mut scattered: OrderedHashMap<MemberPath<'_>, Value<'_>> = Default::default();
248    for (path, values) in path_to_values {
249        // The variable is missing in one or more of the maps.
250        // It cannot be used in the merged block.
251        if values.len() != n_mappings {
252            continue;
253        }
254
255        let merged_value = compute_remapped_variables(
256            &values.iter().collect_vec(),
257            false,
258            &path,
259            remapped_callback,
260        );
261        scattered.insert(path, merged_value);
262    }
263
264    SemanticLoweringMapping { scattered }
265}
266
267/// Given a list of [Value]s that correspond to the same semantic [MemberPath] in different blocks,
268/// compute the [Value] in the merge block.
269///
270/// If all values are the same, no remapping is needed.
271/// If some of the values are [Value::Var] and some are [Value::Scattered], then all the values
272/// inside the [Value::Scattered] values need to be remapped.
273/// If all of them are [Value::Scattered], then it is possible that some of the members require
274/// remapping and some don't.
275///
276/// Pass `require_remapping=true` to indicate that during the recursion we encountered a
277/// [Value::Var], and thus we need to remap all the [Value::Scattered] values.
278/// In particular, once we have `require_remapping=true`, all the recursive calls in the subtree
279/// will have `require_remapping=true`.
280///
281/// For example, suppose `values` consists of two trees:
282/// * `A = Scattered(Scattered(v0, v1), v2)` and
283/// * `B = Scattered(Scattered(v0, v3), v4)`.
284///
285/// Then, the result will be:
286/// * `Scattered(Scattered(v0, new_var), new_var)`,
287///
288/// since `v0` is the same in both trees, but the other nodes are not.
289///
290/// If in addition to `A` and `B`, we have another tree
291/// * `C = Scattered(v5, v6)`,
292///
293/// then `v5` will need to be deconstructed, so `C` can be thought of as
294/// * `C = Scattered(Scattered(?, ?), v6)`.
295///
296/// Now, the node of `v0` also requires remapping, so the result will be:
297/// * `Scattered(Scattered(new_var, new_var), new_var)`.
298///
299/// In the recursion, when we encounter `v5`, we change `require_remapping` to `true` and drop `C`
300/// from the list of values (keeping only the scattered values).
301/// This signals that inside this subtree, all values need to be remapped (because of the children
302/// of `v5`, which are marked by `?` above).
303fn compute_remapped_variables<'db>(
304    values: &[&Value<'db>],
305    require_remapping: bool,
306    parent_path: &MemberPath<'db>,
307    remapped_callback: &mut impl FnMut(&MemberPath<'db>) -> VariableId,
308) -> Value<'db> {
309    if let Some(x) = values.iter().find(|value| matches!(value, Value::MovedVar { .. })) {
310        // If any of the values being merged is a [MovedVar], the result will be a [MovedVar].
311        // Return an arbitrary one of them.
312        return (*x).clone();
313    }
314
315    if !require_remapping {
316        // If all values are the same, no remapping is needed.
317        let first_var = values[0];
318        if values.iter().all(|x| *x == first_var) {
319            return first_var.clone();
320        }
321    }
322
323    // Collect all the `Value::Scattered` values.
324    let only_scattered: Vec<&Box<Scattered<'_>>> =
325        values.iter().filter_map(|value| try_extract_matches!(value, Value::Scattered)).collect();
326
327    if only_scattered.is_empty() {
328        let remapped_var = remapped_callback(parent_path);
329        return Value::Var(remapped_var);
330    }
331
332    // If we encountered a [Value::Var], we need to remap all the [Value::Scattered] values.
333    let require_remapping = require_remapping || only_scattered.len() < values.len();
334
335    let ty = only_scattered[0].ty;
336    let members = only_scattered[0]
337        .members
338        .keys()
339        .map(|kind| {
340            let member_path =
341                MemberPath::Member { parent: parent_path.clone().into(), kind: kind.clone() };
342            // Call `compute_remapped_variables` recursively on the scattered values.
343            // If there is a [Value::Var], `require_remapping` will be set to `true` to account
344            // for it.
345            let member_values =
346                only_scattered.iter().map(|scattered| &scattered.members[kind]).collect_vec();
347
348            (
349                kind.clone(),
350                compute_remapped_variables(
351                    &member_values,
352                    require_remapping,
353                    &member_path,
354                    remapped_callback,
355                ),
356            )
357        })
358        .collect();
359
360    Value::Scattered(Box::new(Scattered { ty, members }))
361}
362
363/// Returns an iterator to all the [MemberPath]s that appear in both mappings and have different
364/// values.
365pub fn find_changed_members<'db, 'a>(
366    semantics0: &'a SemanticLoweringMapping<'db>,
367    semantics1: &'a SemanticLoweringMapping<'db>,
368) -> impl Iterator<Item = MemberPath<'db>> + 'a {
369    semantics0.scattered.iter().filter_map(|(path, value0)| {
370        if let Some(value1) = semantics1.scattered.get(path)
371            && value0 != value1
372        {
373            return Some(path.clone());
374        }
375        None
376    })
377}
378
379/// Represents a non-copyable variable that was moved, and can no longer be used.
380#[derive(Clone, Debug, DebugWithDb, Eq, PartialEq)]
381#[debug_db(ExprFormatter<'db>)]
382pub struct MovedVar<'db> {
383    /// The type of the variable.
384    pub var_id: VariableId,
385    /// The reason it is not copyable.
386    pub inference_error: InferenceError<'db>,
387    /// The location of the last use of the moved variable. This is used to report an error.
388    pub last_use_location: LocationId<'db>,
389}
390
391/// An intermediate value for a member path.
392#[derive(Clone, Debug, DebugWithDb, Eq, PartialEq)]
393#[debug_db(ExprFormatter<'db>)]
394enum Value<'db> {
395    /// The value of member path is stored in a lowered variable.
396    Var(VariableId),
397    /// The value of the member path is not stored. If needed, it should be reconstructed from the
398    /// member values.
399    Scattered(Box<Scattered<'db>>),
400    /// Represents a non-copyable variable that was moved, and can no longer be used.
401    MovedVar(MovedVar<'db>),
402}
403
404impl<'db> std::fmt::Display for Value<'db> {
405    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
406        match self {
407            Value::Var(var) => write!(f, "v{}", var.index()),
408            Value::Scattered(scattered) => {
409                write!(
410                    f,
411                    "Scattered({})",
412                    scattered.members.values().map(|value| value.to_string()).format(", ")
413                )
414            }
415            Value::MovedVar(..) => write!(f, "MovedVar"),
416        }
417    }
418}
419
420/// A value for a non-stored member path. Recursively holds the [Value] for the members.
421#[derive(Clone, Debug, DebugWithDb, Eq, PartialEq)]
422#[debug_db(ExprFormatter<'db>)]
423struct Scattered<'db> {
424    /// The type of the scattered aggregate (a struct or a tuple), used to reconstruct it.
425    ty: semantic::TypeId<'db>,
426    members: OrderedHashMap<MemberAccessKind<'db>, Value<'db>>,
427}
428
429/// Returns the type of the aggregate (a struct or a tuple) that a [MemberAccessKind] accesses a
430/// member of.
431fn aggregate_ty<'db>(
432    db: &'db dyn salsa::Database,
433    kind: &MemberAccessKind<'db>,
434) -> semantic::TypeId<'db> {
435    match kind {
436        MemberAccessKind::Struct { concrete_struct_id, .. } => {
437            TypeLongId::Concrete(ConcreteTypeId::Struct(*concrete_struct_id)).intern(db)
438        }
439        MemberAccessKind::Index { tuple_ty, .. } => *tuple_ty,
440    }
441}
442
443/// Returns the ordered members of an aggregate type (struct or tuple), as pairs of
444/// [MemberAccessKind] and the member type.
445pub(crate) fn aggregate_members<'db>(
446    db: &'db dyn salsa::Database,
447    ty: semantic::TypeId<'db>,
448) -> Vec<(MemberAccessKind<'db>, semantic::TypeId<'db>)> {
449    match ty.long(db) {
450        TypeLongId::Concrete(ConcreteTypeId::Struct(concrete_struct_id)) => db
451            .concrete_struct_members(*concrete_struct_id)
452            .unwrap()
453            .iter()
454            .map(|(_, member)| {
455                (
456                    MemberAccessKind::Struct {
457                        concrete_struct_id: *concrete_struct_id,
458                        member_id: member.id,
459                    },
460                    member.ty,
461                )
462            })
463            .collect(),
464        TypeLongId::Tuple(tys) => tys
465            .iter()
466            .enumerate()
467            .map(|(index, member_ty)| (MemberAccessKind::Index { tuple_ty: ty, index }, *member_ty))
468            .collect(),
469        _ => unreachable!("Tried to scatter a non-aggregate type."),
470    }
471}
472
473/// Returns the snapshot-wrapped lowered types of the members of the aggregate `aggregate_ty`,
474/// along with the index of the member selected by `kind`. Returns `None` if `aggregate_ty` does
475/// not match `kind` (e.g. a tuple index on a non-tuple type).
476pub(crate) fn member_access_components<'db>(
477    ctx: &LoweringContext<'db, '_>,
478    aggregate_ty: semantic::TypeId<'db>,
479    kind: &MemberAccessKind<'db>,
480) -> Option<(Vec<semantic::TypeId<'db>>, usize)> {
481    let (n_snapshots, long_ty) = peel_snapshots(ctx.db, aggregate_ty);
482    match kind {
483        MemberAccessKind::Struct { concrete_struct_id, member_id } => {
484            let members = ctx.db.concrete_struct_members(*concrete_struct_id).ok()?;
485            let member_idx = members.iter().position(|(_, member)| member.id == *member_id)?;
486            let member_tys = members
487                .iter()
488                .map(|(_, member)| wrap_in_snapshots(ctx.db, member.ty, n_snapshots))
489                .collect();
490            Some((member_tys, member_idx))
491        }
492        MemberAccessKind::Index { index, .. } => {
493            let TypeLongId::Tuple(tys) = long_ty else {
494                return None;
495            };
496            let member_tys =
497                tys.iter().map(|ty| wrap_in_snapshots(ctx.db, *ty, n_snapshots)).collect();
498            Some((member_tys, *index))
499        }
500    }
501}