Skip to main content

antlr4_runtime/
prediction.rs

1use std::cmp::Ordering;
2use std::collections::{BTreeMap, BTreeSet, HashMap};
3use std::hash::{BuildHasherDefault, Hash, Hasher};
4use std::mem::size_of;
5
6pub const EMPTY_RETURN_STATE: usize = usize::MAX;
7const COMPACT_EMPTY_RETURN_STATE: u32 = u32::MAX;
8
9/// Lightweight `FxHash`-style hasher used on prediction hot paths.
10#[derive(Debug, Default)]
11pub struct PredictionFxHasher {
12    hash: u64,
13}
14
15const FX_ROT: u32 = 5;
16const FX_SEED: u64 = 0x51_7c_c1_b7_27_22_0a_95;
17
18impl Hasher for PredictionFxHasher {
19    #[inline]
20    fn write(&mut self, bytes: &[u8]) {
21        let mut bytes = bytes;
22        while bytes.len() >= 8 {
23            let (head, rest) = bytes.split_at(8);
24            let word = u64::from_le_bytes(head.try_into().expect("8-byte chunk"));
25            self.hash = (self.hash.rotate_left(FX_ROT) ^ word).wrapping_mul(FX_SEED);
26            bytes = rest;
27        }
28        for &byte in bytes {
29            self.hash = (self.hash.rotate_left(FX_ROT) ^ u64::from(byte)).wrapping_mul(FX_SEED);
30        }
31    }
32
33    #[inline]
34    fn write_u8(&mut self, value: u8) {
35        self.hash = (self.hash.rotate_left(FX_ROT) ^ u64::from(value)).wrapping_mul(FX_SEED);
36    }
37
38    #[inline]
39    fn write_u32(&mut self, value: u32) {
40        self.hash = (self.hash.rotate_left(FX_ROT) ^ u64::from(value)).wrapping_mul(FX_SEED);
41    }
42
43    #[inline]
44    fn write_u64(&mut self, value: u64) {
45        self.hash = (self.hash.rotate_left(FX_ROT) ^ value).wrapping_mul(FX_SEED);
46    }
47
48    #[inline]
49    fn write_usize(&mut self, value: usize) {
50        self.hash = (self.hash.rotate_left(FX_ROT) ^ value as u64).wrapping_mul(FX_SEED);
51    }
52
53    #[inline]
54    fn write_i32(&mut self, value: i32) {
55        self.write_u32(i32::cast_unsigned(value));
56    }
57
58    #[inline]
59    fn finish(&self) -> u64 {
60        self.hash
61    }
62}
63
64type FxHashMap<K, V> = HashMap<K, V, BuildHasherDefault<PredictionFxHasher>>;
65
66/// Store-local identity for one canonical prediction-context graph node.
67#[repr(transparent)]
68#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
69pub struct ContextId(u32);
70
71pub const EMPTY_CONTEXT: ContextId = ContextId(0);
72
73impl ContextId {
74    pub(crate) const fn compact(self) -> u32 {
75        self.0
76    }
77}
78
79#[derive(Clone, Copy, Debug, Eq, PartialEq)]
80enum ContextTag {
81    Empty,
82    Singleton,
83    Array,
84}
85
86#[derive(Clone, Copy, Debug)]
87struct ContextRecord {
88    tag: ContextTag,
89    cached_hash: u64,
90    parent_or_start: u32,
91    return_state_or_len: u32,
92}
93
94/// Allocation and interning totals for one prediction-context arena.
95#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
96pub struct PredictionContextStats {
97    pub contexts_created: usize,
98    pub singleton_contexts: usize,
99    pub array_contexts: usize,
100    pub array_entries: usize,
101    pub interner_hits: usize,
102    pub pooled_bytes: usize,
103    /// Element storage implied by retained capacities, excluding allocator and
104    /// hash-table control metadata.
105    pub retained_bytes: usize,
106    pub context_capacity: usize,
107    pub array_parent_capacity: usize,
108    pub array_return_state_capacity: usize,
109    pub interner_capacity: usize,
110    pub workspace_merge_cache_entries: usize,
111    pub workspace_merge_cache_capacity: usize,
112    pub workspace_entry_capacity: usize,
113    pub outer_context_cache_hits: usize,
114    pub outer_context_cache_misses: usize,
115}
116
117/// Canonical compact storage paired with one learned parser DFA store.
118#[derive(Debug)]
119pub(crate) struct ContextArena {
120    records: Vec<ContextRecord>,
121    array_parents: Vec<ContextId>,
122    array_return_states: Vec<u32>,
123    interner_heads: FxHashMap<u64, ContextId>,
124    interner_next: Vec<Option<ContextId>>,
125    interner_hits: usize,
126    #[cfg(debug_assertions)]
127    generation: u64,
128}
129
130#[cfg(debug_assertions)]
131fn next_context_arena_generation() -> u64 {
132    use std::sync::atomic::{AtomicU64, Ordering as AtomicOrdering};
133
134    static NEXT_GENERATION: AtomicU64 = AtomicU64::new(1);
135    NEXT_GENERATION.fetch_add(1, AtomicOrdering::Relaxed)
136}
137
138impl ContextArena {
139    pub(crate) fn new() -> Self {
140        let empty = ContextRecord {
141            tag: ContextTag::Empty,
142            cached_hash: prediction_context_empty_hash(),
143            parent_or_start: 0,
144            return_state_or_len: 0,
145        };
146        let mut interner_heads = FxHashMap::default();
147        interner_heads.insert(empty.cached_hash, EMPTY_CONTEXT);
148        Self {
149            records: vec![empty],
150            array_parents: Vec::new(),
151            array_return_states: Vec::new(),
152            interner_heads,
153            interner_next: vec![None],
154            interner_hits: 0,
155            #[cfg(debug_assertions)]
156            generation: next_context_arena_generation(),
157        }
158    }
159
160    #[cfg(debug_assertions)]
161    pub(crate) const fn generation(&self) -> u64 {
162        self.generation
163    }
164
165    pub(crate) fn stats(&self) -> PredictionContextStats {
166        let mut singleton_contexts = 0;
167        let mut array_contexts = 0;
168        for record in &self.records {
169            match record.tag {
170                ContextTag::Empty => {}
171                ContextTag::Singleton => singleton_contexts += 1,
172                ContextTag::Array => array_contexts += 1,
173            }
174        }
175        PredictionContextStats {
176            contexts_created: self.records.len(),
177            singleton_contexts,
178            array_contexts,
179            array_entries: self.array_parents.len(),
180            interner_hits: self.interner_hits,
181            pooled_bytes: self.records.len() * size_of::<ContextRecord>()
182                + self.array_parents.len() * size_of::<ContextId>()
183                + self.array_return_states.len() * size_of::<u32>()
184                + self.interner_next.len() * size_of::<Option<ContextId>>(),
185            retained_bytes: self.records.capacity() * size_of::<ContextRecord>()
186                + self.array_parents.capacity() * size_of::<ContextId>()
187                + self.array_return_states.capacity() * size_of::<u32>()
188                + self.interner_heads.capacity() * size_of::<(u64, ContextId)>()
189                + self.interner_next.capacity() * size_of::<Option<ContextId>>(),
190            context_capacity: self.records.capacity(),
191            array_parent_capacity: self.array_parents.capacity(),
192            array_return_state_capacity: self.array_return_states.capacity(),
193            interner_capacity: self.interner_heads.capacity(),
194            workspace_merge_cache_entries: 0,
195            workspace_merge_cache_capacity: 0,
196            workspace_entry_capacity: 0,
197            outer_context_cache_hits: 0,
198            outer_context_cache_misses: 0,
199        }
200    }
201
202    pub(crate) fn singleton(&mut self, parent: ContextId, return_state: usize) -> ContextId {
203        self.assert_valid(parent);
204        if return_state == EMPTY_RETURN_STATE {
205            return EMPTY_CONTEXT;
206        }
207        #[cfg(feature = "perf-counters")]
208        crate::perf::record_context_cache_call();
209        let return_state =
210            u32::try_from(return_state).expect("prediction return state must fit in u32");
211        let cached_hash = prediction_context_singleton_hash(self.cached_hash(parent), return_state);
212        if let Some(existing) = self.find_interned(cached_hash, |record| {
213            record.tag == ContextTag::Singleton
214                && record.parent_or_start == parent.0
215                && record.return_state_or_len == return_state
216        }) {
217            self.interner_hits = self.interner_hits.saturating_add(1);
218            #[cfg(feature = "perf-counters")]
219            crate::perf::record_context_cache_hit();
220            return existing;
221        }
222        #[cfg(feature = "perf-counters")]
223        {
224            crate::perf::record_context_cache_miss();
225            crate::perf::record_context_cache_insert();
226        }
227        self.push_record(ContextRecord {
228            tag: ContextTag::Singleton,
229            cached_hash,
230            parent_or_start: parent.0,
231            return_state_or_len: return_state,
232        })
233    }
234
235    fn intern_entries(&mut self, entries: &[(ContextId, u32)]) -> ContextId {
236        match entries {
237            [] => EMPTY_CONTEXT,
238            [(parent, return_state)] => {
239                if *return_state == COMPACT_EMPTY_RETURN_STATE {
240                    EMPTY_CONTEXT
241                } else {
242                    self.singleton(
243                        *parent,
244                        usize::try_from(*return_state).expect("u32 return state fits in usize"),
245                    )
246                }
247            }
248            _ => {
249                debug_assert!(
250                    entries
251                        .windows(2)
252                        .all(|pair| { compare_entries(pair[0], pair[1]) == Ordering::Less })
253                );
254                #[cfg(feature = "perf-counters")]
255                crate::perf::record_context_cache_call();
256                let cached_hash = prediction_context_array_hash(self, entries);
257                if let Some(existing) = self.find_interned(cached_hash, |record| {
258                    if record.tag != ContextTag::Array
259                        || usize::try_from(record.return_state_or_len).ok() != Some(entries.len())
260                    {
261                        return false;
262                    }
263                    let start =
264                        usize::try_from(record.parent_or_start).expect("u32 pool index fits usize");
265                    let end = start + entries.len();
266                    self.array_parents[start..end]
267                        .iter()
268                        .copied()
269                        .zip(self.array_return_states[start..end].iter().copied())
270                        .eq(entries.iter().copied())
271                }) {
272                    self.interner_hits = self.interner_hits.saturating_add(1);
273                    #[cfg(feature = "perf-counters")]
274                    crate::perf::record_context_cache_hit();
275                    return existing;
276                }
277                #[cfg(feature = "perf-counters")]
278                {
279                    crate::perf::record_context_cache_miss();
280                    crate::perf::record_context_cache_insert();
281                }
282                let start = u32::try_from(self.array_parents.len())
283                    .expect("prediction-context parent pool must fit in u32");
284                let len = u32::try_from(entries.len())
285                    .expect("prediction-context array length must fit in u32");
286                self.array_parents
287                    .extend(entries.iter().map(|(parent, _)| *parent));
288                self.array_return_states
289                    .extend(entries.iter().map(|(_, return_state)| *return_state));
290                self.push_record(ContextRecord {
291                    tag: ContextTag::Array,
292                    cached_hash,
293                    parent_or_start: start,
294                    return_state_or_len: len,
295                })
296            }
297        }
298    }
299
300    fn find_interned(
301        &self,
302        cached_hash: u64,
303        matches: impl Fn(&ContextRecord) -> bool,
304    ) -> Option<ContextId> {
305        let mut candidate = self.interner_heads.get(&cached_hash).copied();
306        while let Some(id) = candidate {
307            let index = usize::try_from(id.0).expect("u32 context ID fits in usize");
308            let record = &self.records[index];
309            if matches(record) {
310                return Some(id);
311            }
312            candidate = self.interner_next[index];
313        }
314        None
315    }
316
317    fn push_record(&mut self, record: ContextRecord) -> ContextId {
318        let id = ContextId(
319            u32::try_from(self.records.len()).expect("prediction-context arena must fit in u32"),
320        );
321        let previous = self.interner_heads.insert(record.cached_hash, id);
322        self.records.push(record);
323        self.interner_next.push(previous);
324        id
325    }
326
327    pub(crate) fn merge(
328        &mut self,
329        left: ContextId,
330        right: ContextId,
331        root_is_wildcard: bool,
332        workspace: &mut PredictionWorkspace,
333    ) -> ContextId {
334        self.assert_valid(left);
335        self.assert_valid(right);
336        #[cfg(feature = "perf-counters")]
337        crate::perf::record_context_merge_call();
338        if left == right {
339            #[cfg(feature = "perf-counters")]
340            crate::perf::record_context_merge_identical();
341            return left;
342        }
343        let key = MergeKey::new(left, right, root_is_wildcard);
344        if let Some(merged) = workspace.merge_cache.get(&key).copied() {
345            #[cfg(feature = "perf-counters")]
346            crate::perf::record_context_merge_cache_hit();
347            return merged;
348        }
349        #[cfg(feature = "perf-counters")]
350        {
351            crate::perf::record_context_merge_cache_miss();
352            crate::perf::record_context_merge_uncached();
353        }
354        let merged = if root_is_wildcard && (left == EMPTY_CONTEXT || right == EMPTY_CONTEXT) {
355            EMPTY_CONTEXT
356        } else {
357            self.merge_uncached(left, right, root_is_wildcard, workspace)
358        };
359        workspace.merge_cache.insert(key, merged);
360        merged
361    }
362
363    fn merge_uncached(
364        &mut self,
365        left: ContextId,
366        right: ContextId,
367        root_is_wildcard: bool,
368        workspace: &mut PredictionWorkspace,
369    ) -> ContextId {
370        match (self.tag(left), self.tag(right)) {
371            (ContextTag::Array, ContextTag::Array) => {
372                self.merge_arrays(left, right, root_is_wildcard, workspace)
373            }
374            (ContextTag::Array, _) => {
375                let entry = self.first_entry(right);
376                self.merge_array_with_entry(left, entry, root_is_wildcard, workspace)
377            }
378            (_, ContextTag::Array) => {
379                let entry = self.first_entry(left);
380                self.merge_array_with_entry(right, entry, root_is_wildcard, workspace)
381            }
382            _ => self.merge_two_entries(
383                self.first_entry(left),
384                self.first_entry(right),
385                root_is_wildcard,
386                workspace,
387            ),
388        }
389    }
390
391    fn merge_two_entries(
392        &mut self,
393        left: (ContextId, u32),
394        right: (ContextId, u32),
395        root_is_wildcard: bool,
396        workspace: &mut PredictionWorkspace,
397    ) -> ContextId {
398        if left.1 == right.1 {
399            let parent = if left.0 == right.0 {
400                left.0
401            } else {
402                self.merge(left.0, right.0, root_is_wildcard, workspace)
403            };
404            return self.intern_entries(&[(parent, left.1)]);
405        }
406
407        let start = workspace.entries.len();
408        if right.1 < left.1 {
409            workspace.entries.extend([right, left]);
410        } else {
411            workspace.entries.extend([left, right]);
412        }
413        self.intern_workspace_entries(workspace, start)
414    }
415
416    fn intern_workspace_entries(
417        &mut self,
418        workspace: &mut PredictionWorkspace,
419        start: usize,
420    ) -> ContextId {
421        let context = self.intern_entries(&workspace.entries[start..]);
422        workspace.entries.truncate(start);
423        context
424    }
425
426    fn merge_array_with_entry(
427        &mut self,
428        array: ContextId,
429        entry: (ContextId, u32),
430        root_is_wildcard: bool,
431        workspace: &mut PredictionWorkspace,
432    ) -> ContextId {
433        let array_len = self.len(array);
434        let mut insert_index = array_len;
435        for index in 0..array_len {
436            let current = self.entry(array, index).expect("array entry in range");
437            match entry.1.cmp(&current.1) {
438                Ordering::Less => {
439                    insert_index = index;
440                    break;
441                }
442                Ordering::Equal => {
443                    let parent = if entry.0 == current.0 {
444                        current.0
445                    } else {
446                        self.merge(entry.0, current.0, root_is_wildcard, workspace)
447                    };
448                    if parent == current.0 {
449                        return array;
450                    }
451
452                    let start = workspace.entries.len();
453                    for entry_index in 0..array_len {
454                        let array_entry = self
455                            .entry(array, entry_index)
456                            .expect("array entry in range");
457                        workspace.entries.push(if entry_index == index {
458                            (parent, current.1)
459                        } else {
460                            array_entry
461                        });
462                    }
463                    return self.intern_workspace_entries(workspace, start);
464                }
465                Ordering::Greater => {}
466            }
467        }
468
469        let start = workspace.entries.len();
470        for index in 0..insert_index {
471            workspace
472                .entries
473                .push(self.entry(array, index).expect("array entry in range"));
474        }
475        workspace.entries.push(entry);
476        for index in insert_index..array_len {
477            workspace
478                .entries
479                .push(self.entry(array, index).expect("array entry in range"));
480        }
481        self.intern_workspace_entries(workspace, start)
482    }
483
484    fn merge_arrays(
485        &mut self,
486        left: ContextId,
487        right: ContextId,
488        root_is_wildcard: bool,
489        workspace: &mut PredictionWorkspace,
490    ) -> ContextId {
491        let start = workspace.entries.len();
492        let left_len = self.len(left);
493        let right_len = self.len(right);
494        let mut left_index = 0;
495        let mut right_index = 0;
496        while left_index < left_len && right_index < right_len {
497            let left_entry = self.entry(left, left_index).expect("array entry in range");
498            let right_entry = self
499                .entry(right, right_index)
500                .expect("array entry in range");
501            match left_entry.1.cmp(&right_entry.1) {
502                Ordering::Less => {
503                    workspace.entries.push(left_entry);
504                    left_index += 1;
505                }
506                Ordering::Greater => {
507                    workspace.entries.push(right_entry);
508                    right_index += 1;
509                }
510                Ordering::Equal => {
511                    let parent = if left_entry.0 == right_entry.0 {
512                        left_entry.0
513                    } else {
514                        self.merge(left_entry.0, right_entry.0, root_is_wildcard, workspace)
515                    };
516                    workspace.entries.push((parent, left_entry.1));
517                    left_index += 1;
518                    right_index += 1;
519                }
520            }
521        }
522        while left_index < left_len {
523            workspace
524                .entries
525                .push(self.entry(left, left_index).expect("array entry in range"));
526            left_index += 1;
527        }
528        while right_index < right_len {
529            workspace.entries.push(
530                self.entry(right, right_index)
531                    .expect("array entry in range"),
532            );
533            right_index += 1;
534        }
535        self.intern_workspace_entries(workspace, start)
536    }
537
538    pub(crate) fn len(&self, context: ContextId) -> usize {
539        let record = self.record(context);
540        match record.tag {
541            ContextTag::Empty | ContextTag::Singleton => 1,
542            ContextTag::Array => usize::try_from(record.return_state_or_len)
543                .expect("u32 context length fits in usize"),
544        }
545    }
546
547    pub(crate) fn is_empty(&self, context: ContextId) -> bool {
548        self.assert_valid(context);
549        context == EMPTY_CONTEXT
550    }
551
552    pub(crate) fn has_empty_path(&self, context: ContextId) -> bool {
553        if context == EMPTY_CONTEXT {
554            return true;
555        }
556        let record = self.record(context);
557        match record.tag {
558            ContextTag::Empty => true,
559            ContextTag::Singleton => false,
560            ContextTag::Array => {
561                let len = usize::try_from(record.return_state_or_len)
562                    .expect("u32 context length fits in usize");
563                let start = usize::try_from(record.parent_or_start)
564                    .expect("u32 context pool index fits in usize");
565                self.array_return_states[start + len - 1] == COMPACT_EMPTY_RETURN_STATE
566            }
567        }
568    }
569
570    pub(crate) fn return_state(&self, context: ContextId, index: usize) -> Option<usize> {
571        let (_, return_state) = self.entry(context, index)?;
572        Some(expand_return_state(return_state))
573    }
574
575    pub(crate) fn parent(&self, context: ContextId, index: usize) -> Option<ContextId> {
576        if context == EMPTY_CONTEXT {
577            self.assert_valid(context);
578            return None;
579        }
580        self.entry(context, index).map(|(parent, _)| parent)
581    }
582
583    fn first_entry(&self, context: ContextId) -> (ContextId, u32) {
584        self.entry(context, 0)
585            .expect("empty and singleton contexts have one logical entry")
586    }
587
588    fn entry(&self, context: ContextId, index: usize) -> Option<(ContextId, u32)> {
589        let record = self.record(context);
590        match record.tag {
591            ContextTag::Empty if index == 0 => Some((EMPTY_CONTEXT, COMPACT_EMPTY_RETURN_STATE)),
592            ContextTag::Singleton if index == 0 => Some((
593                ContextId(record.parent_or_start),
594                record.return_state_or_len,
595            )),
596            ContextTag::Array => {
597                let len = usize::try_from(record.return_state_or_len).ok()?;
598                if index >= len {
599                    return None;
600                }
601                let start = usize::try_from(record.parent_or_start).ok()?;
602                Some((
603                    self.array_parents[start + index],
604                    self.array_return_states[start + index],
605                ))
606            }
607            ContextTag::Empty | ContextTag::Singleton => None,
608        }
609    }
610
611    fn tag(&self, context: ContextId) -> ContextTag {
612        self.record(context).tag
613    }
614
615    fn cached_hash(&self, context: ContextId) -> u64 {
616        self.record(context).cached_hash
617    }
618
619    fn record(&self, context: ContextId) -> &ContextRecord {
620        self.assert_valid(context);
621        &self.records[usize::try_from(context.0).expect("u32 context ID fits in usize")]
622    }
623
624    pub(crate) fn assert_valid(&self, context: ContextId) {
625        assert!(
626            usize::try_from(context.0).is_ok_and(|index| index < self.records.len()),
627            "prediction ContextId does not belong to this store"
628        );
629    }
630
631    pub(crate) fn import_all(
632        &mut self,
633        source: &Self,
634        workspace: &mut PredictionWorkspace,
635    ) -> Vec<ContextId> {
636        workspace.entries.clear();
637        let mut remap = Vec::with_capacity(source.records.len());
638        remap.push(EMPTY_CONTEXT);
639        for source_index in 1..source.records.len() {
640            let source_id = ContextId(
641                u32::try_from(source_index).expect("source prediction-context ID fits in u32"),
642            );
643            let imported = match source.tag(source_id) {
644                ContextTag::Empty => EMPTY_CONTEXT,
645                ContextTag::Singleton => {
646                    let (parent, return_state) = source.first_entry(source_id);
647                    let parent_index =
648                        usize::try_from(parent.0).expect("u32 context ID fits usize");
649                    assert!(
650                        parent_index < remap.len(),
651                        "prediction contexts must reference earlier arena records"
652                    );
653                    self.singleton(remap[parent_index], expand_return_state(return_state))
654                }
655                ContextTag::Array => {
656                    let start = workspace.entries.len();
657                    for entry_index in 0..source.len(source_id) {
658                        let (parent, return_state) = source
659                            .entry(source_id, entry_index)
660                            .expect("source array entry in range");
661                        let parent_index =
662                            usize::try_from(parent.0).expect("u32 context ID fits usize");
663                        assert!(
664                            parent_index < remap.len(),
665                            "prediction contexts must reference earlier arena records"
666                        );
667                        workspace.entries.push((remap[parent_index], return_state));
668                    }
669                    workspace
670                        .entries
671                        .sort_unstable_by(|left, right| compare_entries(*left, *right));
672                    workspace.entries.dedup();
673                    self.intern_workspace_entries(workspace, start)
674                }
675            };
676            remap.push(imported);
677        }
678        remap
679    }
680}
681
682impl Default for ContextArena {
683    fn default() -> Self {
684        Self::new()
685    }
686}
687
688fn compare_entries(left: (ContextId, u32), right: (ContextId, u32)) -> Ordering {
689    left.1.cmp(&right.1)
690}
691
692fn expand_return_state(return_state: u32) -> usize {
693    if return_state == COMPACT_EMPTY_RETURN_STATE {
694        EMPTY_RETURN_STATE
695    } else {
696        usize::try_from(return_state).expect("u32 return state fits in usize")
697    }
698}
699
700fn prediction_context_empty_hash() -> u64 {
701    let mut hasher = PredictionFxHasher::default();
702    hasher.write_u8(0);
703    hasher.finish()
704}
705
706fn prediction_context_singleton_hash(parent_hash: u64, return_state: u32) -> u64 {
707    let mut hasher = PredictionFxHasher::default();
708    hasher.write_u8(1);
709    hasher.write_u64(parent_hash);
710    hasher.write_u32(return_state);
711    hasher.finish()
712}
713
714fn prediction_context_array_hash(arena: &ContextArena, entries: &[(ContextId, u32)]) -> u64 {
715    let mut hasher = PredictionFxHasher::default();
716    hasher.write_u8(2);
717    hasher.write_usize(entries.len());
718    for (parent, _) in entries {
719        hasher.write_u64(arena.cached_hash(*parent));
720    }
721    hasher.write_usize(entries.len());
722    for (_, return_state) in entries {
723        hasher.write_u32(*return_state);
724    }
725    hasher.finish()
726}
727
728#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
729struct MergeKey {
730    left: ContextId,
731    right: ContextId,
732    root_is_wildcard: bool,
733}
734
735impl MergeKey {
736    fn new(left: ContextId, right: ContextId, root_is_wildcard: bool) -> Self {
737        let (left, right) = if right < left {
738            (right, left)
739        } else {
740            (left, right)
741        };
742        Self {
743            left,
744            right,
745            root_is_wildcard,
746        }
747    }
748}
749
750const MAX_RETAINED_MERGE_CACHE_ENTRIES: usize = 65_536;
751const MAX_RETAINED_CONTEXT_ENTRIES: usize = 16_384;
752
753/// Reusable per-prediction merge cache and temporary compact entry storage.
754#[derive(Debug, Default)]
755pub(crate) struct PredictionWorkspace {
756    merge_cache: FxHashMap<MergeKey, ContextId>,
757    entries: Vec<(ContextId, u32)>,
758}
759
760impl PredictionWorkspace {
761    pub(crate) fn reset(&mut self) {
762        if self.merge_cache.capacity() > MAX_RETAINED_MERGE_CACHE_ENTRIES {
763            self.merge_cache = FxHashMap::default();
764        } else {
765            self.merge_cache.clear();
766        }
767        if self.entries.capacity() > MAX_RETAINED_CONTEXT_ENTRIES {
768            self.entries = Vec::new();
769        } else {
770            self.entries.clear();
771        }
772    }
773
774    pub(crate) fn merge_cache_capacity(&self) -> usize {
775        self.merge_cache.capacity()
776    }
777
778    pub(crate) fn merge_cache_len(&self) -> usize {
779        self.merge_cache.len()
780    }
781
782    pub(crate) const fn entry_capacity(&self) -> usize {
783        self.entries.capacity()
784    }
785
786    pub(crate) fn retained_bytes(&self) -> usize {
787        self.merge_cache.capacity() * size_of::<(MergeKey, ContextId)>()
788            + self.entries.capacity() * size_of::<(ContextId, u32)>()
789    }
790}
791
792#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
793pub enum SemanticContext {
794    None,
795    Predicate {
796        rule_index: usize,
797        pred_index: usize,
798        context_dependent: bool,
799    },
800    Precedence {
801        precedence: i32,
802    },
803    And(Vec<Self>),
804    Or(Vec<Self>),
805}
806
807impl SemanticContext {
808    pub const fn none() -> Self {
809        Self::None
810    }
811
812    pub fn and(left: Self, right: Self) -> Self {
813        combine_semantic_context(left, right, true)
814    }
815
816    pub fn or(left: Self, right: Self) -> Self {
817        combine_semantic_context(left, right, false)
818    }
819
820    pub const fn is_none(&self) -> bool {
821        matches!(self, Self::None)
822    }
823}
824
825fn combine_semantic_context(
826    left: SemanticContext,
827    right: SemanticContext,
828    and: bool,
829) -> SemanticContext {
830    if left == right {
831        return left;
832    }
833    if left.is_none() {
834        return right;
835    }
836    if right.is_none() {
837        return left;
838    }
839    let mut entries = Vec::new();
840    for context in [left, right] {
841        match (and, context) {
842            (true, SemanticContext::And(children)) | (false, SemanticContext::Or(children)) => {
843                entries.extend(children);
844            }
845            (_, other) => entries.push(other),
846        }
847    }
848    entries.sort();
849    entries.dedup();
850    if and {
851        SemanticContext::And(entries)
852    } else {
853        SemanticContext::Or(entries)
854    }
855}
856
857#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
858pub(crate) struct PredictionRuleCall {
859    pub(crate) source_state: usize,
860    pub(crate) rule_index: usize,
861}
862
863#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
864pub(crate) struct PredictionPredicateCall {
865    pub(crate) rule_index: usize,
866    pub(crate) pred_index: usize,
867    pub(crate) rule_calls: Vec<PredictionRuleCall>,
868}
869
870#[derive(Clone, Debug, Default, Eq, Hash, Ord, PartialEq, PartialOrd)]
871struct PredictionSemanticProvenance {
872    active_rule_calls: Vec<PredictionRuleCall>,
873    predicate_calls: Vec<PredictionPredicateCall>,
874}
875
876#[repr(transparent)]
877#[derive(Clone, Copy, Debug, Default, Eq, Hash, Ord, PartialEq, PartialOrd)]
878pub(crate) struct PredictionSemanticProvenanceId(u32);
879
880#[derive(Debug, Default)]
881pub(crate) struct PredictionSemanticProvenanceArena {
882    records: Vec<PredictionSemanticProvenance>,
883    interner_heads: FxHashMap<u64, PredictionSemanticProvenanceId>,
884    interner_next: Vec<Option<PredictionSemanticProvenanceId>>,
885}
886
887impl PredictionSemanticProvenanceArena {
888    pub(crate) fn enter_rule(
889        &mut self,
890        id: PredictionSemanticProvenanceId,
891        source_state: usize,
892        rule_index: usize,
893    ) -> PredictionSemanticProvenanceId {
894        let mut provenance = self.get(id).cloned().unwrap_or_default();
895        provenance.active_rule_calls.push(PredictionRuleCall {
896            source_state,
897            rule_index,
898        });
899        self.intern(provenance)
900    }
901
902    pub(crate) fn exit_rule(
903        &mut self,
904        id: PredictionSemanticProvenanceId,
905    ) -> PredictionSemanticProvenanceId {
906        let Some(mut provenance) = self.get(id).cloned() else {
907            return PredictionSemanticProvenanceId::default();
908        };
909        provenance.active_rule_calls.pop();
910        self.intern(provenance)
911    }
912
913    pub(crate) fn record_predicate(
914        &mut self,
915        id: PredictionSemanticProvenanceId,
916        rule_index: usize,
917        pred_index: usize,
918    ) -> PredictionSemanticProvenanceId {
919        let mut provenance = self.get(id).cloned().unwrap_or_default();
920        let call = PredictionPredicateCall {
921            rule_index,
922            pred_index,
923            rule_calls: provenance.active_rule_calls.clone(),
924        };
925        if !provenance.predicate_calls.contains(&call) {
926            provenance.predicate_calls.push(call);
927        }
928        self.intern(provenance)
929    }
930
931    pub(crate) fn predicate_calls(
932        &self,
933        id: PredictionSemanticProvenanceId,
934    ) -> &[PredictionPredicateCall] {
935        self.get(id)
936            .map_or(&[], |provenance| provenance.predicate_calls.as_slice())
937    }
938
939    fn get(&self, id: PredictionSemanticProvenanceId) -> Option<&PredictionSemanticProvenance> {
940        let index = id.0.checked_sub(1)?;
941        self.records.get(usize::try_from(index).ok()?)
942    }
943
944    fn find_interned(
945        &self,
946        cached_hash: u64,
947        provenance: &PredictionSemanticProvenance,
948    ) -> Option<PredictionSemanticProvenanceId> {
949        let mut candidate = self.interner_heads.get(&cached_hash).copied();
950        while let Some(id) = candidate {
951            let index = usize::try_from(id.0.checked_sub(1)?).ok()?;
952            if self.records.get(index) == Some(provenance) {
953                return Some(id);
954            }
955            candidate = self.interner_next.get(index).copied().flatten();
956        }
957        None
958    }
959
960    fn intern(
961        &mut self,
962        provenance: PredictionSemanticProvenance,
963    ) -> PredictionSemanticProvenanceId {
964        if provenance.active_rule_calls.is_empty() && provenance.predicate_calls.is_empty() {
965            return PredictionSemanticProvenanceId::default();
966        }
967        let mut hasher = PredictionFxHasher::default();
968        provenance.hash(&mut hasher);
969        let cached_hash = hasher.finish();
970        if let Some(id) = self.find_interned(cached_hash, &provenance) {
971            return id;
972        }
973        let id = PredictionSemanticProvenanceId(
974            u32::try_from(self.records.len() + 1)
975                .expect("prediction semantic provenance arena exhausted"),
976        );
977        assert!(
978            id.0 <= ATN_CONFIG_PROVENANCE_MASK,
979            "prediction semantic provenance arena exhausted"
980        );
981        let previous = self.interner_heads.insert(cached_hash, id);
982        self.records.push(provenance);
983        self.interner_next.push(previous);
984        id
985    }
986}
987
988const ATN_CONFIG_PRECEDENCE_FILTER_SUPPRESSED: u32 = 1 << 31;
989const ATN_CONFIG_PROVENANCE_MASK: u32 = !ATN_CONFIG_PRECEDENCE_FILTER_SUPPRESSED;
990
991#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
992pub(crate) struct AtnConfig {
993    pub(crate) state: usize,
994    pub(crate) alt: usize,
995    pub(crate) context: ContextId,
996    pub(crate) semantic_context: SemanticContext,
997    pub(crate) reaches_into_outer_context: usize,
998    semantic_provenance_and_flags: u32,
999    #[cfg(debug_assertions)]
1000    context_generation: u64,
1001}
1002
1003impl AtnConfig {
1004    pub(crate) fn new(state: usize, alt: usize, context: ContextId, arena: &ContextArena) -> Self {
1005        arena.assert_valid(context);
1006        Self {
1007            state,
1008            alt,
1009            context,
1010            semantic_context: SemanticContext::None,
1011            reaches_into_outer_context: 0,
1012            semantic_provenance_and_flags: 0,
1013            #[cfg(debug_assertions)]
1014            context_generation: arena.generation(),
1015        }
1016    }
1017
1018    #[must_use]
1019    #[cfg(test)]
1020    pub(crate) fn with_semantic_context(mut self, semantic_context: SemanticContext) -> Self {
1021        self.semantic_context = semantic_context;
1022        self
1023    }
1024
1025    pub(crate) fn set_context(&mut self, context: ContextId, arena: &ContextArena) {
1026        arena.assert_valid(context);
1027        self.context = context;
1028        #[cfg(debug_assertions)]
1029        {
1030            self.context_generation = arena.generation();
1031        }
1032    }
1033
1034    pub(crate) fn moved_to(&self, state: usize, context: ContextId, arena: &ContextArena) -> Self {
1035        let mut moved = Self::new(state, self.alt, context, arena);
1036        moved.semantic_context = self.semantic_context.clone();
1037        moved.reaches_into_outer_context = self.reaches_into_outer_context;
1038        moved.semantic_provenance_and_flags = self.semantic_provenance_and_flags;
1039        moved
1040    }
1041
1042    pub(crate) fn enter_prediction_rule(
1043        &mut self,
1044        arena: &mut PredictionSemanticProvenanceArena,
1045        source_state: usize,
1046        rule_index: usize,
1047    ) {
1048        let id = arena.enter_rule(self.semantic_provenance_id(), source_state, rule_index);
1049        self.set_semantic_provenance_id(id);
1050    }
1051
1052    pub(crate) fn exit_prediction_rule(&mut self, arena: &mut PredictionSemanticProvenanceArena) {
1053        let id = arena.exit_rule(self.semantic_provenance_id());
1054        self.set_semantic_provenance_id(id);
1055    }
1056
1057    pub(crate) fn record_prediction_predicate(
1058        &mut self,
1059        arena: &mut PredictionSemanticProvenanceArena,
1060        rule_index: usize,
1061        pred_index: usize,
1062    ) {
1063        let id = arena.record_predicate(self.semantic_provenance_id(), rule_index, pred_index);
1064        self.set_semantic_provenance_id(id);
1065    }
1066
1067    pub(crate) const fn semantic_provenance_id(&self) -> PredictionSemanticProvenanceId {
1068        PredictionSemanticProvenanceId(
1069            self.semantic_provenance_and_flags & ATN_CONFIG_PROVENANCE_MASK,
1070        )
1071    }
1072
1073    pub(crate) const fn semantic_provenance_and_flags(&self) -> u32 {
1074        self.semantic_provenance_and_flags
1075    }
1076
1077    fn set_semantic_provenance_id(&mut self, id: PredictionSemanticProvenanceId) {
1078        debug_assert_eq!(id.0 & ATN_CONFIG_PRECEDENCE_FILTER_SUPPRESSED, 0);
1079        self.semantic_provenance_and_flags =
1080            (self.semantic_provenance_and_flags & ATN_CONFIG_PRECEDENCE_FILTER_SUPPRESSED) | id.0;
1081    }
1082
1083    pub(crate) const fn precedence_filter_suppressed(&self) -> bool {
1084        self.semantic_provenance_and_flags & ATN_CONFIG_PRECEDENCE_FILTER_SUPPRESSED != 0
1085    }
1086
1087    pub(crate) const fn suppress_precedence_filter(&mut self) {
1088        self.semantic_provenance_and_flags |= ATN_CONFIG_PRECEDENCE_FILTER_SUPPRESSED;
1089    }
1090
1091    pub(crate) const fn merge_precedence_filter_suppression(&mut self, other: &Self) {
1092        if other.precedence_filter_suppressed() {
1093            self.suppress_precedence_filter();
1094        }
1095    }
1096
1097    pub(crate) fn assert_store(&self, arena: &ContextArena) {
1098        arena.assert_valid(self.context);
1099        #[cfg(debug_assertions)]
1100        assert_eq!(
1101            self.context_generation,
1102            arena.generation(),
1103            "ATN config carries a ContextId from another prediction store"
1104        );
1105    }
1106}
1107
1108#[derive(Clone, Debug, Default)]
1109pub(crate) struct AtnConfigSet {
1110    configs: Vec<AtnConfig>,
1111    config_index: FxHashMap<AtnConfigKey, usize>,
1112    full_context: bool,
1113    unique_alt: Option<usize>,
1114    conflicting_alts: BTreeSet<usize>,
1115    has_semantic_context: bool,
1116    dips_into_outer_context: bool,
1117    readonly: bool,
1118}
1119
1120impl AtnConfigSet {
1121    pub(crate) fn new() -> Self {
1122        Self::default()
1123    }
1124
1125    pub(crate) fn new_full_context(full_context: bool) -> Self {
1126        Self {
1127            configs: Vec::new(),
1128            config_index: FxHashMap::default(),
1129            full_context,
1130            unique_alt: None,
1131            conflicting_alts: BTreeSet::new(),
1132            has_semantic_context: false,
1133            dips_into_outer_context: false,
1134            readonly: false,
1135        }
1136    }
1137
1138    /// Adds a configuration, merging contexts for equivalent config keys.
1139    pub(crate) fn add(
1140        &mut self,
1141        config: AtnConfig,
1142        arena: &mut ContextArena,
1143        workspace: &mut PredictionWorkspace,
1144    ) -> bool {
1145        assert!(!self.readonly, "cannot mutate readonly ATN config set");
1146        config.assert_store(arena);
1147        #[cfg(feature = "perf-counters")]
1148        crate::perf::record_config_add_call();
1149        if !config.semantic_context.is_none() {
1150            self.has_semantic_context = true;
1151        }
1152        if config.reaches_into_outer_context > 0 {
1153            self.dips_into_outer_context = true;
1154        }
1155        let key = AtnConfigKey::from(&config);
1156        if let Some(existing_index) = self.config_index.get(&key).copied() {
1157            #[cfg(feature = "perf-counters")]
1158            crate::perf::record_config_merge();
1159            let existing = &mut self.configs[existing_index];
1160            existing.assert_store(arena);
1161            existing.context = arena.merge(
1162                existing.context,
1163                config.context,
1164                !self.full_context,
1165                workspace,
1166            );
1167            existing.reaches_into_outer_context = existing
1168                .reaches_into_outer_context
1169                .max(config.reaches_into_outer_context);
1170            existing.merge_precedence_filter_suppression(&config);
1171            self.conflicting_alts.clear();
1172            false
1173        } else {
1174            let index = self.configs.len();
1175            self.config_index.insert(key, index);
1176            self.configs.push(config);
1177            #[cfg(feature = "perf-counters")]
1178            crate::perf::record_config_insert(self.configs.len());
1179            self.unique_alt = None;
1180            self.conflicting_alts.clear();
1181            true
1182        }
1183    }
1184
1185    pub(crate) fn configs(&self) -> &[AtnConfig] {
1186        &self.configs
1187    }
1188
1189    pub(crate) fn into_configs(self) -> Vec<AtnConfig> {
1190        self.configs
1191    }
1192
1193    pub(crate) const fn is_empty(&self) -> bool {
1194        self.configs.is_empty()
1195    }
1196
1197    pub(crate) const fn len(&self) -> usize {
1198        self.configs.len()
1199    }
1200
1201    pub(crate) fn set_readonly(&mut self, readonly: bool) {
1202        self.readonly = readonly;
1203        if readonly {
1204            self.config_index = FxHashMap::default();
1205            self.conflicting_alts.clear();
1206        }
1207    }
1208
1209    pub(crate) const fn full_context(&self) -> bool {
1210        self.full_context
1211    }
1212
1213    pub(crate) const fn has_semantic_context(&self) -> bool {
1214        self.has_semantic_context
1215    }
1216
1217    pub(crate) fn unique_alt(&mut self) -> Option<usize> {
1218        if self.unique_alt.is_none() {
1219            self.unique_alt = unique_alt(self.configs());
1220        }
1221        self.unique_alt
1222    }
1223
1224    pub(crate) fn alts(&self) -> BTreeSet<usize> {
1225        self.configs.iter().map(|config| config.alt).collect()
1226    }
1227
1228    pub(crate) fn conflicting_alt_subsets(&self) -> Vec<BTreeSet<usize>> {
1229        conflicting_alt_subsets(self.configs())
1230    }
1231
1232    pub(crate) fn conflicting_alts(&mut self) -> BTreeSet<usize> {
1233        if self.conflicting_alts.is_empty() {
1234            self.conflicting_alts = self
1235                .conflicting_alt_subsets()
1236                .into_iter()
1237                .filter(|alts| alts.len() > 1)
1238                .flatten()
1239                .collect();
1240        }
1241        self.conflicting_alts.clone()
1242    }
1243
1244    pub(crate) fn remap_contexts(&mut self, remap: &[ContextId], arena: &ContextArena) {
1245        for config in &mut self.configs {
1246            let index = usize::try_from(config.context.0).expect("u32 context ID fits usize");
1247            config.set_context(
1248                *remap
1249                    .get(index)
1250                    .expect("every imported context ID has a remap"),
1251                arena,
1252            );
1253        }
1254        self.config_index.clear();
1255        if !self.readonly {
1256            for (index, config) in self.configs.iter().enumerate() {
1257                self.config_index.insert(AtnConfigKey::from(config), index);
1258            }
1259        }
1260    }
1261
1262    pub(crate) fn fingerprint(&self) -> u64 {
1263        let mut hasher = PredictionFxHasher::default();
1264        self.configs.hash(&mut hasher);
1265        self.full_context.hash(&mut hasher);
1266        self.has_semantic_context.hash(&mut hasher);
1267        self.dips_into_outer_context.hash(&mut hasher);
1268        hasher.finish()
1269    }
1270
1271    pub(crate) fn retained_bytes(&self) -> usize {
1272        self.configs.capacity() * size_of::<AtnConfig>()
1273            + self.config_index.capacity() * size_of::<(AtnConfigKey, usize)>()
1274    }
1275}
1276
1277impl PartialEq for AtnConfigSet {
1278    fn eq(&self, other: &Self) -> bool {
1279        self.configs == other.configs
1280            && self.full_context == other.full_context
1281            && self.has_semantic_context == other.has_semantic_context
1282            && self.dips_into_outer_context == other.dips_into_outer_context
1283    }
1284}
1285
1286impl Eq for AtnConfigSet {}
1287
1288impl Ord for AtnConfigSet {
1289    fn cmp(&self, other: &Self) -> Ordering {
1290        self.configs
1291            .cmp(&other.configs)
1292            .then_with(|| self.full_context.cmp(&other.full_context))
1293            .then_with(|| self.has_semantic_context.cmp(&other.has_semantic_context))
1294            .then_with(|| {
1295                self.dips_into_outer_context
1296                    .cmp(&other.dips_into_outer_context)
1297            })
1298    }
1299}
1300
1301impl PartialOrd for AtnConfigSet {
1302    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
1303        Some(self.cmp(other))
1304    }
1305}
1306
1307#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
1308struct AtnConfigKey {
1309    state: usize,
1310    alt: usize,
1311    semantic_context: SemanticContext,
1312    semantic_provenance: PredictionSemanticProvenanceId,
1313}
1314
1315impl From<&AtnConfig> for AtnConfigKey {
1316    fn from(config: &AtnConfig) -> Self {
1317        Self {
1318            state: config.state,
1319            alt: config.alt,
1320            semantic_context: config.semantic_context.clone(),
1321            semantic_provenance: config.semantic_provenance_id(),
1322        }
1323    }
1324}
1325
1326pub(crate) fn unique_alt(configs: &[AtnConfig]) -> Option<usize> {
1327    let mut alt = None;
1328    for config in configs {
1329        match alt {
1330            None => alt = Some(config.alt),
1331            Some(existing) if existing == config.alt => {}
1332            Some(_) => return None,
1333        }
1334    }
1335    alt
1336}
1337
1338pub(crate) fn conflicting_alt_subsets(configs: &[AtnConfig]) -> Vec<BTreeSet<usize>> {
1339    let mut by_state_context = FxHashMap::<(usize, ContextId), BTreeSet<usize>>::default();
1340    for config in configs {
1341        by_state_context
1342            .entry((config.state, config.context))
1343            .or_default()
1344            .insert(config.alt);
1345    }
1346    by_state_context.into_values().collect()
1347}
1348
1349pub(crate) fn all_subsets_conflict(alt_subsets: &[BTreeSet<usize>]) -> bool {
1350    alt_subsets.iter().all(|alts| alts.len() > 1)
1351}
1352
1353pub(crate) fn all_subsets_equal(alt_subsets: &[BTreeSet<usize>]) -> bool {
1354    let mut subsets = alt_subsets.iter();
1355    let Some(first) = subsets.next() else {
1356        return true;
1357    };
1358    subsets.all(|alts| alts == first)
1359}
1360
1361pub(crate) fn single_viable_alt(alt_subsets: &[BTreeSet<usize>]) -> Option<usize> {
1362    let mut result = None;
1363    for alts in alt_subsets {
1364        let min_alt = alts.iter().next().copied()?;
1365        match result {
1366            None => result = Some(min_alt),
1367            Some(existing) if existing == min_alt => {}
1368            Some(_) => return None,
1369        }
1370    }
1371    result
1372}
1373
1374pub(crate) fn has_sll_conflict_terminating_prediction(
1375    configs: &AtnConfigSet,
1376    is_rule_stop_state: impl Fn(usize) -> bool,
1377) -> bool {
1378    if configs
1379        .configs()
1380        .iter()
1381        .all(|config| is_rule_stop_state(config.state))
1382    {
1383        return true;
1384    }
1385    let alt_subsets = configs.conflicting_alt_subsets();
1386    alt_subsets.iter().any(|alts| alts.len() > 1)
1387        && !has_state_associated_with_one_alt(configs.configs())
1388}
1389
1390fn has_state_associated_with_one_alt(configs: &[AtnConfig]) -> bool {
1391    let mut by_state = BTreeMap::<usize, BTreeSet<usize>>::new();
1392    for config in configs {
1393        by_state.entry(config.state).or_default().insert(config.alt);
1394    }
1395    by_state.values().any(|alts| alts.len() == 1)
1396}
1397
1398#[cfg(test)]
1399mod tests {
1400    use super::*;
1401
1402    #[test]
1403    fn arena_interns_singletons_without_per_context_objects() {
1404        let mut arena = ContextArena::new();
1405        let first = arena.singleton(EMPTY_CONTEXT, 7);
1406        let second = arena.singleton(EMPTY_CONTEXT, 7);
1407
1408        assert_eq!(first, second);
1409        assert_eq!(arena.stats().singleton_contexts, 1);
1410        assert_eq!(arena.stats().interner_hits, 1);
1411    }
1412
1413    #[test]
1414    fn array_interner_verifies_payload_after_hash_collision() {
1415        let mut arena = ContextArena::new();
1416        let first_parent = arena.singleton(EMPTY_CONTEXT, 1);
1417        let second_parent = arena.singleton(EMPTY_CONTEXT, 2);
1418        let expected = [(first_parent, 10), (second_parent, 20)];
1419        let colliding = [(second_parent, 10), (first_parent, 20)];
1420        let cached_hash = prediction_context_array_hash(&arena, &expected);
1421        let start = u32::try_from(arena.array_parents.len()).expect("pool index fits u32");
1422        arena
1423            .array_parents
1424            .extend(colliding.iter().map(|(parent, _)| *parent));
1425        arena
1426            .array_return_states
1427            .extend(colliding.iter().map(|(_, return_state)| *return_state));
1428        let collision = arena.push_record(ContextRecord {
1429            tag: ContextTag::Array,
1430            cached_hash,
1431            parent_or_start: start,
1432            return_state_or_len: 2,
1433        });
1434
1435        let interned = arena.intern_entries(&expected);
1436
1437        assert_ne!(interned, collision);
1438        assert_eq!(arena.entry(interned, 0), Some(expected[0]));
1439        assert_eq!(arena.entry(interned, 1), Some(expected[1]));
1440    }
1441
1442    #[test]
1443    fn merge_with_empty_preserves_full_context_empty_path() {
1444        let mut arena = ContextArena::new();
1445        let mut workspace = PredictionWorkspace::default();
1446        let singleton = arena.singleton(EMPTY_CONTEXT, 42);
1447
1448        let merged = arena.merge(singleton, EMPTY_CONTEXT, false, &mut workspace);
1449
1450        assert_eq!(arena.len(merged), 2);
1451        assert_eq!(arena.return_state(merged, 0), Some(42));
1452        assert_eq!(arena.parent(merged, 0), Some(EMPTY_CONTEXT));
1453        assert_eq!(arena.return_state(merged, 1), Some(EMPTY_RETURN_STATE));
1454        assert!(arena.has_empty_path(merged));
1455    }
1456
1457    #[test]
1458    fn wildcard_merge_collapses_to_empty() {
1459        let mut arena = ContextArena::new();
1460        let mut workspace = PredictionWorkspace::default();
1461        let singleton = arena.singleton(EMPTY_CONTEXT, 42);
1462
1463        assert_eq!(
1464            arena.merge(singleton, EMPTY_CONTEXT, true, &mut workspace),
1465            EMPTY_CONTEXT
1466        );
1467    }
1468
1469    #[test]
1470    fn merge_is_order_independent() {
1471        let mut arena = ContextArena::new();
1472        let mut workspace = PredictionWorkspace::default();
1473        let left_parent = arena.singleton(EMPTY_CONTEXT, 100);
1474        let right_parent = arena.singleton(EMPTY_CONTEXT, 200);
1475        let left = arena.singleton(left_parent, 7);
1476        let right = arena.singleton(right_parent, 7);
1477
1478        let left_right = arena.merge(left, right, false, &mut workspace);
1479        workspace.reset();
1480        let right_left = arena.merge(right, left, false, &mut workspace);
1481
1482        assert_eq!(left_right, right_left);
1483        assert_eq!(arena.len(left_right), 1);
1484        let merged_parent = arena.parent(left_right, 0).expect("merged parent");
1485        assert_eq!(arena.len(merged_parent), 2);
1486        assert_eq!(arena.return_state(merged_parent, 0), Some(100));
1487        assert_eq!(arena.return_state(merged_parent, 1), Some(200));
1488    }
1489
1490    #[test]
1491    fn import_remaps_contexts_into_destination_arena() {
1492        let mut source = ContextArena::new();
1493        let parent = source.singleton(EMPTY_CONTEXT, 3);
1494        let child = source.singleton(parent, 9);
1495        let mut destination = ContextArena::new();
1496        let mut workspace = PredictionWorkspace::default();
1497
1498        let remap = destination.import_all(&source, &mut workspace);
1499        let imported = remap[usize::try_from(child.0).expect("context ID fits usize")];
1500
1501        assert_eq!(destination.return_state(imported, 0), Some(9));
1502        let imported_parent = destination.parent(imported, 0).expect("parent");
1503        assert_eq!(destination.return_state(imported_parent, 0), Some(3));
1504    }
1505
1506    #[test]
1507    fn config_set_merges_context_ids() {
1508        let mut arena = ContextArena::new();
1509        let mut workspace = PredictionWorkspace::default();
1510        let left = arena.singleton(EMPTY_CONTEXT, 1);
1511        let right = arena.singleton(EMPTY_CONTEXT, 2);
1512        let mut set = AtnConfigSet::new_full_context(true);
1513
1514        assert!(set.add(
1515            AtnConfig::new(1, 1, left, &arena),
1516            &mut arena,
1517            &mut workspace
1518        ));
1519        assert!(!set.add(
1520            AtnConfig::new(1, 1, right, &arena),
1521            &mut arena,
1522            &mut workspace
1523        ));
1524        assert_eq!(set.len(), 1);
1525        assert_eq!(arena.len(set.configs()[0].context), 2);
1526    }
1527
1528    #[test]
1529    fn predicate_provenance_is_idempotent_per_rule_path() {
1530        let arena = ContextArena::new();
1531        let mut provenance = PredictionSemanticProvenanceArena::default();
1532        let mut config = AtnConfig::new(1, 1, EMPTY_CONTEXT, &arena);
1533        config.enter_prediction_rule(&mut provenance, 4, 2);
1534        config.record_prediction_predicate(&mut provenance, 2, 3);
1535        let after_first = provenance
1536            .predicate_calls(config.semantic_provenance_id())
1537            .to_vec();
1538
1539        config.record_prediction_predicate(&mut provenance, 2, 3);
1540
1541        assert_eq!(
1542            provenance.predicate_calls(config.semantic_provenance_id()),
1543            after_first,
1544            "revisiting one predicate on the same rule path must not grow closure keys"
1545        );
1546    }
1547
1548    #[test]
1549    fn provenance_arena_stores_records_once_and_verifies_hash_collisions() {
1550        let mut arena = PredictionSemanticProvenanceArena::default();
1551        let first = PredictionSemanticProvenance {
1552            active_rule_calls: vec![PredictionRuleCall {
1553                source_state: 4,
1554                rule_index: 2,
1555            }],
1556            predicate_calls: Vec::new(),
1557        };
1558        let first_id = arena.intern(first.clone());
1559
1560        assert_eq!(arena.intern(first), first_id);
1561        assert_eq!(arena.records.len(), 1);
1562        assert_eq!(arena.interner_next.len(), 1);
1563
1564        let second = PredictionSemanticProvenance {
1565            active_rule_calls: vec![PredictionRuleCall {
1566                source_state: 5,
1567                rule_index: 3,
1568            }],
1569            predicate_calls: Vec::new(),
1570        };
1571        let mut hasher = PredictionFxHasher::default();
1572        second.hash(&mut hasher);
1573        arena.interner_heads.insert(hasher.finish(), first_id);
1574
1575        let second_id = arena.intern(second.clone());
1576        assert_ne!(second_id, first_id);
1577        assert_eq!(arena.intern(second), second_id);
1578        assert_eq!(arena.records.len(), 2);
1579        assert_eq!(arena.interner_next.len(), 2);
1580    }
1581
1582    #[test]
1583    fn config_set_keeps_distinct_prediction_provenance() {
1584        let mut arena = ContextArena::new();
1585        let mut provenance = PredictionSemanticProvenanceArena::default();
1586        let mut workspace = PredictionWorkspace::default();
1587        let mut first = AtnConfig::new(1, 1, EMPTY_CONTEXT, &arena);
1588        first.enter_prediction_rule(&mut provenance, 4, 2);
1589        let mut second = AtnConfig::new(1, 1, EMPTY_CONTEXT, &arena);
1590        second.enter_prediction_rule(&mut provenance, 5, 3);
1591        let mut set = AtnConfigSet::new();
1592
1593        assert!(set.add(first.clone(), &mut arena, &mut workspace));
1594        assert!(set.add(second, &mut arena, &mut workspace));
1595        assert!(!set.add(first, &mut arena, &mut workspace));
1596        assert_eq!(set.len(), 2);
1597        assert_eq!(set.config_index.len(), set.len());
1598
1599        set.remap_contexts(&[EMPTY_CONTEXT], &arena);
1600        assert_eq!(set.config_index.len(), set.len());
1601    }
1602
1603    #[cfg(target_pointer_width = "64")]
1604    #[test]
1605    fn parser_config_hot_path_layout_stays_compact() {
1606        let debug_generation = if cfg!(debug_assertions) {
1607            size_of::<u64>()
1608        } else {
1609            0
1610        };
1611        assert!(size_of::<AtnConfig>() <= 64 + debug_generation);
1612        assert!(size_of::<AtnConfigKey>() <= 56);
1613    }
1614
1615    #[test]
1616    fn workspace_drops_pathological_capacity() {
1617        let mut workspace = PredictionWorkspace::default();
1618        workspace
1619            .merge_cache
1620            .reserve(MAX_RETAINED_MERGE_CACHE_ENTRIES.saturating_mul(2));
1621        workspace
1622            .entries
1623            .reserve(MAX_RETAINED_CONTEXT_ENTRIES.saturating_mul(2));
1624        workspace.reset();
1625
1626        assert!(workspace.merge_cache.capacity() <= MAX_RETAINED_MERGE_CACHE_ENTRIES);
1627        assert!(workspace.entries.capacity() <= MAX_RETAINED_CONTEXT_ENTRIES);
1628    }
1629
1630    mod upstream_graph_nodes {
1631        use super::*;
1632        use std::collections::{BTreeSet, HashMap, VecDeque};
1633        use std::fmt::Write;
1634
1635        const EMPTY_WILDCARD_DOT: &str = concat!(
1636            "digraph G {\n",
1637            "rankdir=LR;\n",
1638            "  s0[label=\"*\"];\n",
1639            "}\n",
1640        );
1641        const EMPTY_FULL_CONTEXT_DOT: &str = concat!(
1642            "digraph G {\n",
1643            "rankdir=LR;\n",
1644            "  s0[label=\"$\"];\n",
1645            "}\n",
1646        );
1647        const X_EMPTY_FULL_CONTEXT_DOT: &str = concat!(
1648            "digraph G {\n",
1649            "rankdir=LR;\n",
1650            "  s0[shape=record, label=\"<p0>|<p1>$\"];\n",
1651            "  s1[label=\"$\"];\n",
1652            "  s0:p0->s1[label=\"9\"];\n",
1653            "}\n",
1654        );
1655        const A_DOT: &str = concat!(
1656            "digraph G {\n",
1657            "rankdir=LR;\n",
1658            "  s0[label=\"0\"];\n",
1659            "  s1[label=\"*\"];\n",
1660            "  s0->s1[label=\"1\"];\n",
1661            "}\n",
1662        );
1663        const A_EMPTY_AX_FULL_CONTEXT_DOT: &str = concat!(
1664            "digraph G {\n",
1665            "rankdir=LR;\n",
1666            "  s0[label=\"0\"];\n",
1667            "  s1[shape=record, label=\"<p0>|<p1>$\"];\n",
1668            "  s2[label=\"$\"];\n",
1669            "  s0->s1[label=\"1\"];\n",
1670            "  s1:p0->s2[label=\"9\"];\n",
1671            "}\n",
1672        );
1673        const NESTED_FULL_CONTEXT_DOT: &str = concat!(
1674            "digraph G {\n",
1675            "rankdir=LR;\n",
1676            "  s0[shape=record, label=\"<p0>|<p1>$\"];\n",
1677            "  s1[shape=record, label=\"<p0>|<p1>$\"];\n",
1678            "  s2[label=\"$\"];\n",
1679            "  s0:p0->s1[label=\"8\"];\n",
1680            "  s1:p0->s2[label=\"8\"];\n",
1681            "}\n",
1682        );
1683        const A_B_DOT: &str = concat!(
1684            "digraph G {\n",
1685            "rankdir=LR;\n",
1686            "  s0[shape=record, label=\"<p0>|<p1>\"];\n",
1687            "  s1[label=\"*\"];\n",
1688            "  s0:p0->s1[label=\"1\"];\n",
1689            "  s0:p1->s1[label=\"2\"];\n",
1690            "}\n",
1691        );
1692        const AX_AX_DOT: &str = concat!(
1693            "digraph G {\n",
1694            "rankdir=LR;\n",
1695            "  s0[label=\"0\"];\n",
1696            "  s1[label=\"1\"];\n",
1697            "  s2[label=\"*\"];\n",
1698            "  s0->s1[label=\"1\"];\n",
1699            "  s1->s2[label=\"9\"];\n",
1700            "}\n",
1701        );
1702        const ABX_ABX_DOT: &str = concat!(
1703            "digraph G {\n",
1704            "rankdir=LR;\n",
1705            "  s0[label=\"0\"];\n",
1706            "  s1[label=\"1\"];\n",
1707            "  s2[label=\"2\"];\n",
1708            "  s3[label=\"*\"];\n",
1709            "  s0->s1[label=\"1\"];\n",
1710            "  s1->s2[label=\"2\"];\n",
1711            "  s2->s3[label=\"9\"];\n",
1712            "}\n",
1713        );
1714        const ABX_ACX_DOT: &str = concat!(
1715            "digraph G {\n",
1716            "rankdir=LR;\n",
1717            "  s0[label=\"0\"];\n",
1718            "  s1[shape=record, label=\"<p0>|<p1>\"];\n",
1719            "  s2[label=\"2\"];\n",
1720            "  s3[label=\"*\"];\n",
1721            "  s0->s1[label=\"1\"];\n",
1722            "  s1:p0->s2[label=\"2\"];\n",
1723            "  s1:p1->s2[label=\"3\"];\n",
1724            "  s2->s3[label=\"9\"];\n",
1725            "}\n",
1726        );
1727        const AX_BX_DOT: &str = concat!(
1728            "digraph G {\n",
1729            "rankdir=LR;\n",
1730            "  s0[shape=record, label=\"<p0>|<p1>\"];\n",
1731            "  s1[label=\"1\"];\n",
1732            "  s2[label=\"*\"];\n",
1733            "  s0:p0->s1[label=\"1\"];\n",
1734            "  s0:p1->s1[label=\"2\"];\n",
1735            "  s1->s2[label=\"9\"];\n",
1736            "}\n",
1737        );
1738        const AX_BY_DOT: &str = concat!(
1739            "digraph G {\n",
1740            "rankdir=LR;\n",
1741            "  s0[shape=record, label=\"<p0>|<p1>\"];\n",
1742            "  s2[label=\"2\"];\n",
1743            "  s3[label=\"*\"];\n",
1744            "  s1[label=\"1\"];\n",
1745            "  s0:p0->s1[label=\"1\"];\n",
1746            "  s0:p1->s2[label=\"2\"];\n",
1747            "  s2->s3[label=\"10\"];\n",
1748            "  s1->s3[label=\"9\"];\n",
1749            "}\n",
1750        );
1751        const A_EMPTY_BX_DOT: &str = concat!(
1752            "digraph G {\n",
1753            "rankdir=LR;\n",
1754            "  s0[shape=record, label=\"<p0>|<p1>\"];\n",
1755            "  s2[label=\"2\"];\n",
1756            "  s1[label=\"*\"];\n",
1757            "  s0:p0->s1[label=\"1\"];\n",
1758            "  s0:p1->s2[label=\"2\"];\n",
1759            "  s2->s1[label=\"9\"];\n",
1760            "}\n",
1761        );
1762        const A_EMPTY_BX_FULL_CONTEXT_DOT: &str = concat!(
1763            "digraph G {\n",
1764            "rankdir=LR;\n",
1765            "  s0[shape=record, label=\"<p0>|<p1>\"];\n",
1766            "  s2[label=\"2\"];\n",
1767            "  s1[label=\"$\"];\n",
1768            "  s0:p0->s1[label=\"1\"];\n",
1769            "  s0:p1->s2[label=\"2\"];\n",
1770            "  s2->s1[label=\"9\"];\n",
1771            "}\n",
1772        );
1773        const AEX_BFX_DOT: &str = concat!(
1774            "digraph G {\n",
1775            "rankdir=LR;\n",
1776            "  s0[shape=record, label=\"<p0>|<p1>\"];\n",
1777            "  s2[label=\"2\"];\n",
1778            "  s3[label=\"3\"];\n",
1779            "  s4[label=\"*\"];\n",
1780            "  s1[label=\"1\"];\n",
1781            "  s0:p0->s1[label=\"1\"];\n",
1782            "  s0:p1->s2[label=\"2\"];\n",
1783            "  s2->s3[label=\"6\"];\n",
1784            "  s3->s4[label=\"9\"];\n",
1785            "  s1->s3[label=\"5\"];\n",
1786            "}\n",
1787        );
1788        const A_B_C_DOT: &str = concat!(
1789            "digraph G {\n",
1790            "rankdir=LR;\n",
1791            "  s0[shape=record, label=\"<p0>|<p1>|<p2>\"];\n",
1792            "  s1[label=\"*\"];\n",
1793            "  s0:p0->s1[label=\"1\"];\n",
1794            "  s0:p1->s1[label=\"2\"];\n",
1795            "  s0:p2->s1[label=\"3\"];\n",
1796            "}\n",
1797        );
1798        const AAX_AAY_DOT: &str = concat!(
1799            "digraph G {\n",
1800            "rankdir=LR;\n",
1801            "  s0[label=\"0\"];\n",
1802            "  s1[shape=record, label=\"<p0>|<p1>\"];\n",
1803            "  s2[label=\"*\"];\n",
1804            "  s0->s1[label=\"1\"];\n",
1805            "  s1:p0->s2[label=\"9\"];\n",
1806            "  s1:p1->s2[label=\"10\"];\n",
1807            "}\n",
1808        );
1809        const AAXC_AAYD_DOT: &str = concat!(
1810            "digraph G {\n",
1811            "rankdir=LR;\n",
1812            "  s0[shape=record, label=\"<p0>|<p1>|<p2>\"];\n",
1813            "  s2[label=\"*\"];\n",
1814            "  s1[shape=record, label=\"<p0>|<p1>\"];\n",
1815            "  s0:p0->s1[label=\"1\"];\n",
1816            "  s0:p1->s2[label=\"3\"];\n",
1817            "  s0:p2->s2[label=\"4\"];\n",
1818            "  s1:p0->s2[label=\"9\"];\n",
1819            "  s1:p1->s2[label=\"10\"];\n",
1820            "}\n",
1821        );
1822        const AAUBV_ACWDX_DOT: &str = concat!(
1823            "digraph G {\n",
1824            "rankdir=LR;\n",
1825            "  s0[shape=record, label=\"<p0>|<p1>|<p2>|<p3>\"];\n",
1826            "  s4[label=\"4\"];\n",
1827            "  s5[label=\"*\"];\n",
1828            "  s3[label=\"3\"];\n",
1829            "  s2[label=\"2\"];\n",
1830            "  s1[label=\"1\"];\n",
1831            "  s0:p0->s1[label=\"1\"];\n",
1832            "  s0:p1->s2[label=\"2\"];\n",
1833            "  s0:p2->s3[label=\"3\"];\n",
1834            "  s0:p3->s4[label=\"4\"];\n",
1835            "  s4->s5[label=\"9\"];\n",
1836            "  s3->s5[label=\"8\"];\n",
1837            "  s2->s5[label=\"7\"];\n",
1838            "  s1->s5[label=\"6\"];\n",
1839            "}\n",
1840        );
1841        const AAUBV_ABVDX_DOT: &str = concat!(
1842            "digraph G {\n",
1843            "rankdir=LR;\n",
1844            "  s0[shape=record, label=\"<p0>|<p1>|<p2>\"];\n",
1845            "  s3[label=\"3\"];\n",
1846            "  s4[label=\"*\"];\n",
1847            "  s2[label=\"2\"];\n",
1848            "  s1[label=\"1\"];\n",
1849            "  s0:p0->s1[label=\"1\"];\n",
1850            "  s0:p1->s2[label=\"2\"];\n",
1851            "  s0:p2->s3[label=\"4\"];\n",
1852            "  s3->s4[label=\"9\"];\n",
1853            "  s2->s4[label=\"7\"];\n",
1854            "  s1->s4[label=\"6\"];\n",
1855            "}\n",
1856        );
1857        const AAUBV_ABWDX_DOT: &str = concat!(
1858            "digraph G {\n",
1859            "rankdir=LR;\n",
1860            "  s0[shape=record, label=\"<p0>|<p1>|<p2>\"];\n",
1861            "  s3[label=\"3\"];\n",
1862            "  s4[label=\"*\"];\n",
1863            "  s2[shape=record, label=\"<p0>|<p1>\"];\n",
1864            "  s1[label=\"1\"];\n",
1865            "  s0:p0->s1[label=\"1\"];\n",
1866            "  s0:p1->s2[label=\"2\"];\n",
1867            "  s0:p2->s3[label=\"4\"];\n",
1868            "  s3->s4[label=\"9\"];\n",
1869            "  s2:p0->s4[label=\"7\"];\n",
1870            "  s2:p1->s4[label=\"8\"];\n",
1871            "  s1->s4[label=\"6\"];\n",
1872            "}\n",
1873        );
1874        const AAUBV_ABVDU_DOT: &str = concat!(
1875            "digraph G {\n",
1876            "rankdir=LR;\n",
1877            "  s0[shape=record, label=\"<p0>|<p1>|<p2>\"];\n",
1878            "  s2[label=\"2\"];\n",
1879            "  s3[label=\"*\"];\n",
1880            "  s1[label=\"1\"];\n",
1881            "  s0:p0->s1[label=\"1\"];\n",
1882            "  s0:p1->s2[label=\"2\"];\n",
1883            "  s0:p2->s1[label=\"4\"];\n",
1884            "  s2->s3[label=\"7\"];\n",
1885            "  s1->s3[label=\"6\"];\n",
1886            "}\n",
1887        );
1888        const AAUBU_ACUDU_DOT: &str = concat!(
1889            "digraph G {\n",
1890            "rankdir=LR;\n",
1891            "  s0[shape=record, label=\"<p0>|<p1>|<p2>|<p3>\"];\n",
1892            "  s1[label=\"1\"];\n",
1893            "  s2[label=\"*\"];\n",
1894            "  s0:p0->s1[label=\"1\"];\n",
1895            "  s0:p1->s1[label=\"2\"];\n",
1896            "  s0:p2->s1[label=\"3\"];\n",
1897            "  s0:p3->s1[label=\"4\"];\n",
1898            "  s1->s2[label=\"6\"];\n",
1899            "}\n",
1900        );
1901
1902        #[derive(Clone, Copy)]
1903        enum ContextSpec {
1904            Empty,
1905            Chain(&'static [usize]),
1906            Array(&'static [&'static [usize]]),
1907        }
1908
1909        #[derive(Clone, Copy)]
1910        enum Scenario {
1911            Merge {
1912                left: ContextSpec,
1913                right: ContextSpec,
1914            },
1915            NestedFullContext,
1916        }
1917
1918        struct GraphCase {
1919            source_test: &'static str,
1920            logical_id: &'static str,
1921            scenario: Scenario,
1922            root_is_wildcard: bool,
1923            expected: &'static str,
1924        }
1925
1926        impl GraphCase {
1927            const fn merge(
1928                source_test: &'static str,
1929                logical_id: &'static str,
1930                left: ContextSpec,
1931                right: ContextSpec,
1932                root_is_wildcard: bool,
1933                expected: &'static str,
1934            ) -> Self {
1935                Self {
1936                    source_test,
1937                    logical_id,
1938                    scenario: Scenario::Merge { left, right },
1939                    root_is_wildcard,
1940                    expected,
1941                }
1942            }
1943
1944            const fn nested_full_context(
1945                source_test: &'static str,
1946                logical_id: &'static str,
1947                expected: &'static str,
1948            ) -> Self {
1949                Self {
1950                    source_test,
1951                    logical_id,
1952                    scenario: Scenario::NestedFullContext,
1953                    root_is_wildcard: false,
1954                    expected,
1955                }
1956            }
1957        }
1958
1959        const CASES: &[GraphCase] = &[
1960            GraphCase::merge(
1961                "test_$_$",
1962                "testgraphnodes-test-9ea85e6b69",
1963                ContextSpec::Empty,
1964                ContextSpec::Empty,
1965                true,
1966                EMPTY_WILDCARD_DOT,
1967            ),
1968            GraphCase::merge(
1969                "test_$_$_fullctx",
1970                "testgraphnodes-test-fullctx-3a6b2d8201",
1971                ContextSpec::Empty,
1972                ContextSpec::Empty,
1973                false,
1974                EMPTY_FULL_CONTEXT_DOT,
1975            ),
1976            GraphCase::merge(
1977                "test_x_$",
1978                "testgraphnodes-test-x-546922b23c",
1979                ContextSpec::Chain(&[9]),
1980                ContextSpec::Empty,
1981                true,
1982                EMPTY_WILDCARD_DOT,
1983            ),
1984            GraphCase::merge(
1985                "test_x_$_fullctx",
1986                "testgraphnodes-test-x-fullctx-7fdaaf473e",
1987                ContextSpec::Chain(&[9]),
1988                ContextSpec::Empty,
1989                false,
1990                X_EMPTY_FULL_CONTEXT_DOT,
1991            ),
1992            GraphCase::merge(
1993                "test_$_x",
1994                "testgraphnodes-test-x-546922b23c",
1995                ContextSpec::Empty,
1996                ContextSpec::Chain(&[9]),
1997                true,
1998                EMPTY_WILDCARD_DOT,
1999            ),
2000            GraphCase::merge(
2001                "test_$_x_fullctx",
2002                "testgraphnodes-test-x-fullctx-7fdaaf473e",
2003                ContextSpec::Empty,
2004                ContextSpec::Chain(&[9]),
2005                false,
2006                X_EMPTY_FULL_CONTEXT_DOT,
2007            ),
2008            GraphCase::merge(
2009                "test_a_a",
2010                "testgraphnodes-test-a-a-429589e373",
2011                ContextSpec::Chain(&[1]),
2012                ContextSpec::Chain(&[1]),
2013                true,
2014                A_DOT,
2015            ),
2016            GraphCase::merge(
2017                "test_a$_ax",
2018                "testgraphnodes-test-a-ax-fd976a340d",
2019                ContextSpec::Chain(&[1]),
2020                ContextSpec::Chain(&[9, 1]),
2021                true,
2022                A_DOT,
2023            ),
2024            GraphCase::merge(
2025                "test_a$_ax_fullctx",
2026                "testgraphnodes-test-a-ax-fullctx-502155fcf9",
2027                ContextSpec::Chain(&[1]),
2028                ContextSpec::Chain(&[9, 1]),
2029                false,
2030                A_EMPTY_AX_FULL_CONTEXT_DOT,
2031            ),
2032            GraphCase::merge(
2033                "test_ax$_a$",
2034                "testgraphnodes-test-ax-a-62a48f251b",
2035                ContextSpec::Chain(&[9, 1]),
2036                ContextSpec::Chain(&[1]),
2037                true,
2038                A_DOT,
2039            ),
2040            GraphCase::nested_full_context(
2041                "test_aa$_a$_$_fullCtx",
2042                "testgraphnodes-test-aa-a-fullctx-8e728ea773",
2043                NESTED_FULL_CONTEXT_DOT,
2044            ),
2045            GraphCase::merge(
2046                "test_ax$_a$_fullctx",
2047                "testgraphnodes-test-ax-a-fullctx-7ef9c1d6b2",
2048                ContextSpec::Chain(&[9, 1]),
2049                ContextSpec::Chain(&[1]),
2050                false,
2051                A_EMPTY_AX_FULL_CONTEXT_DOT,
2052            ),
2053            GraphCase::merge(
2054                "test_a_b",
2055                "testgraphnodes-test-a-b-080058428f",
2056                ContextSpec::Chain(&[1]),
2057                ContextSpec::Chain(&[2]),
2058                true,
2059                A_B_DOT,
2060            ),
2061            GraphCase::merge(
2062                "test_ax_ax_same",
2063                "testgraphnodes-test-ax-ax-same-1504dc3dd3",
2064                ContextSpec::Chain(&[9, 1]),
2065                ContextSpec::Chain(&[9, 1]),
2066                true,
2067                AX_AX_DOT,
2068            ),
2069            GraphCase::merge(
2070                "test_ax_ax",
2071                "testgraphnodes-test-ax-ax-48f57578fa",
2072                ContextSpec::Chain(&[9, 1]),
2073                ContextSpec::Chain(&[9, 1]),
2074                true,
2075                AX_AX_DOT,
2076            ),
2077            GraphCase::merge(
2078                "test_abx_abx",
2079                "testgraphnodes-test-abx-abx-77366e32e9",
2080                ContextSpec::Chain(&[9, 2, 1]),
2081                ContextSpec::Chain(&[9, 2, 1]),
2082                true,
2083                ABX_ABX_DOT,
2084            ),
2085            GraphCase::merge(
2086                "test_abx_acx",
2087                "testgraphnodes-test-abx-acx-a3af7f90fa",
2088                ContextSpec::Chain(&[9, 2, 1]),
2089                ContextSpec::Chain(&[9, 3, 1]),
2090                true,
2091                ABX_ACX_DOT,
2092            ),
2093            GraphCase::merge(
2094                "test_ax_bx_same",
2095                "testgraphnodes-test-ax-bx-same-d0506bf7a9",
2096                ContextSpec::Chain(&[9, 1]),
2097                ContextSpec::Chain(&[9, 2]),
2098                true,
2099                AX_BX_DOT,
2100            ),
2101            GraphCase::merge(
2102                "test_ax_bx",
2103                "testgraphnodes-test-ax-bx-1ea2df9a04",
2104                ContextSpec::Chain(&[9, 1]),
2105                ContextSpec::Chain(&[9, 2]),
2106                true,
2107                AX_BX_DOT,
2108            ),
2109            GraphCase::merge(
2110                "test_ax_by",
2111                "testgraphnodes-test-ax-by-47815d59d2",
2112                ContextSpec::Chain(&[9, 1]),
2113                ContextSpec::Chain(&[10, 2]),
2114                true,
2115                AX_BY_DOT,
2116            ),
2117            GraphCase::merge(
2118                "test_a$_bx",
2119                "testgraphnodes-test-a-bx-b15f7b876f",
2120                ContextSpec::Chain(&[1]),
2121                ContextSpec::Chain(&[9, 2]),
2122                true,
2123                A_EMPTY_BX_DOT,
2124            ),
2125            GraphCase::merge(
2126                "test_a$_bx_fullctx",
2127                "testgraphnodes-test-a-bx-fullctx-a35242b6cf",
2128                ContextSpec::Chain(&[1]),
2129                ContextSpec::Chain(&[9, 2]),
2130                false,
2131                A_EMPTY_BX_FULL_CONTEXT_DOT,
2132            ),
2133            GraphCase::merge(
2134                "test_aex_bfx",
2135                "testgraphnodes-test-aex-bfx-07ad9de126",
2136                ContextSpec::Chain(&[9, 5, 1]),
2137                ContextSpec::Chain(&[9, 6, 2]),
2138                true,
2139                AEX_BFX_DOT,
2140            ),
2141            GraphCase::merge(
2142                "test_A$_A$_fullctx",
2143                "testgraphnodes-test-a-a-fullctx-b023f64b6c",
2144                ContextSpec::Array(&[&[]]),
2145                ContextSpec::Array(&[&[]]),
2146                false,
2147                EMPTY_FULL_CONTEXT_DOT,
2148            ),
2149            GraphCase::merge(
2150                "test_Aab_Ac",
2151                "testgraphnodes-test-aab-ac-139c5b709d",
2152                ContextSpec::Array(&[&[1], &[2]]),
2153                ContextSpec::Array(&[&[3]]),
2154                true,
2155                A_B_C_DOT,
2156            ),
2157            GraphCase::merge(
2158                "test_Aa_Aa",
2159                "testgraphnodes-test-aa-aa-0a175c83db",
2160                ContextSpec::Array(&[&[1]]),
2161                ContextSpec::Array(&[&[1]]),
2162                true,
2163                A_DOT,
2164            ),
2165            GraphCase::merge(
2166                "test_Aa_Abc",
2167                "testgraphnodes-test-aa-abc-db12d99894",
2168                ContextSpec::Array(&[&[1]]),
2169                ContextSpec::Array(&[&[2], &[3]]),
2170                true,
2171                A_B_C_DOT,
2172            ),
2173            GraphCase::merge(
2174                "test_Aac_Ab",
2175                "testgraphnodes-test-aac-ab-ef785e17e7",
2176                ContextSpec::Array(&[&[1], &[3]]),
2177                ContextSpec::Array(&[&[2]]),
2178                true,
2179                A_B_C_DOT,
2180            ),
2181            GraphCase::merge(
2182                "test_Aab_Aa",
2183                "testgraphnodes-test-aab-aa-d90d8d54f0",
2184                ContextSpec::Array(&[&[1], &[2]]),
2185                ContextSpec::Array(&[&[1]]),
2186                true,
2187                A_B_DOT,
2188            ),
2189            GraphCase::merge(
2190                "test_Aab_Ab",
2191                "testgraphnodes-test-aab-ab-e2d46352b4",
2192                ContextSpec::Array(&[&[1], &[2]]),
2193                ContextSpec::Array(&[&[2]]),
2194                true,
2195                A_B_DOT,
2196            ),
2197            GraphCase::merge(
2198                "test_Aax_Aby",
2199                "testgraphnodes-test-aax-aby-cccf935759",
2200                ContextSpec::Array(&[&[9, 1]]),
2201                ContextSpec::Array(&[&[10, 2]]),
2202                true,
2203                AX_BY_DOT,
2204            ),
2205            GraphCase::merge(
2206                "test_Aax_Aay",
2207                "testgraphnodes-test-aax-aay-c0f9b80842",
2208                ContextSpec::Array(&[&[9, 1]]),
2209                ContextSpec::Array(&[&[10, 1]]),
2210                true,
2211                AAX_AAY_DOT,
2212            ),
2213            GraphCase::merge(
2214                "test_Aaxc_Aayd",
2215                "testgraphnodes-test-aaxc-aayd-a73533f64d",
2216                ContextSpec::Array(&[&[9, 1], &[3]]),
2217                ContextSpec::Array(&[&[10, 1], &[4]]),
2218                true,
2219                AAXC_AAYD_DOT,
2220            ),
2221            GraphCase::merge(
2222                "test_Aaubv_Acwdx",
2223                "testgraphnodes-test-aaubv-acwdx-f479c849df",
2224                ContextSpec::Array(&[&[6, 1], &[7, 2]]),
2225                ContextSpec::Array(&[&[8, 3], &[9, 4]]),
2226                true,
2227                AAUBV_ACWDX_DOT,
2228            ),
2229            GraphCase::merge(
2230                "test_Aaubv_Abvdx",
2231                "testgraphnodes-test-aaubv-abvdx-01eb5714fe",
2232                ContextSpec::Array(&[&[6, 1], &[7, 2]]),
2233                ContextSpec::Array(&[&[7, 2], &[9, 4]]),
2234                true,
2235                AAUBV_ABVDX_DOT,
2236            ),
2237            GraphCase::merge(
2238                "test_Aaubv_Abwdx",
2239                "testgraphnodes-test-aaubv-abwdx-7953c9b489",
2240                ContextSpec::Array(&[&[6, 1], &[7, 2]]),
2241                ContextSpec::Array(&[&[8, 2], &[9, 4]]),
2242                true,
2243                AAUBV_ABWDX_DOT,
2244            ),
2245            GraphCase::merge(
2246                "test_Aaubv_Abvdu",
2247                "testgraphnodes-test-aaubv-abvdu-ecc8850384",
2248                ContextSpec::Array(&[&[6, 1], &[7, 2]]),
2249                ContextSpec::Array(&[&[7, 2], &[6, 4]]),
2250                true,
2251                AAUBV_ABVDU_DOT,
2252            ),
2253            GraphCase::merge(
2254                "test_Aaubu_Acudu",
2255                "testgraphnodes-test-aaubu-acudu-7cb798b616",
2256                ContextSpec::Array(&[&[6, 1], &[6, 2]]),
2257                ContextSpec::Array(&[&[6, 3], &[6, 4]]),
2258                true,
2259                AAUBU_ACUDU_DOT,
2260            ),
2261        ];
2262
2263        fn build_chain(arena: &mut ContextArena, return_states: &[usize]) -> ContextId {
2264            let mut context = EMPTY_CONTEXT;
2265            for &return_state in return_states {
2266                context = arena.singleton(context, return_state);
2267            }
2268            context
2269        }
2270
2271        fn build_context(arena: &mut ContextArena, spec: ContextSpec) -> ContextId {
2272            match spec {
2273                ContextSpec::Empty => EMPTY_CONTEXT,
2274                ContextSpec::Chain(return_states) => build_chain(arena, return_states),
2275                ContextSpec::Array(chains) => {
2276                    let mut entries = Vec::with_capacity(chains.len());
2277                    for return_states in chains {
2278                        let context = build_chain(arena, return_states);
2279                        entries.push(arena.first_entry(context));
2280                    }
2281                    arena.intern_entries(&entries)
2282                }
2283            }
2284        }
2285
2286        fn run_case(case: &GraphCase) -> String {
2287            let mut arena = ContextArena::new();
2288            let mut workspace = PredictionWorkspace::default();
2289            let merged = match case.scenario {
2290                Scenario::Merge { left, right } => {
2291                    let left = build_context(&mut arena, left);
2292                    let right = build_context(&mut arena, right);
2293                    arena.merge(left, right, case.root_is_wildcard, &mut workspace)
2294                }
2295                Scenario::NestedFullContext => {
2296                    let child = arena.singleton(EMPTY_CONTEXT, 8);
2297                    let right = arena.merge(EMPTY_CONTEXT, child, false, &mut workspace);
2298                    let left = arena.singleton(right, 8);
2299                    arena.merge(left, right, false, &mut workspace)
2300                }
2301            };
2302            render_dot(&arena, merged, case.root_is_wildcard)
2303        }
2304
2305        fn render_dot(arena: &ContextArena, context: ContextId, root_is_wildcard: bool) -> String {
2306            let mut nodes = String::new();
2307            let mut edges = String::new();
2308            let mut context_ids = HashMap::new();
2309            let mut work_list = VecDeque::new();
2310            context_ids.insert(context, 0);
2311            work_list.push_back(context);
2312
2313            while let Some(current) = work_list.pop_front() {
2314                let current_id = context_ids[&current];
2315                let len = arena.len(current);
2316                write!(&mut nodes, "  s{current_id}[").expect("write to string");
2317                if len > 1 {
2318                    nodes.push_str("shape=record, ");
2319                }
2320                nodes.push_str("label=\"");
2321                if arena.is_empty(current) {
2322                    nodes.push(if root_is_wildcard { '*' } else { '$' });
2323                } else if len > 1 {
2324                    for index in 0..len {
2325                        if index > 0 {
2326                            nodes.push('|');
2327                        }
2328                        write!(&mut nodes, "<p{index}>").expect("write to string");
2329                        if arena.return_state(current, index) == Some(EMPTY_RETURN_STATE) {
2330                            nodes.push(if root_is_wildcard { '*' } else { '$' });
2331                        }
2332                    }
2333                } else {
2334                    write!(&mut nodes, "{current_id}").expect("write to string");
2335                }
2336                nodes.push_str("\"];\n");
2337
2338                if arena.is_empty(current) {
2339                    continue;
2340                }
2341                for index in 0..len {
2342                    let return_state = arena
2343                        .return_state(current, index)
2344                        .expect("context entry in range");
2345                    if return_state == EMPTY_RETURN_STATE {
2346                        continue;
2347                    }
2348                    let parent = arena.parent(current, index).expect("non-empty parent");
2349                    let parent_id = if let Some(&parent_id) = context_ids.get(&parent) {
2350                        parent_id
2351                    } else {
2352                        let parent_id = context_ids.len();
2353                        context_ids.insert(parent, parent_id);
2354                        work_list.push_front(parent);
2355                        parent_id
2356                    };
2357
2358                    write!(&mut edges, "  s{current_id}").expect("write to string");
2359                    if len > 1 {
2360                        write!(&mut edges, ":p{index}").expect("write to string");
2361                    }
2362                    writeln!(&mut edges, "->s{parent_id}[label=\"{return_state}\"];")
2363                        .expect("write to string");
2364                }
2365            }
2366
2367            let mut dot = String::from("digraph G {\nrankdir=LR;\n");
2368            dot.push_str(&nodes);
2369            dot.push_str(&edges);
2370            dot.push_str("}\n");
2371            dot
2372        }
2373
2374        #[test]
2375        fn pinned_upstream_test_graph_nodes_matches_dot() {
2376            assert_eq!(CASES.len(), 38, "pinned Java source case inventory drifted");
2377            let source_tests = CASES
2378                .iter()
2379                .map(|case| case.source_test)
2380                .collect::<BTreeSet<_>>();
2381            assert_eq!(
2382                source_tests.len(),
2383                38,
2384                "pinned Java source test names must be unique"
2385            );
2386            let logical_ids = CASES
2387                .iter()
2388                .map(|case| case.logical_id)
2389                .collect::<BTreeSet<_>>();
2390            assert_eq!(
2391                logical_ids.len(),
2392                36,
2393                "pinned upstream logical row inventory drifted"
2394            );
2395
2396            let selector = std::env::var("ANTLR_GRAPH_NODE_CASE").ok();
2397            let selected = CASES
2398                .iter()
2399                .filter(|case| {
2400                    selector
2401                        .as_deref()
2402                        .is_none_or(|logical_id| case.logical_id == logical_id)
2403                })
2404                .collect::<Vec<_>>();
2405            assert!(
2406                !selected.is_empty(),
2407                "ANTLR_GRAPH_NODE_CASE={:?} matched no logical row",
2408                selector.as_deref().unwrap_or_default()
2409            );
2410
2411            let mut mismatches = Vec::new();
2412            for case in &selected {
2413                let actual = run_case(case);
2414                if actual != case.expected {
2415                    mismatches.push(format!(
2416                        "logical_id={}\nsource_test={}\n--- expected\n{}--- actual\n{}",
2417                        case.logical_id, case.source_test, case.expected, actual
2418                    ));
2419                }
2420            }
2421
2422            assert!(
2423                mismatches.is_empty(),
2424                "TestGraphNodes DOT mismatches ({}/{} source cases):\n\n{}",
2425                mismatches.len(),
2426                selected.len(),
2427                mismatches.join("\n")
2428            );
2429        }
2430    }
2431}