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#[derive(Clone, Debug)]
20pub struct ClosureInfo<'db> {
21 pub members: OrderedHashMap<MemberPath<'db>, semantic::TypeId<'db>>,
24 pub snapshots: OrderedHashMap<MemberPath<'db>, semantic::TypeId<'db>>,
26}
27
28pub enum AssembleValueError<'db> {
30 Moved(MovedVar<'db>),
32 Missing,
34}
35
36#[derive(Clone, Default, Debug)]
37pub struct SemanticLoweringMapping<'db> {
38 scattered: OrderedHashMap<MemberPath<'db>, Value<'db>>,
40}
41impl<'db> SemanticLoweringMapping<'db> {
42 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 let value = self.break_into_value(ctx, path)?;
104 *value = Value::Var(var);
105 Some(())
106 }
107
108 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 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 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
220pub 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 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 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
267fn 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 return (*x).clone();
313 }
314
315 if !require_remapping {
316 let first_var = values[0];
318 if values.iter().all(|x| *x == first_var) {
319 return first_var.clone();
320 }
321 }
322
323 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 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 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
363pub 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#[derive(Clone, Debug, DebugWithDb, Eq, PartialEq)]
381#[debug_db(ExprFormatter<'db>)]
382pub struct MovedVar<'db> {
383 pub var_id: VariableId,
385 pub inference_error: InferenceError<'db>,
387 pub last_use_location: LocationId<'db>,
389}
390
391#[derive(Clone, Debug, DebugWithDb, Eq, PartialEq)]
393#[debug_db(ExprFormatter<'db>)]
394enum Value<'db> {
395 Var(VariableId),
397 Scattered(Box<Scattered<'db>>),
400 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#[derive(Clone, Debug, DebugWithDb, Eq, PartialEq)]
422#[debug_db(ExprFormatter<'db>)]
423struct Scattered<'db> {
424 ty: semantic::TypeId<'db>,
426 members: OrderedHashMap<MemberAccessKind<'db>, Value<'db>>,
427}
428
429fn 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
443pub(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
473pub(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}