Skip to main content

antlr4_runtime/
prediction.rs

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