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