1use 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#[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
40pub(crate) const CALL_NODE_NAME: &str = "#CALL";
47
48#[derive(Debug, Clone, PartialEq, Eq, Hash)]
50pub enum AstNodeData {
51 Literal(ValueRef),
53
54 Omitted,
56
57 Reference {
59 original_id: StringId, ref_type: CompactRefType, },
62
63 UnaryOp { op_id: StringId, expr_id: AstNodeId },
65
66 BinaryOp {
68 op_id: StringId,
69 left_id: AstNodeId,
70 right_id: AstNodeId,
71 },
72
73 Function {
75 name_id: StringId,
76 args_offset: u32, args_count: u16, },
79
80 Array {
82 rows: u16,
83 cols: u16,
84 elements_offset: u32, },
86}
87
88#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
90pub enum SheetKey {
91 Id(u16),
92 Name(StringId),
93}
94
95#[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 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 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#[derive(Debug, Clone, PartialEq, Eq, Hash)]
157pub(crate) struct AstNodeEntry {
158 pub(crate) data: AstNodeData,
159 pub(crate) meta: AstNodeMetadata,
160 pub(crate) subtotal_calls: u8,
166}
167
168pub(crate) const SUBTOTAL_CALL: u8 = 1;
170pub(crate) const AGGREGATE_CALL: u8 = 2;
172
173pub(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#[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#[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 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 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
273pub struct AstArena {
275 nodes: Vec<AstNodeEntry>,
277
278 dedup_map: FxHashMap<u64, AstNodeId>,
280
281 function_args: Vec<AstNodeId>,
283
284 array_elements: Vec<AstNodeId>,
286
287 strings: StringInterner,
289
290 table_specs: Vec<TableSpecifier>,
292 table_spec_dedup: FxHashMap<u64, TableSpecId>,
293
294 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), 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 pub fn insert(&mut self, node: AstNodeData) -> AstNodeId {
330 self.insert_entry(node, AstNodeMetadata::default())
331 }
332
333 #[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 let hash = self.hash_node(&node);
351
352 if let Some(&id) = self.dedup_map.get(&hash) {
354 if self.nodes[id.0 as usize].data == node {
356 self.dedup_hits += 1;
357 return id;
358 }
359 }
360
361 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 pub fn insert_literal(&mut self, value: ValueRef) -> AstNodeId {
375 self.insert(AstNodeData::Literal(value))
376 }
377
378 pub(crate) fn insert_omitted(&mut self) -> AstNodeId {
380 self.insert(AstNodeData::Omitted)
381 }
382
383 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 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 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 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 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 pub fn get(&self, id: AstNodeId) -> Option<&AstNodeData> {
446 self.nodes.get(id.0 as usize).map(|entry| &entry.data)
447 }
448
449 #[allow(dead_code)]
451 pub(crate) fn entry(&self, id: AstNodeId) -> Option<&AstNodeEntry> {
452 self.nodes.get(id.0 as usize)
453 }
454
455 #[allow(dead_code)]
457 pub(crate) fn metadata(&self, id: AstNodeId) -> Option<AstNodeMetadata> {
458 self.entry(id).map(|entry| entry.meta)
459 }
460
461 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 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 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 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 pub fn resolve_string(&self, id: StringId) -> &str {
547 self.strings.resolve(id)
548 }
549
550 pub fn strings(&self) -> &StringInterner {
552 &self.strings
553 }
554
555 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 fn hash_node(&self, node: &AstNodeData) -> u64 {
588 let mut hasher = DefaultHasher::new();
589 node.hash(&mut hasher);
590 hasher.finish()
591 }
592
593 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 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 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 pub(crate) fn node_count(&self) -> usize {
791 self.nodes.len()
792 }
793
794 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 pub fn memory_usage(&self) -> usize {
808 self.nodes.capacity() * std::mem::size_of::<AstNodeEntry>()
809 + self.dedup_map.capacity() * (8 + 4) + 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 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#[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 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); 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 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 let one = arena.insert_literal(ValueRef::small_int(1).unwrap());
956 let expr1 = arena.insert_binary_op("+", a1_ref, one);
957
958 let two = arena.insert_literal(ValueRef::small_int(2).unwrap());
960 let expr2 = arena.insert_binary_op("*", a1_ref, two);
961
962 assert_eq!(arena.stats().node_count, 5); 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 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 let sum = arena.insert_function("SUM", vec![range]);
1028
1029 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 let zero = arena.insert_literal(ValueRef::small_int(0).unwrap());
1043
1044 let condition = arena.insert_binary_op(">", b1, zero);
1046
1047 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 let if_expr = arena.insert_function("IF", vec![condition, c1, d1]);
1071
1072 let final_expr = arena.insert_binary_op("+", sum, if_expr);
1074
1075 assert!(arena.get(final_expr).is_some());
1077 assert_eq!(arena.stats().node_count, 9); }
1082
1083 #[test]
1084 fn test_ast_arena_string_deduplication() {
1085 let mut arena = AstArena::new();
1086
1087 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 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}