Skip to main content

formualizer_eval/engine/arena/
ast.rs

1/// AST arena with structural sharing and deduplication
2/// Stores formula AST nodes efficiently with content-addressable storage
3use super::string_interner::{StringId, StringInterner};
4use super::value_ref::ValueRef;
5use formualizer_parse::parser::{ExternalRefKind, TableSpecifier};
6use rustc_hash::FxHashMap;
7use std::collections::hash_map::DefaultHasher;
8use std::fmt;
9use std::hash::{Hash, Hasher};
10
11/// Reference to an AST node in the arena
12#[derive(Copy, Clone, Debug, Eq, PartialEq, Hash)]
13pub struct AstNodeId(u32);
14
15impl AstNodeId {
16    pub fn as_u32(self) -> u32 {
17        self.0
18    }
19
20    pub(crate) const fn from_u32(raw: u32) -> Self {
21        Self(raw)
22    }
23}
24
25impl fmt::Display for AstNodeId {
26    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
27        write!(f, "AstNode({})", self.0)
28    }
29}
30
31#[derive(Copy, Clone, Debug, Eq, PartialEq, Hash)]
32pub struct TableSpecId(u32);
33
34impl TableSpecId {
35    pub fn as_u32(self) -> u32 {
36        self.0
37    }
38}
39
40/// Function-node name that encodes a postfix call (`LAMBDA(x,x+1)(B1)`).
41///
42/// A call is stored as `Function { name: CALL_NODE_NAME, args: [callee, args..] }`
43/// so the arena keeps its callee and arguments without a new `AstNodeData`
44/// variant. The tokenizer never produces `#` in a function name, so no parsed
45/// function can collide with it.
46pub(crate) const CALL_NODE_NAME: &str = "#CALL";
47
48/// Compact representation of AST nodes in the arena
49#[derive(Debug, Clone, PartialEq, Eq, Hash)]
50pub enum AstNodeData {
51    /// Literal value
52    Literal(ValueRef),
53
54    /// Explicitly omitted function argument.
55    Omitted,
56
57    /// Cell or range reference
58    Reference {
59        original_id: StringId,    // Original reference string
60        ref_type: CompactRefType, // Compact reference representation
61    },
62
63    /// Unary operation
64    UnaryOp { op_id: StringId, expr_id: AstNodeId },
65
66    /// Binary operation
67    BinaryOp {
68        op_id: StringId,
69        left_id: AstNodeId,
70        right_id: AstNodeId,
71    },
72
73    /// Function call
74    Function {
75        name_id: StringId,
76        args_offset: u32, // Index into args array
77        args_count: u16,  // Number of arguments
78    },
79
80    /// Array literal
81    Array {
82        rows: u16,
83        cols: u16,
84        elements_offset: u32, // Index into elements array
85    },
86}
87
88/// Identifies a sheet either by stable registry id or by unresolved name.
89#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
90pub enum SheetKey {
91    Id(u16),
92    Name(StringId),
93}
94
95/// Compact representation of reference types
96#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
97pub enum CompactRefType {
98    Cell {
99        sheet: Option<SheetKey>,
100        row: u32,
101        col: u32,
102        row_abs: bool,
103        col_abs: bool,
104    },
105    Range {
106        sheet: Option<SheetKey>,
107        start_row: u32,
108        start_col: u32,
109        end_row: u32,
110        end_col: u32,
111        start_row_abs: bool,
112        start_col_abs: bool,
113        end_row_abs: bool,
114        end_col_abs: bool,
115    },
116    External {
117        raw_id: StringId,
118        book_id: StringId,
119        sheet_id: StringId,
120        kind: ExternalRefKind,
121    },
122    NamedRange(StringId),
123    Table {
124        name_id: StringId,
125        specifier_id: Option<TableSpecId>,
126    },
127    /// 3D cell reference (`Sheet1:Sheet3!A1`).
128    Cell3D {
129        sheet_first: StringId,
130        sheet_last: StringId,
131        row: u32,
132        col: u32,
133        row_abs: bool,
134        col_abs: bool,
135    },
136    /// 3D range reference (`Sheet1:Sheet3!A1:B2`).
137    Range3D {
138        sheet_first: StringId,
139        sheet_last: StringId,
140        start_row: u32,
141        start_col: u32,
142        end_row: u32,
143        end_col: u32,
144        start_row_abs: bool,
145        start_col_abs: bool,
146        end_row_abs: bool,
147        end_col_abs: bool,
148    },
149}
150
151/// Arena entry containing structural node data plus canonical metadata.
152///
153/// Phase 1 keeps metadata at its default value for existing raw interning
154/// paths. Future canonical interning paths populate `meta` via the arena
155/// canonicalization helpers.
156#[derive(Debug, Clone, PartialEq, Eq, Hash)]
157pub(crate) struct AstNodeEntry {
158    pub(crate) data: AstNodeData,
159    pub(crate) meta: AstNodeMetadata,
160    /// Which of SUBTOTAL ([`SUBTOTAL_CALL`]) and AGGREGATE
161    /// ([`AGGREGATE_CALL`]) the subtree calls: a cell whose formula calls
162    /// them is skipped by an enclosing SUBTOTAL/AGGREGATE (Excel's "nested
163    /// subtotals are ignored"). Derived from the node's data and its
164    /// children's bits, so dedup and compaction carry it unchanged.
165    pub(crate) subtotal_calls: u8,
166}
167
168/// [`AstNodeEntry::subtotal_calls`] bit: the subtree calls SUBTOTAL.
169pub(crate) const SUBTOTAL_CALL: u8 = 1;
170/// [`AstNodeEntry::subtotal_calls`] bit: the subtree calls AGGREGATE.
171pub(crate) const AGGREGATE_CALL: u8 = 2;
172
173/// The [`AstNodeEntry::subtotal_calls`] bit of a function name
174/// (case-insensitive, with or without the `_xlfn.` prefix).
175pub(crate) fn subtotal_call_bit(name: &str) -> u8 {
176    let name = match name.get(..6) {
177        Some(prefix) if prefix.eq_ignore_ascii_case("_xlfn.") => &name[6..],
178        _ => name,
179    };
180    if name.eq_ignore_ascii_case("SUBTOTAL") {
181        SUBTOTAL_CALL
182    } else if name.eq_ignore_ascii_case("AGGREGATE") {
183        AGGREGATE_CALL
184    } else {
185        0
186    }
187}
188
189/// Canonical metadata associated with an arena AST node.
190#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
191pub(crate) struct AstNodeMetadata {
192    pub(crate) canonical_hash: u64,
193    pub(crate) labels: CanonicalLabels,
194    pub(crate) reference_returning_admission: ReferenceReturningAdmission,
195}
196
197#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
198pub(crate) struct ReferenceReturningAdmission(u8);
199
200impl ReferenceReturningAdmission {
201    pub(crate) const fn new(safe: bool, scalar: bool) -> Self {
202        Self((safe as u8) | ((scalar as u8) << 1))
203    }
204
205    pub(crate) const fn safe(self) -> bool {
206        self.0 & 1 != 0
207    }
208
209    pub(crate) const fn scalar(self) -> bool {
210        self.0 & 2 != 0
211    }
212}
213
214/// Compact bitset labels for arena-native canonicalization.
215#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
216pub(crate) struct CanonicalLabels {
217    pub(crate) flags: u64,
218    pub(crate) rejects: u64,
219}
220
221#[allow(dead_code)]
222impl CanonicalLabels {
223    pub(crate) const FLAG_RELATIVE_ONLY: u64 = 1 << 0;
224    pub(crate) const FLAG_ABSOLUTE_ONLY: u64 = 1 << 1;
225    pub(crate) const FLAG_MIXED_ANCHORS: u64 = 1 << 2;
226    pub(crate) const FLAG_VOLATILE: u64 = 1 << 3;
227    pub(crate) const FLAG_DYNAMIC: u64 = 1 << 4;
228    pub(crate) const FLAG_CONTAINS_STRUCTURED_REF: u64 = 1 << 5;
229    pub(crate) const FLAG_NEEDS_PLACEMENT_REWRITE: u64 = 1 << 6;
230    pub(crate) const FLAG_CONTAINS_NAME: u64 = 1 << 7;
231    pub(crate) const FLAG_CONTAINS_TABLE: u64 = 1 << 8;
232    pub(crate) const FLAG_CONTAINS_RANGE: u64 = 1 << 9;
233    pub(crate) const FLAG_CONTAINS_ARRAY: u64 = 1 << 10;
234    pub(crate) const FLAG_CONTAINS_LET_LAMBDA: u64 = 1 << 11;
235    pub(crate) const FLAG_CONTAINS_FUNCTION: u64 = 1 << 12;
236    pub(crate) const FLAG_EXPLICIT_SHEET: u64 = 1 << 13;
237    pub(crate) const FLAG_CURRENT_SHEET: u64 = 1 << 14;
238
239    // Reject bits mirror `CanonicalRejectReason` variants in
240    // `formula_plane/template_canonical.rs` by kind/variant order.
241    pub(crate) const REJECT_INVALID_PLACEMENT_ANCHOR: u64 = 1 << 0;
242    pub(crate) const REJECT_DYNAMIC_REFERENCE: u64 = 1 << 1;
243    pub(crate) const REJECT_UNKNOWN_OR_CUSTOM_FUNCTION: u64 = 1 << 2;
244    pub(crate) const REJECT_LOCAL_ENVIRONMENT: u64 = 1 << 3;
245    pub(crate) const REJECT_PARSER_VOLATILE_FLAG: u64 = 1 << 4;
246    pub(crate) const REJECT_VOLATILE_FUNCTION: u64 = 1 << 5;
247    pub(crate) const REJECT_REFERENCE_RETURNING_FUNCTION: u64 = 1 << 6;
248    pub(crate) const REJECT_ARRAY_OR_SPILL_FUNCTION: u64 = 1 << 7;
249    pub(crate) const REJECT_ARRAY_LITERAL: u64 = 1 << 8;
250    pub(crate) const REJECT_SPILL_REFERENCE: u64 = 1 << 9;
251    pub(crate) const REJECT_SPILL_RESULT_REGION_OPERATOR: u64 = 1 << 10;
252    pub(crate) const REJECT_IMPLICIT_INTERSECTION_OPERATOR: u64 = 1 << 11;
253    pub(crate) const REJECT_CALL_EXPRESSION: u64 = 1 << 12;
254    // Bit 13 (REJECT_NAMED_REFERENCE) retired: named references canonicalize
255    // by identity and are accepted/rejected at read-projection time instead.
256    pub(crate) const REJECT_STRUCTURED_REFERENCE: u64 = 1 << 14;
257    pub(crate) const REJECT_STRUCTURED_REFERENCE_CURRENT_ROW: u64 = 1 << 15;
258    pub(crate) const REJECT_THREE_D_REFERENCE: u64 = 1 << 16;
259    pub(crate) const REJECT_EXTERNAL_REFERENCE: u64 = 1 << 17;
260    pub(crate) const REJECT_OPEN_RANGE_REFERENCE: u64 = 1 << 18;
261    pub(crate) const REJECT_WHOLE_AXIS_REFERENCE: u64 = 1 << 19;
262    pub(crate) const REJECT_UNSUPPORTED_REFERENCE: u64 = 1 << 20;
263
264    pub(crate) fn has_flag(self, flag: u64) -> bool {
265        self.flags & flag != 0
266    }
267
268    pub(crate) fn has_reject(self, reject: u64) -> bool {
269        self.rejects & reject != 0
270    }
271}
272
273/// Arena for storing AST nodes with deduplication
274pub struct AstArena {
275    /// Node storage
276    nodes: Vec<AstNodeEntry>,
277
278    /// Hash -> node index for deduplication
279    dedup_map: FxHashMap<u64, AstNodeId>,
280
281    /// Function arguments storage (flattened)
282    function_args: Vec<AstNodeId>,
283
284    /// Array elements storage (flattened)
285    array_elements: Vec<AstNodeId>,
286
287    /// String pool for operators and function names
288    strings: StringInterner,
289
290    /// Structured table specifiers
291    table_specs: Vec<TableSpecifier>,
292    table_spec_dedup: FxHashMap<u64, TableSpecId>,
293
294    /// Statistics
295    dedup_hits: usize,
296}
297
298impl AstArena {
299    pub fn new() -> Self {
300        Self {
301            nodes: Vec::new(),
302            dedup_map: FxHashMap::default(),
303            function_args: Vec::new(),
304            array_elements: Vec::new(),
305            strings: StringInterner::new(),
306            table_specs: Vec::new(),
307            table_spec_dedup: FxHashMap::default(),
308            dedup_hits: 0,
309        }
310    }
311
312    pub fn with_capacity(node_cap: usize) -> Self {
313        Self {
314            nodes: Vec::with_capacity(node_cap),
315            dedup_map: FxHashMap::with_capacity_and_hasher(node_cap, Default::default()),
316            function_args: Vec::with_capacity(node_cap * 2), // Assume avg 2 args
317            array_elements: Vec::with_capacity(node_cap),
318            strings: StringInterner::with_capacity(node_cap / 10),
319            table_specs: Vec::new(),
320            table_spec_dedup: FxHashMap::default(),
321            dedup_hits: 0,
322        }
323    }
324
325    /// Insert a node, deduplicating if it already exists.
326    ///
327    /// Phase 1 preserves raw/literal interning semantics: metadata is filled
328    /// with zeros and is not part of the dedup key.
329    pub fn insert(&mut self, node: AstNodeData) -> AstNodeId {
330        self.insert_entry(node, AstNodeMetadata::default())
331    }
332
333    /// Insert a node with precomputed canonical metadata.
334    ///
335    /// This is unused in Phase 1 and reserved for the future canonical
336    /// interning path. Deduplication remains structural on `AstNodeData` only;
337    /// if a matching node already exists, the existing entry (and metadata)
338    /// wins.
339    #[allow(dead_code)]
340    pub(crate) fn insert_with_meta(
341        &mut self,
342        node: AstNodeData,
343        meta: AstNodeMetadata,
344    ) -> AstNodeId {
345        self.insert_entry(node, meta)
346    }
347
348    fn insert_entry(&mut self, node: AstNodeData, meta: AstNodeMetadata) -> AstNodeId {
349        // Compute structural hash. Metadata is deliberately excluded.
350        let hash = self.hash_node(&node);
351
352        // Check for existing node
353        if let Some(&id) = self.dedup_map.get(&hash) {
354            // Verify it's actually the same (handle hash collisions)
355            if self.nodes[id.0 as usize].data == node {
356                self.dedup_hits += 1;
357                return id;
358            }
359        }
360
361        // Add new node
362        let id = AstNodeId(self.nodes.len() as u32);
363        let subtotal_calls = self.node_subtotal_calls(&node);
364        self.nodes.push(AstNodeEntry {
365            data: node,
366            meta,
367            subtotal_calls,
368        });
369        self.dedup_map.insert(hash, id);
370        id
371    }
372
373    /// Insert a literal node
374    pub fn insert_literal(&mut self, value: ValueRef) -> AstNodeId {
375        self.insert(AstNodeData::Literal(value))
376    }
377
378    /// Insert an explicitly omitted argument node.
379    pub(crate) fn insert_omitted(&mut self) -> AstNodeId {
380        self.insert(AstNodeData::Omitted)
381    }
382
383    /// Insert a reference node
384    pub fn insert_reference(&mut self, original: &str, ref_type: CompactRefType) -> AstNodeId {
385        let original_id = self.strings.intern(original);
386        self.insert(AstNodeData::Reference {
387            original_id,
388            ref_type,
389        })
390    }
391
392    /// Insert a unary operation node
393    pub fn insert_unary_op(&mut self, op: &str, expr: AstNodeId) -> AstNodeId {
394        let op_id = self.strings.intern(op);
395        self.insert(AstNodeData::UnaryOp {
396            op_id,
397            expr_id: expr,
398        })
399    }
400
401    /// Insert a binary operation node
402    pub fn insert_binary_op(&mut self, op: &str, left: AstNodeId, right: AstNodeId) -> AstNodeId {
403        let op_id = self.strings.intern(op);
404        self.insert(AstNodeData::BinaryOp {
405            op_id,
406            left_id: left,
407            right_id: right,
408        })
409    }
410
411    /// Insert a function call node
412    pub fn insert_function(&mut self, name: &str, args: Vec<AstNodeId>) -> AstNodeId {
413        let name_id = self.strings.intern(name);
414        let args_offset = self.function_args.len() as u32;
415        let args_count = args.len() as u16;
416
417        self.function_args.extend(args);
418
419        self.insert(AstNodeData::Function {
420            name_id,
421            args_offset,
422            args_count,
423        })
424    }
425
426    /// Insert an array literal node
427    pub fn insert_array(&mut self, rows: u16, cols: u16, elements: Vec<AstNodeId>) -> AstNodeId {
428        assert_eq!(
429            elements.len(),
430            (rows * cols) as usize,
431            "Array dimensions don't match element count"
432        );
433
434        let elements_offset = self.array_elements.len() as u32;
435        self.array_elements.extend(elements);
436
437        self.insert(AstNodeData::Array {
438            rows,
439            cols,
440            elements_offset,
441        })
442    }
443
444    /// Get a node by ID
445    pub fn get(&self, id: AstNodeId) -> Option<&AstNodeData> {
446        self.nodes.get(id.0 as usize).map(|entry| &entry.data)
447    }
448
449    /// Get an arena entry by ID.
450    #[allow(dead_code)]
451    pub(crate) fn entry(&self, id: AstNodeId) -> Option<&AstNodeEntry> {
452        self.nodes.get(id.0 as usize)
453    }
454
455    /// Get canonical metadata for a node by ID.
456    #[allow(dead_code)]
457    pub(crate) fn metadata(&self, id: AstNodeId) -> Option<AstNodeMetadata> {
458        self.entry(id).map(|entry| entry.meta)
459    }
460
461    /// The SUBTOTAL/AGGREGATE call bits of the subtree rooted at `id`
462    /// (precomputed at insertion; O(1)).
463    pub(crate) fn subtotal_calls(&self, id: AstNodeId) -> u8 {
464        self.nodes
465            .get(id.0 as usize)
466            .map_or(0, |entry| entry.subtotal_calls)
467    }
468
469    /// `subtotal_calls` of a node about to be inserted: its own function
470    /// name's bit and its children's (children are inserted first).
471    fn node_subtotal_calls(&self, node: &AstNodeData) -> u8 {
472        let children = |ids: &[AstNodeId]| {
473            ids.iter()
474                .fold(0, |bits, id| bits | self.subtotal_calls(*id))
475        };
476        match node {
477            AstNodeData::Function {
478                name_id,
479                args_offset,
480                args_count,
481            } => {
482                let start = *args_offset as usize;
483                subtotal_call_bit(self.strings.resolve(*name_id))
484                    | children(&self.function_args[start..start + *args_count as usize])
485            }
486            AstNodeData::UnaryOp { expr_id, .. } => self.subtotal_calls(*expr_id),
487            AstNodeData::BinaryOp {
488                left_id, right_id, ..
489            } => self.subtotal_calls(*left_id) | self.subtotal_calls(*right_id),
490            AstNodeData::Array {
491                rows,
492                cols,
493                elements_offset,
494            } => {
495                let start = *elements_offset as usize;
496                children(&self.array_elements[start..start + (*rows as usize) * (*cols as usize)])
497            }
498            AstNodeData::Literal(_) | AstNodeData::Omitted | AstNodeData::Reference { .. } => 0,
499        }
500    }
501
502    /// Get function arguments for a function node
503    pub fn get_function_args(&self, id: AstNodeId) -> Option<&[AstNodeId]> {
504        match self.get(id)? {
505            AstNodeData::Function {
506                args_offset,
507                args_count,
508                ..
509            } => {
510                let start = *args_offset as usize;
511                let end = start + *args_count as usize;
512                Some(&self.function_args[start..end])
513            }
514            _ => None,
515        }
516    }
517
518    /// Get array elements for an array node
519    pub fn get_array_elements(&self, id: AstNodeId) -> Option<&[AstNodeId]> {
520        match self.get(id)? {
521            AstNodeData::Array {
522                rows,
523                cols,
524                elements_offset,
525            } => {
526                let start = *elements_offset as usize;
527                let count = (*rows * *cols) as usize;
528                let end = start + count;
529                Some(&self.array_elements[start..end])
530            }
531            _ => None,
532        }
533    }
534
535    pub fn get_array_elements_info(&self, id: AstNodeId) -> Option<(u16, u16, &[AstNodeId])> {
536        match self.get(id)? {
537            AstNodeData::Array { rows, cols, .. } => {
538                let elements = self.get_array_elements(id)?;
539                Some((*rows, *cols, elements))
540            }
541            _ => None,
542        }
543    }
544
545    /// Resolve a string ID to its content
546    pub fn resolve_string(&self, id: StringId) -> &str {
547        self.strings.resolve(id)
548    }
549
550    /// Get the string interner (for external use)
551    pub fn strings(&self) -> &StringInterner {
552        &self.strings
553    }
554
555    /// Get mutable access to the string interner
556    pub fn strings_mut(&mut self) -> &mut StringInterner {
557        &mut self.strings
558    }
559
560    pub fn intern_table_specifier(&mut self, specifier: &TableSpecifier) -> TableSpecId {
561        let hash = {
562            let mut hasher = DefaultHasher::new();
563            specifier.hash(&mut hasher);
564            hasher.finish()
565        };
566
567        if let Some(&id) = self.table_spec_dedup.get(&hash)
568            && self
569                .table_specs
570                .get(id.0 as usize)
571                .is_some_and(|existing| existing == specifier)
572        {
573            return id;
574        }
575
576        let id = TableSpecId(self.table_specs.len() as u32);
577        self.table_specs.push(specifier.clone());
578        self.table_spec_dedup.insert(hash, id);
579        id
580    }
581
582    pub fn resolve_table_specifier(&self, id: TableSpecId) -> Option<&TableSpecifier> {
583        self.table_specs.get(id.0 as usize)
584    }
585
586    /// Compute hash for a node
587    fn hash_node(&self, node: &AstNodeData) -> u64 {
588        let mut hasher = DefaultHasher::new();
589        node.hash(&mut hasher);
590        hasher.finish()
591    }
592
593    /// Keep only the nodes reachable from `roots` (Program 2 compression:
594    /// family members' own trees become garbage once they reference their
595    /// template). Returns the remap, indexed by old id: `u32::MAX` for a
596    /// dropped node. Strings and table specifiers are kept as they are;
597    /// node metadata travels with its node. Every holder of an id must be
598    /// remapped by the caller.
599    pub(crate) fn compact(
600        &mut self,
601        roots: impl IntoIterator<Item = AstNodeId>,
602    ) -> (Vec<u32>, super::string_interner::StringGarbage) {
603        let old_len = self.nodes.len();
604        let mut remap = vec![u32::MAX; old_len];
605        let mut nodes: Vec<AstNodeEntry> = Vec::new();
606        let mut dedup_map: FxHashMap<u64, AstNodeId> = FxHashMap::default();
607        let mut function_args: Vec<AstNodeId> = Vec::new();
608        let mut array_elements: Vec<AstNodeId> = Vec::new();
609        // Iterative post-order: (id, children pushed).
610        let mut stack: Vec<(u32, bool)> = Vec::new();
611        for root in roots {
612            let r = root.0 as usize;
613            if r >= old_len || remap[r] != u32::MAX {
614                continue;
615            }
616            stack.push((root.0, false));
617            while let Some((id, expanded)) = stack.pop() {
618                let i = id as usize;
619                if remap[i] != u32::MAX {
620                    continue;
621                }
622                let entry = &self.nodes[i];
623                let children: smallvec::SmallVec<[AstNodeId; 8]> = match &entry.data {
624                    AstNodeData::UnaryOp { expr_id, .. } => smallvec::smallvec![*expr_id],
625                    AstNodeData::BinaryOp {
626                        left_id, right_id, ..
627                    } => smallvec::smallvec![*left_id, *right_id],
628                    AstNodeData::Function {
629                        args_offset,
630                        args_count,
631                        ..
632                    } => self.function_args
633                        [*args_offset as usize..*args_offset as usize + *args_count as usize]
634                        .iter()
635                        .copied()
636                        .collect(),
637                    AstNodeData::Array {
638                        rows,
639                        cols,
640                        elements_offset,
641                    } => {
642                        let n = *rows as usize * *cols as usize;
643                        self.array_elements
644                            [*elements_offset as usize..*elements_offset as usize + n]
645                            .iter()
646                            .copied()
647                            .collect()
648                    }
649                    AstNodeData::Literal(_)
650                    | AstNodeData::Omitted
651                    | AstNodeData::Reference { .. } => smallvec::SmallVec::new(),
652                };
653                if !expanded && children.iter().any(|c| remap[c.0 as usize] == u32::MAX) {
654                    stack.push((id, true));
655                    for c in children.iter().rev() {
656                        if remap[c.0 as usize] == u32::MAX {
657                            stack.push((c.0, false));
658                        }
659                    }
660                    continue;
661                }
662                let m = |c: AstNodeId| AstNodeId(remap[c.0 as usize]);
663                let data = match &entry.data {
664                    AstNodeData::UnaryOp { op_id, expr_id } => AstNodeData::UnaryOp {
665                        op_id: *op_id,
666                        expr_id: m(*expr_id),
667                    },
668                    AstNodeData::BinaryOp {
669                        op_id,
670                        left_id,
671                        right_id,
672                    } => AstNodeData::BinaryOp {
673                        op_id: *op_id,
674                        left_id: m(*left_id),
675                        right_id: m(*right_id),
676                    },
677                    AstNodeData::Function {
678                        name_id,
679                        args_count,
680                        ..
681                    } => {
682                        let args_offset = function_args.len() as u32;
683                        function_args.extend(children.iter().map(|&c| m(c)));
684                        AstNodeData::Function {
685                            name_id: *name_id,
686                            args_offset,
687                            args_count: *args_count,
688                        }
689                    }
690                    AstNodeData::Array { rows, cols, .. } => {
691                        let elements_offset = array_elements.len() as u32;
692                        array_elements.extend(children.iter().map(|&c| m(c)));
693                        AstNodeData::Array {
694                            rows: *rows,
695                            cols: *cols,
696                            elements_offset,
697                        }
698                    }
699                    other => other.clone(),
700                };
701                let meta = entry.meta;
702                let subtotal_calls = entry.subtotal_calls;
703                let hash = self.hash_node(&data);
704                let new_id = match dedup_map.get(&hash) {
705                    Some(&existing) if nodes[existing.0 as usize].data == data => existing,
706                    _ => {
707                        let new_id = AstNodeId(nodes.len() as u32);
708                        nodes.push(AstNodeEntry {
709                            data,
710                            meta,
711                            subtotal_calls,
712                        });
713                        dedup_map.insert(hash, new_id);
714                        new_id
715                    }
716                };
717                remap[i] = new_id.0;
718            }
719        }
720        // Free the texts no kept node names (reference texts of dropped
721        // members, mostly). Ids stay stable: the authority's tokens hold
722        // operator, function and name ids of live templates.
723        let mut live = vec![false; self.strings.len()];
724        let mut mark = |id: super::string_interner::StringId| {
725            if let Some(slot) = live.get_mut(id.as_u32() as usize) {
726                *slot = true;
727            }
728        };
729        for entry in &nodes {
730            match &entry.data {
731                AstNodeData::Reference {
732                    original_id,
733                    ref_type,
734                } => {
735                    mark(*original_id);
736                    match ref_type {
737                        CompactRefType::Cell { sheet, .. }
738                        | CompactRefType::Range { sheet, .. } => {
739                            if let Some(SheetKey::Name(id)) = sheet {
740                                mark(*id);
741                            }
742                        }
743                        CompactRefType::External {
744                            raw_id,
745                            book_id,
746                            sheet_id,
747                            ..
748                        } => {
749                            mark(*raw_id);
750                            mark(*book_id);
751                            mark(*sheet_id);
752                        }
753                        CompactRefType::NamedRange(id) => mark(*id),
754                        CompactRefType::Table { name_id, .. } => mark(*name_id),
755                        CompactRefType::Cell3D {
756                            sheet_first,
757                            sheet_last,
758                            ..
759                        }
760                        | CompactRefType::Range3D {
761                            sheet_first,
762                            sheet_last,
763                            ..
764                        } => {
765                            mark(*sheet_first);
766                            mark(*sheet_last);
767                        }
768                    }
769                }
770                AstNodeData::UnaryOp { op_id, .. } | AstNodeData::BinaryOp { op_id, .. } => {
771                    mark(*op_id)
772                }
773                AstNodeData::Function { name_id, .. } => mark(*name_id),
774                AstNodeData::Literal(_) | AstNodeData::Omitted | AstNodeData::Array { .. } => {}
775            }
776        }
777        let garbage = self.strings.free_dead(&live);
778        nodes.shrink_to_fit();
779        function_args.shrink_to_fit();
780        array_elements.shrink_to_fit();
781        dedup_map.shrink_to_fit();
782        self.nodes = nodes;
783        self.dedup_map = dedup_map;
784        self.function_args = function_args;
785        self.array_elements = array_elements;
786        (remap, garbage)
787    }
788
789    /// Number of stored nodes.
790    pub(crate) fn node_count(&self) -> usize {
791        self.nodes.len()
792    }
793
794    /// Get statistics about the arena
795    pub fn stats(&self) -> AstArenaStats {
796        AstArenaStats {
797            node_count: self.nodes.len(),
798            dedup_hits: self.dedup_hits,
799            string_count: self.strings.len(),
800            table_spec_count: self.table_specs.len(),
801            total_args: self.function_args.len(),
802            total_array_elements: self.array_elements.len(),
803        }
804    }
805
806    /// Returns memory usage in bytes (approximate)
807    pub fn memory_usage(&self) -> usize {
808        self.nodes.capacity() * std::mem::size_of::<AstNodeEntry>()
809            + self.dedup_map.capacity() * (8 + 4) // hash + id
810            + self.function_args.capacity() * 4
811            + self.array_elements.capacity() * 4
812            + self.strings.memory_usage()
813            + self.table_specs.capacity() * std::mem::size_of::<TableSpecifier>()
814            + self.table_spec_dedup.capacity() * (8 + 4)
815    }
816
817    /// Clear all nodes from the arena
818    pub fn clear(&mut self) {
819        self.nodes.clear();
820        self.dedup_map.clear();
821        self.function_args.clear();
822        self.array_elements.clear();
823        self.strings.clear();
824        self.table_specs.clear();
825        self.table_spec_dedup.clear();
826        self.dedup_hits = 0;
827    }
828}
829
830impl Default for AstArena {
831    fn default() -> Self {
832        Self::new()
833    }
834}
835
836impl fmt::Debug for AstArena {
837    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
838        f.debug_struct("AstArena")
839            .field("nodes", &self.nodes.len())
840            .field("dedup_hits", &self.dedup_hits)
841            .field("strings", &self.strings.len())
842            .finish()
843    }
844}
845
846/// Statistics about the AST arena
847#[derive(Debug, Clone)]
848pub struct AstArenaStats {
849    pub node_count: usize,
850    pub dedup_hits: usize,
851    pub string_count: usize,
852    pub table_spec_count: usize,
853    pub total_args: usize,
854    pub total_array_elements: usize,
855}
856
857#[cfg(test)]
858mod tests {
859    use super::*;
860
861    #[test]
862    fn test_ast_arena_literal() {
863        let mut arena = AstArena::new();
864
865        let lit1 = arena.insert_literal(ValueRef::small_int(42).unwrap());
866        let lit2 = arena.insert_literal(ValueRef::boolean(true));
867
868        assert_ne!(lit1, lit2);
869
870        match arena.get(lit1) {
871            Some(AstNodeData::Literal(v)) => {
872                assert_eq!(v.as_small_int(), Some(42));
873            }
874            _ => panic!("Expected literal node"),
875        }
876    }
877
878    #[test]
879    fn test_ast_arena_deduplication() {
880        let mut arena = AstArena::new();
881
882        // Insert same literal twice
883        let lit1 = arena.insert_literal(ValueRef::small_int(42).unwrap());
884        let lit2 = arena.insert_literal(ValueRef::small_int(42).unwrap());
885
886        assert_eq!(lit1, lit2); // Should be deduplicated
887        assert_eq!(arena.stats().dedup_hits, 1);
888    }
889
890    #[test]
891    fn test_ast_arena_binary_op() {
892        let mut arena = AstArena::new();
893
894        let left = arena.insert_literal(ValueRef::small_int(1).unwrap());
895        let right = arena.insert_literal(ValueRef::small_int(2).unwrap());
896        let add = arena.insert_binary_op("+", left, right);
897
898        match arena.get(add) {
899            Some(AstNodeData::BinaryOp {
900                op_id,
901                left_id,
902                right_id,
903            }) => {
904                assert_eq!(arena.resolve_string(*op_id), "+");
905                assert_eq!(*left_id, left);
906                assert_eq!(*right_id, right);
907            }
908            _ => panic!("Expected binary op node"),
909        }
910    }
911
912    #[test]
913    fn test_ast_arena_function() {
914        let mut arena = AstArena::new();
915
916        let arg1 = arena.insert_literal(ValueRef::small_int(10).unwrap());
917        let arg2 = arena.insert_literal(ValueRef::small_int(20).unwrap());
918        let arg3 = arena.insert_literal(ValueRef::small_int(30).unwrap());
919
920        let func = arena.insert_function("SUM", vec![arg1, arg2, arg3]);
921
922        match arena.get(func) {
923            Some(AstNodeData::Function {
924                name_id,
925                args_count,
926                ..
927            }) => {
928                assert_eq!(arena.resolve_string(*name_id), "SUM");
929                assert_eq!(*args_count, 3);
930            }
931            _ => panic!("Expected function node"),
932        }
933
934        let args = arena.get_function_args(func).unwrap();
935        assert_eq!(args, &[arg1, arg2, arg3]);
936    }
937
938    #[test]
939    fn test_ast_arena_structural_sharing() {
940        let mut arena = AstArena::new();
941
942        // Create "A1" reference that will be shared
943        let a1_ref = arena.insert_reference(
944            "A1",
945            CompactRefType::Cell {
946                sheet: None,
947                row: 1,
948                col: 1,
949                row_abs: false,
950                col_abs: false,
951            },
952        );
953
954        // Create "A1 + 1"
955        let one = arena.insert_literal(ValueRef::small_int(1).unwrap());
956        let expr1 = arena.insert_binary_op("+", a1_ref, one);
957
958        // Create "A1 * 2"
959        let two = arena.insert_literal(ValueRef::small_int(2).unwrap());
960        let expr2 = arena.insert_binary_op("*", a1_ref, two);
961
962        // A1 reference should be shared
963        assert_eq!(arena.stats().node_count, 5); // A1, 1, +expr, 2, *expr
964
965        // Try to insert A1 again - should be deduplicated
966        let a1_ref2 = arena.insert_reference(
967            "A1",
968            CompactRefType::Cell {
969                sheet: None,
970                row: 1,
971                col: 1,
972                row_abs: false,
973                col_abs: false,
974            },
975        );
976        assert_eq!(a1_ref, a1_ref2);
977    }
978
979    #[test]
980    fn test_ast_arena_array() {
981        let mut arena = AstArena::new();
982
983        let elements = vec![
984            arena.insert_literal(ValueRef::small_int(1).unwrap()),
985            arena.insert_literal(ValueRef::small_int(2).unwrap()),
986            arena.insert_literal(ValueRef::small_int(3).unwrap()),
987            arena.insert_literal(ValueRef::small_int(4).unwrap()),
988        ];
989
990        let array = arena.insert_array(2, 2, elements.clone());
991
992        match arena.get(array) {
993            Some(AstNodeData::Array { rows, cols, .. }) => {
994                assert_eq!(*rows, 2);
995                assert_eq!(*cols, 2);
996            }
997            _ => panic!("Expected array node"),
998        }
999
1000        let stored_elements = arena.get_array_elements(array).unwrap();
1001        assert_eq!(stored_elements, &elements[..]);
1002    }
1003
1004    #[test]
1005    fn test_ast_arena_complex_expression() {
1006        let mut arena = AstArena::new();
1007
1008        // Build: SUM(A1:A10) + IF(B1 > 0, C1, D1)
1009
1010        // A1:A10 range
1011        let range = arena.insert_reference(
1012            "A1:A10",
1013            CompactRefType::Range {
1014                sheet: None,
1015                start_row: 1,
1016                start_col: 1,
1017                end_row: 10,
1018                end_col: 1,
1019                start_row_abs: false,
1020                start_col_abs: false,
1021                end_row_abs: false,
1022                end_col_abs: false,
1023            },
1024        );
1025
1026        // SUM(A1:A10)
1027        let sum = arena.insert_function("SUM", vec![range]);
1028
1029        // B1 reference
1030        let b1 = arena.insert_reference(
1031            "B1",
1032            CompactRefType::Cell {
1033                sheet: None,
1034                row: 1,
1035                col: 2,
1036                row_abs: false,
1037                col_abs: false,
1038            },
1039        );
1040
1041        // 0 literal
1042        let zero = arena.insert_literal(ValueRef::small_int(0).unwrap());
1043
1044        // B1 > 0
1045        let condition = arena.insert_binary_op(">", b1, zero);
1046
1047        // C1 and D1 references
1048        let c1 = arena.insert_reference(
1049            "C1",
1050            CompactRefType::Cell {
1051                sheet: None,
1052                row: 1,
1053                col: 3,
1054                row_abs: false,
1055                col_abs: false,
1056            },
1057        );
1058        let d1 = arena.insert_reference(
1059            "D1",
1060            CompactRefType::Cell {
1061                sheet: None,
1062                row: 1,
1063                col: 4,
1064                row_abs: false,
1065                col_abs: false,
1066            },
1067        );
1068
1069        // IF(B1 > 0, C1, D1)
1070        let if_expr = arena.insert_function("IF", vec![condition, c1, d1]);
1071
1072        // Final: SUM(...) + IF(...)
1073        let final_expr = arena.insert_binary_op("+", sum, if_expr);
1074
1075        // Verify structure
1076        assert!(arena.get(final_expr).is_some());
1077        // Note: zero literal gets deduplicated if used multiple times
1078        // We have: range, sum, b1, zero, condition(>), c1, d1, if_expr, final_expr(+)
1079        // That's 9 unique nodes (zero is deduplicated)
1080        assert_eq!(arena.stats().node_count, 9); // All unique nodes except deduplicated zero
1081    }
1082
1083    #[test]
1084    fn test_ast_arena_string_deduplication() {
1085        let mut arena = AstArena::new();
1086
1087        // Use same operator multiple times
1088        let one = arena.insert_literal(ValueRef::small_int(1).unwrap());
1089        let two = arena.insert_literal(ValueRef::small_int(2).unwrap());
1090        let three = arena.insert_literal(ValueRef::small_int(3).unwrap());
1091
1092        let add1 = arena.insert_binary_op("+", one, two);
1093        let add2 = arena.insert_binary_op("+", two, three);
1094        let add3 = arena.insert_binary_op("+", one, three);
1095
1096        // "+" should be interned only once
1097        assert_eq!(arena.strings().len(), 1);
1098    }
1099
1100    #[test]
1101    fn test_ast_arena_clear() {
1102        let mut arena = AstArena::new();
1103
1104        arena.insert_literal(ValueRef::small_int(1).unwrap());
1105        arena.insert_literal(ValueRef::small_int(2).unwrap());
1106        let left = arena.insert_literal(ValueRef::small_int(3).unwrap());
1107        let right = arena.insert_literal(ValueRef::small_int(4).unwrap());
1108        arena.insert_binary_op("+", left, right);
1109
1110        assert_eq!(arena.stats().node_count, 5);
1111
1112        arena.clear();
1113
1114        assert_eq!(arena.stats().node_count, 0);
1115        assert_eq!(arena.strings().len(), 0);
1116    }
1117}