1use 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#[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#[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#[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 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#[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(¤t.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#[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 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[¤t];
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}