Skip to main content

infotheory/backends/
ctw.rs

1//! Context Tree Weighting (CTW) and Factorized Action-Conditional CTW (FAC-CTW).
2//!
3//! This implementation stores maximal non-root unary runs as chains while keeping
4//! the root and true branching points explicit. Updates and reverts rebuild only
5//! the touched context path from exact per-node state, which preserves predictive
6//! semantics while avoiding the `O(depth)` explicit-node blow-up for singleton
7//! paths.
8
9use std::cell::{Cell, RefCell};
10use std::f64;
11use std::mem::size_of;
12use std::sync::OnceLock;
13
14type Symbol = bool;
15const HISTORY_WORD_BITS: usize = u64::BITS as usize;
16
17#[inline(always)]
18fn history_word_len(bits: usize) -> usize {
19    bits.div_ceil(HISTORY_WORD_BITS)
20}
21
22trait HistoryAccess {
23    fn len(&self) -> usize;
24    fn bit(&self, index: usize) -> Symbol;
25
26    #[inline(always)]
27    fn recent_bit(&self, _depth: usize) -> Option<Symbol> {
28        None
29    }
30
31    #[inline(always)]
32    fn recent_path_bits(&self, _depth: usize, _len: usize) -> Option<u64> {
33        None
34    }
35}
36
37impl HistoryAccess for [Symbol] {
38    #[inline(always)]
39    fn len(&self) -> usize {
40        <[Symbol]>::len(self)
41    }
42
43    #[inline(always)]
44    fn bit(&self, index: usize) -> Symbol {
45        debug_assert!(index < self.len());
46        unsafe { *self.get_unchecked(index) }
47    }
48}
49
50impl HistoryAccess for Vec<Symbol> {
51    #[inline(always)]
52    fn len(&self) -> usize {
53        self.as_slice().len()
54    }
55
56    #[inline(always)]
57    fn bit(&self, index: usize) -> Symbol {
58        self.as_slice().bit(index)
59    }
60}
61
62#[derive(Clone, Debug, Default, PartialEq, Eq)]
63struct BitHistory {
64    words: Vec<u64>,
65    len: usize,
66    recent: u64,
67}
68
69impl BitHistory {
70    #[inline(always)]
71    fn len(&self) -> usize {
72        self.len
73    }
74
75    #[inline(always)]
76    fn is_empty(&self) -> bool {
77        self.len == 0
78    }
79
80    #[inline(always)]
81    fn memory_usage(&self) -> usize {
82        self.words.capacity() * size_of::<u64>()
83    }
84
85    #[inline]
86    fn rebuild_recent(&mut self) {
87        self.recent = 0;
88        let tail_len = self.len.min(HISTORY_WORD_BITS);
89        for depth in 0..tail_len {
90            let idx = self.len - depth - 1;
91            if self.bit(idx) {
92                self.recent |= 1u64 << depth;
93            }
94        }
95    }
96
97    #[inline]
98    fn reserve_exact(&mut self, additional_bits: usize) {
99        let required_bits = self.len.saturating_add(additional_bits);
100        let required_words = history_word_len(required_bits);
101        if required_words > self.words.capacity() {
102            self.words
103                .reserve_exact(required_words.saturating_sub(self.words.len()));
104        }
105    }
106
107    #[inline]
108    fn push(&mut self, bit: Symbol) {
109        let word_idx = self.len / HISTORY_WORD_BITS;
110        let bit_idx = self.len % HISTORY_WORD_BITS;
111        if word_idx == self.words.len() {
112            self.words.push(0);
113        }
114        let mask = 1u64 << bit_idx;
115        if bit {
116            self.words[word_idx] |= mask;
117        } else {
118            self.words[word_idx] &= !mask;
119        }
120        self.recent = (self.recent << 1) | (bit as u64);
121        self.len += 1;
122    }
123
124    #[inline]
125    fn pop(&mut self) -> Option<Symbol> {
126        if self.len == 0 {
127            return None;
128        }
129        let next_len = self.len - 1;
130        let word_idx = next_len / HISTORY_WORD_BITS;
131        let bit_idx = next_len % HISTORY_WORD_BITS;
132        let mask = 1u64 << bit_idx;
133        let bit = (self.words[word_idx] & mask) != 0;
134        if bit {
135            self.words[word_idx] &= !mask;
136        }
137        self.len = next_len;
138        if bit_idx == 0 {
139            self.words.truncate(word_idx);
140        }
141        self.recent >>= 1;
142        if self.len >= HISTORY_WORD_BITS {
143            let exposed_idx = self.len - HISTORY_WORD_BITS;
144            if self.bit(exposed_idx) {
145                self.recent |= 1u64 << (HISTORY_WORD_BITS - 1);
146            }
147        }
148        Some(bit)
149    }
150
151    #[inline]
152    fn extend_from_slice(&mut self, symbols: &[Symbol]) {
153        self.reserve_exact(symbols.len());
154        for &symbol in symbols {
155            self.push(symbol);
156        }
157    }
158
159    #[inline]
160    fn truncate(&mut self, new_len: usize) {
161        if new_len >= self.len {
162            return;
163        }
164        self.len = new_len;
165        self.words.truncate(history_word_len(new_len));
166        let rem = new_len % HISTORY_WORD_BITS;
167        if rem != 0 {
168            let mask = (1u64 << rem) - 1;
169            if let Some(last) = self.words.last_mut() {
170                *last &= mask;
171            }
172        }
173        self.rebuild_recent();
174    }
175
176    #[inline]
177    fn clear(&mut self) {
178        self.words.clear();
179        self.len = 0;
180        self.recent = 0;
181    }
182
183    #[cfg(test)]
184    fn to_vec(&self) -> Vec<Symbol> {
185        (0..self.len).map(|idx| self.bit(idx)).collect()
186    }
187}
188
189impl HistoryAccess for BitHistory {
190    #[inline(always)]
191    fn len(&self) -> usize {
192        self.len
193    }
194
195    #[inline(always)]
196    fn bit(&self, index: usize) -> Symbol {
197        debug_assert!(index < self.len);
198        let word = unsafe { *self.words.get_unchecked(index / HISTORY_WORD_BITS) };
199        ((word >> (index % HISTORY_WORD_BITS)) & 1) != 0
200    }
201
202    #[inline(always)]
203    fn recent_bit(&self, depth: usize) -> Option<Symbol> {
204        if depth < self.len.min(HISTORY_WORD_BITS) {
205            Some(((self.recent >> depth) & 1) != 0)
206        } else {
207            None
208        }
209    }
210
211    #[inline(always)]
212    fn recent_path_bits(&self, depth: usize, len: usize) -> Option<u64> {
213        let available = self.len.saturating_sub(depth).min(len);
214        if available == 0 {
215            return Some(0);
216        }
217        if depth + available <= self.len.min(HISTORY_WORD_BITS) {
218            let mask = if available >= HISTORY_WORD_BITS {
219                u64::MAX
220            } else {
221                (1u64 << available) - 1
222            };
223            Some((self.recent >> depth) & mask)
224        } else {
225            None
226        }
227    }
228}
229
230const CTW_LOG_CACHE_LIMIT: usize = 1 << 24;
231const CTW_HOT_PREFIX_DEPTH_DEFAULT: usize = 12;
232const CTW_LOG_OVERFLOW_CACHE_SLOTS: usize = 1 << 14;
233
234#[cfg(not(test))]
235#[inline(always)]
236fn ctw_log_cache_limit() -> usize {
237    CTW_LOG_CACHE_LIMIT
238}
239
240#[cfg(test)]
241#[inline(always)]
242fn ctw_log_cache_limit() -> usize {
243    CTW_TEST_LOG_CACHE_LIMIT.with(|limit| limit.borrow().unwrap_or(CTW_LOG_CACHE_LIMIT))
244}
245
246fn ctw_hot_prefix_depth_limit() -> usize {
247    static HOT_PREFIX_DEPTH: OnceLock<usize> = OnceLock::new();
248    *HOT_PREFIX_DEPTH.get_or_init(|| {
249        std::env::var("INFOTHEORY_CTW_HOT_PREFIX_DEPTH")
250            .ok()
251            .and_then(|raw| raw.parse::<usize>().ok())
252            .unwrap_or(CTW_HOT_PREFIX_DEPTH_DEFAULT)
253    })
254}
255
256#[cfg(not(test))]
257#[inline(always)]
258fn ctw_log_overflow_cache_slots() -> usize {
259    CTW_LOG_OVERFLOW_CACHE_SLOTS
260}
261
262#[cfg(test)]
263#[inline(always)]
264fn ctw_log_overflow_cache_slots() -> usize {
265    CTW_TEST_LOG_OVERFLOW_CACHE_SLOTS
266        .with(|slots| slots.borrow().unwrap_or(CTW_LOG_OVERFLOW_CACHE_SLOTS))
267}
268
269#[cfg(test)]
270#[inline(always)]
271fn ensure_log_caches(log_int: &mut Vec<f64>, log_half: &mut Vec<f64>, upto: usize) {
272    if upto < log_int.len() {
273        return;
274    }
275    let required_len = upto + 1;
276    log_int.reserve(required_len - log_int.len());
277    log_half.reserve(required_len - log_half.len());
278    push_log_cache_entries(log_int, log_half, upto);
279}
280
281#[inline(always)]
282fn reserve_bounded_log_cache(cache: &mut Vec<f64>, required_len: usize, max_len: usize) {
283    if cache.capacity() >= required_len {
284        return;
285    }
286    let target_capacity = cache
287        .capacity()
288        .saturating_mul(2)
289        .max(required_len)
290        .min(max_len);
291    cache.reserve_exact(target_capacity - cache.len());
292}
293
294#[inline(always)]
295fn ensure_bounded_log_caches(
296    log_int: &mut Vec<f64>,
297    log_half: &mut Vec<f64>,
298    upto: usize,
299    limit: usize,
300) {
301    let target = upto.min(limit);
302    if target < log_int.len() {
303        return;
304    }
305    let required_len = target + 1;
306    let max_len = limit + 1;
307    reserve_bounded_log_cache(log_int, required_len, max_len);
308    reserve_bounded_log_cache(log_half, required_len, max_len);
309    push_log_cache_entries(log_int, log_half, target);
310}
311
312#[inline(always)]
313fn push_log_cache_entries(log_int: &mut Vec<f64>, log_half: &mut Vec<f64>, upto: usize) {
314    let start = log_int.len();
315    for n in start..=upto {
316        if n == 0 {
317            log_int.push(f64::NEG_INFINITY);
318        } else {
319            log_int.push((n as f64).ln());
320        }
321        log_half.push((n as f64 + 0.5).ln());
322    }
323}
324
325struct LogCacheSlot {
326    key: Cell<usize>,
327    value: Cell<f64>,
328}
329
330impl LogCacheSlot {
331    #[inline(always)]
332    fn empty() -> Self {
333        Self {
334            key: Cell::new(usize::MAX),
335            value: Cell::new(0.0),
336        }
337    }
338}
339
340#[inline(always)]
341fn ensure_overflow_log_cache(cache: &mut Vec<LogCacheSlot>, slots: usize) {
342    if slots == 0 {
343        cache.clear();
344        cache.shrink_to(0);
345        return;
346    }
347    if cache.len() == slots {
348        return;
349    }
350    cache.clear();
351    if cache.capacity() < slots {
352        cache.reserve_exact(slots);
353    }
354    cache.resize_with(slots, LogCacheSlot::empty);
355    if cache.capacity() > slots {
356        cache.shrink_to(slots);
357    }
358}
359
360trait CtLogAccess: Copy {
361    fn log_int(self, n: usize) -> f64;
362    fn log_half(self, n: usize) -> f64;
363}
364
365#[derive(Clone, Copy)]
366struct CachedLogs<'a> {
367    log_int: &'a [f64],
368    log_half: &'a [f64],
369}
370
371impl<'a> CachedLogs<'a> {
372    #[cfg(test)]
373    #[inline(always)]
374    fn new(log_int: &'a [f64], log_half: &'a [f64]) -> Self {
375        Self { log_int, log_half }
376    }
377}
378
379impl CtLogAccess for CachedLogs<'_> {
380    #[inline(always)]
381    fn log_int(self, n: usize) -> f64 {
382        debug_assert!(n < self.log_int.len());
383        // SAFETY: `CachedLogs` is used only when the shared prefix cache already
384        // contains every value through the current visit bound.
385        unsafe { *self.log_int.get_unchecked(n) }
386    }
387
388    #[inline(always)]
389    fn log_half(self, n: usize) -> f64 {
390        debug_assert!(n < self.log_half.len());
391        // SAFETY: `CachedLogs` is used only when the shared prefix cache already
392        // contains every value through the current visit bound.
393        unsafe { *self.log_half.get_unchecked(n) }
394    }
395}
396
397#[derive(Clone, Copy)]
398struct BoundedLogs<'a> {
399    log_int: &'a [f64],
400    log_half: &'a [f64],
401    overflow_log_int: &'a [LogCacheSlot],
402    overflow_log_half: &'a [LogCacheSlot],
403}
404
405impl<'a> BoundedLogs<'a> {
406    #[cfg(test)]
407    #[inline(always)]
408    fn new(log_int: &'a [f64], log_half: &'a [f64]) -> Self {
409        Self {
410            log_int,
411            log_half,
412            overflow_log_int: &[],
413            overflow_log_half: &[],
414        }
415    }
416
417    #[inline(always)]
418    fn with_overflow(
419        log_int: &'a [f64],
420        log_half: &'a [f64],
421        overflow_log_int: &'a [LogCacheSlot],
422        overflow_log_half: &'a [LogCacheSlot],
423    ) -> Self {
424        Self {
425            log_int,
426            log_half,
427            overflow_log_int,
428            overflow_log_half,
429        }
430    }
431
432    #[inline(always)]
433    fn lookup_overflow(slots: &'a [LogCacheSlot], n: usize, compute: impl FnOnce() -> f64) -> f64 {
434        if slots.is_empty() {
435            return compute();
436        }
437        let len = slots.len();
438        let slot_idx = if len.is_power_of_two() {
439            n & (len - 1)
440        } else {
441            n % len
442        };
443        let slot = &slots[slot_idx];
444        if slot.key.get() == n {
445            slot.value.get()
446        } else {
447            let value = compute();
448            slot.key.set(n);
449            slot.value.set(value);
450            value
451        }
452    }
453}
454
455impl CtLogAccess for BoundedLogs<'_> {
456    #[inline(always)]
457    fn log_int(self, n: usize) -> f64 {
458        if n < self.log_int.len() {
459            self.log_int[n]
460        } else if n == 0 {
461            f64::NEG_INFINITY
462        } else {
463            Self::lookup_overflow(self.overflow_log_int, n, || (n as f64).ln())
464        }
465    }
466
467    #[inline(always)]
468    fn log_half(self, n: usize) -> f64 {
469        if n < self.log_half.len() {
470            self.log_half[n]
471        } else {
472            Self::lookup_overflow(self.overflow_log_half, n, || (n as f64 + 0.5).ln())
473        }
474    }
475}
476
477#[derive(Default)]
478struct SharedLogCache {
479    log_int: Vec<f64>,
480    log_half: Vec<f64>,
481    overflow_log_int: Vec<LogCacheSlot>,
482    overflow_log_half: Vec<LogCacheSlot>,
483}
484
485impl SharedLogCache {
486    fn new() -> Self {
487        Self {
488            log_int: vec![f64::NEG_INFINITY],
489            log_half: vec![(0.5f64).ln()],
490            overflow_log_int: Vec::new(),
491            overflow_log_half: Vec::new(),
492        }
493    }
494
495    #[inline(always)]
496    fn ensure(&mut self, upto: usize) {
497        let limit = ctw_log_cache_limit();
498        ensure_bounded_log_caches(&mut self.log_int, &mut self.log_half, upto, limit);
499        if upto > limit {
500            let slots = ctw_log_overflow_cache_slots();
501            ensure_overflow_log_cache(&mut self.overflow_log_int, slots);
502            ensure_overflow_log_cache(&mut self.overflow_log_half, slots);
503        }
504    }
505
506    #[inline(always)]
507    fn memory_usage(&self) -> usize {
508        self.log_int.capacity() * size_of::<f64>()
509            + self.log_half.capacity() * size_of::<f64>()
510            + self.overflow_log_int.capacity() * size_of::<LogCacheSlot>()
511            + self.overflow_log_half.capacity() * size_of::<LogCacheSlot>()
512    }
513}
514
515thread_local! {
516    static CTW_SHARED_LOG_CACHE: RefCell<SharedLogCache> =
517        RefCell::new(SharedLogCache::new());
518}
519
520#[cfg(test)]
521thread_local! {
522    static CTW_TEST_LOG_CACHE_LIMIT: RefCell<Option<usize>> = const { RefCell::new(None) };
523}
524
525#[cfg(test)]
526thread_local! {
527    static CTW_TEST_LOG_OVERFLOW_CACHE_SLOTS: RefCell<Option<usize>> = const { RefCell::new(None) };
528}
529
530#[inline]
531fn with_shared_cached_logs<R>(upto: usize, f: impl FnOnce(CachedLogs<'_>) -> R) -> R {
532    debug_assert!(upto <= ctw_log_cache_limit());
533    CTW_SHARED_LOG_CACHE.with(|cache_cell| {
534        let mut cache = cache_cell.borrow_mut();
535        cache.ensure(upto);
536        f(CachedLogs {
537            log_int: &cache.log_int,
538            log_half: &cache.log_half,
539        })
540    })
541}
542
543#[inline]
544fn with_shared_bounded_logs<R>(upto: usize, f: impl FnOnce(BoundedLogs<'_>) -> R) -> R {
545    CTW_SHARED_LOG_CACHE.with(|cache_cell| {
546        let mut cache = cache_cell.borrow_mut();
547        cache.ensure(upto);
548        f(BoundedLogs::with_overflow(
549            &cache.log_int,
550            &cache.log_half,
551            &cache.overflow_log_int,
552            &cache.overflow_log_half,
553        ))
554    })
555}
556
557#[inline]
558fn shared_log_cache_memory_usage() -> usize {
559    CTW_SHARED_LOG_CACHE.with(|cache_cell| cache_cell.borrow().memory_usage())
560}
561
562#[cfg(test)]
563#[inline]
564fn shared_log_cache_lens() -> (usize, usize) {
565    CTW_SHARED_LOG_CACHE.with(|cache_cell| {
566        let cache = cache_cell.borrow();
567        (cache.log_int.len(), cache.log_half.len())
568    })
569}
570
571#[cfg(test)]
572#[inline]
573fn shared_log_overflow_cache_lens() -> (usize, usize) {
574    CTW_SHARED_LOG_CACHE.with(|cache_cell| {
575        let cache = cache_cell.borrow();
576        (cache.overflow_log_int.len(), cache.overflow_log_half.len())
577    })
578}
579
580#[cfg(test)]
581pub(crate) fn reset_shared_log_cache_for_test() {
582    CTW_SHARED_LOG_CACHE.with(|cache_cell| {
583        *cache_cell.borrow_mut() = SharedLogCache::new();
584    });
585}
586
587#[inline(always)]
588fn history_symbol<H: HistoryAccess + ?Sized>(history: &H, depth: usize) -> Symbol {
589    if let Some(bit) = history.recent_bit(depth) {
590        return bit;
591    }
592    let idx = history.len().wrapping_sub(depth + 1);
593    if depth < history.len() {
594        history.bit(idx)
595    } else {
596        false
597    }
598}
599
600#[inline(always)]
601fn history_at_or_zero<H: HistoryAccess + ?Sized>(
602    history: &H,
603    history_len: isize,
604    idx: isize,
605) -> Symbol {
606    if idx >= 0 && idx < history_len {
607        history.bit(idx as usize)
608    } else {
609        false
610    }
611}
612
613const INDEX_BITS: u32 = 31;
614const INDEX_LIMIT: usize = 1usize << INDEX_BITS;
615const CHILD_SEGMENT_TAG: u32 = 1u32 << INDEX_BITS;
616const CHILD_INDEX_MASK: u32 = CHILD_SEGMENT_TAG - 1;
617const SEG_META_MODE_SHIFT: u32 = 30;
618const SEG_META_MODE_MASK: u32 = 0b11 << SEG_META_MODE_SHIFT;
619const SEG_LEN_MASK: u32 = !SEG_META_MODE_MASK;
620const SEG_MODE_EXACT: u32 = 0 << SEG_META_MODE_SHIFT;
621const SEG_MODE_HISTORY: u32 = 1 << SEG_META_MODE_SHIFT;
622const SEG_MODE_HISTORY_INVERT: u32 = 2 << SEG_META_MODE_SHIFT;
623const SEG_MODE_CONST: u32 = 3 << SEG_META_MODE_SHIFT;
624const SEG_EXACT_MAX_LEN: u32 = 64;
625
626/// Index into the explicit-node arena.
627#[derive(Clone, Copy, Debug, PartialEq, Eq)]
628pub struct NodeIndex(u32);
629
630impl NodeIndex {
631    #[cold]
632    #[inline(never)]
633    fn overflow() -> ! {
634        panic!("ctw node index overflow");
635    }
636
637    #[inline(always)]
638    fn from_usize(idx: usize) -> Self {
639        if idx >= INDEX_LIMIT {
640            Self::overflow();
641        }
642        Self(idx as u32)
643    }
644
645    #[inline(always)]
646    fn get(self) -> usize {
647        self.0 as usize
648    }
649}
650
651/// Index into the unary-segment arena.
652#[derive(Clone, Copy, Debug, PartialEq, Eq)]
653struct SegmentIndex(u32);
654
655impl SegmentIndex {
656    #[cold]
657    #[inline(never)]
658    fn overflow() -> ! {
659        panic!("ctw segment index overflow");
660    }
661
662    #[inline(always)]
663    fn from_usize(idx: usize) -> Self {
664        if idx >= INDEX_LIMIT {
665            Self::overflow();
666        }
667        Self(idx as u32)
668    }
669
670    #[inline(always)]
671    fn get(self) -> usize {
672        self.0 as usize
673    }
674}
675
676/// Tagged child reference: `None`, explicit node, or unary segment.
677#[derive(Clone, Copy, Debug, PartialEq, Eq)]
678struct ChildRef(u32);
679
680impl ChildRef {
681    const NONE: ChildRef = ChildRef(u32::MAX);
682
683    #[inline(always)]
684    fn from_node(idx: NodeIndex) -> Self {
685        debug_assert!(idx.0 < CHILD_SEGMENT_TAG);
686        Self(idx.0)
687    }
688
689    #[inline(always)]
690    fn from_segment(idx: SegmentIndex) -> Self {
691        debug_assert!(idx.0 < CHILD_SEGMENT_TAG);
692        Self(CHILD_SEGMENT_TAG | idx.0)
693    }
694
695    #[inline(always)]
696    fn is_none(self) -> bool {
697        self.0 == u32::MAX
698    }
699
700    #[inline(always)]
701    fn is_some(self) -> bool {
702        self.0 != u32::MAX
703    }
704
705    #[inline(always)]
706    fn as_node(self) -> Option<NodeIndex> {
707        if self.is_none() || (self.0 & CHILD_SEGMENT_TAG) != 0 {
708            None
709        } else {
710            Some(NodeIndex(self.0))
711        }
712    }
713
714    #[inline(always)]
715    fn as_segment(self) -> Option<SegmentIndex> {
716        if self.is_none() || (self.0 & CHILD_SEGMENT_TAG) == 0 {
717            None
718        } else {
719            Some(SegmentIndex(self.0 & CHILD_INDEX_MASK))
720        }
721    }
722}
723
724impl Default for ChildRef {
725    fn default() -> Self {
726        Self::NONE
727    }
728}
729
730#[derive(Clone, Copy, Debug, Default)]
731struct SegmentPayload {
732    repr_lo: u32,
733    repr_hi: u32,
734    meta: u32,
735}
736
737impl SegmentPayload {
738    #[inline(always)]
739    fn exact(bits: u64, len: u32) -> Self {
740        debug_assert!(len <= SEG_EXACT_MAX_LEN);
741        debug_assert!(len <= SEG_LEN_MASK);
742        Self {
743            repr_lo: bits as u32,
744            repr_hi: (bits >> 32) as u32,
745            meta: SEG_MODE_EXACT | len,
746        }
747    }
748
749    #[inline(always)]
750    fn history(anchor: u32, len: u32, invert: bool) -> Self {
751        debug_assert!(len <= SEG_LEN_MASK);
752        Self {
753            repr_lo: anchor,
754            repr_hi: 0,
755            meta: if invert {
756                SEG_MODE_HISTORY_INVERT | len
757            } else {
758                SEG_MODE_HISTORY | len
759            },
760        }
761    }
762
763    #[inline(always)]
764    fn constant(bit: bool, len: u32) -> Self {
765        debug_assert!(len <= SEG_LEN_MASK);
766        Self {
767            repr_lo: bit as u32,
768            repr_hi: 0,
769            meta: SEG_MODE_CONST | len,
770        }
771    }
772
773    #[inline(always)]
774    fn len(self) -> u32 {
775        self.meta & SEG_LEN_MASK
776    }
777
778    #[inline(always)]
779    fn set_len(&mut self, len: u32) {
780        debug_assert!(len <= SEG_LEN_MASK);
781        self.meta = (self.meta & SEG_META_MODE_MASK) | len;
782    }
783
784    #[inline(always)]
785    fn mode(self) -> u32 {
786        self.meta & SEG_META_MODE_MASK
787    }
788
789    #[inline(always)]
790    fn is_exact(self) -> bool {
791        self.mode() == SEG_MODE_EXACT
792    }
793
794    #[inline(always)]
795    fn exact_bits(self) -> u64 {
796        (self.repr_lo as u64) | ((self.repr_hi as u64) << 32)
797    }
798
799    #[inline(always)]
800    fn anchor_or_const(self) -> u32 {
801        self.repr_lo
802    }
803
804    #[inline(always)]
805    fn const_bit(self) -> bool {
806        (self.repr_lo & 1) != 0
807    }
808
809    #[inline(always)]
810    fn prepend_exact(self, edge: usize) -> Option<Self> {
811        if !self.is_exact() || self.len() >= SEG_EXACT_MAX_LEN {
812            return None;
813        }
814        let len = self.len() + 1;
815        let bits = ((edge as u64) & 1) | (self.exact_bits() << 1);
816        Some(Self::exact(bits, len))
817    }
818
819    #[inline(always)]
820    fn prefix(self, len: u32) -> Self {
821        debug_assert!(len <= self.len());
822        match self.mode() {
823            SEG_MODE_EXACT => Self::exact(self.exact_bits() & low_bits_mask_u64(len), len),
824            SEG_MODE_HISTORY | SEG_MODE_HISTORY_INVERT => Self {
825                meta: (self.meta & SEG_META_MODE_MASK) | len,
826                ..self
827            },
828            SEG_MODE_CONST => Self::constant(self.const_bit(), len),
829            _ => unreachable!("invalid ctw segment payload mode"),
830        }
831    }
832
833    #[inline(always)]
834    fn suffix_after(self, skip: u32) -> Self {
835        debug_assert!(skip <= self.len());
836        let new_len = self.len() - skip;
837        match self.mode() {
838            SEG_MODE_EXACT => Self::exact(self.exact_bits() >> skip, new_len),
839            SEG_MODE_HISTORY | SEG_MODE_HISTORY_INVERT => Self {
840                repr_lo: self
841                    .anchor_or_const()
842                    .checked_sub(skip)
843                    .expect("ctw history segment anchor underflow"),
844                meta: (self.meta & SEG_META_MODE_MASK) | new_len,
845                ..self
846            },
847            SEG_MODE_CONST => Self::constant(self.const_bit(), new_len),
848            _ => unreachable!("invalid ctw segment payload mode"),
849        }
850    }
851
852    #[inline(always)]
853    fn from_path<H: HistoryAccess + ?Sized>(history: &H, depth: usize, len: u32) -> Option<Self> {
854        if len > SEG_EXACT_MAX_LEN {
855            return None;
856        }
857        Some(Self::exact(
858            path_bits_from_history(history, depth, len as usize),
859            len,
860        ))
861    }
862}
863
864#[derive(Clone, Copy, Debug, Default)]
865struct LevelState {
866    symbol_count: [u32; 2],
867    log_prob_kt: f64,
868    sibling: ChildRef,
869}
870
871#[cfg(test)]
872#[allow(dead_code)]
873#[derive(Clone, Copy, Debug)]
874struct PredictEntry {
875    symbol_count: [u32; 2],
876    log_prob_kt: f64,
877    log_prob_weighted: f64,
878    sibling_weight: f64,
879    has_sibling: bool,
880}
881
882#[derive(Clone, Copy, Debug)]
883enum Detach {
884    NodeChild { node: NodeIndex, edge: usize },
885    SegmentNext { segment: SegmentIndex, new_len: u32 },
886}
887
888#[derive(Clone, Copy, Debug, PartialEq, Eq)]
889enum ExistingSource {
890    None,
891    Node(NodeIndex),
892    Segment(SegmentIndex, u32),
893}
894
895#[derive(Clone, Copy, Debug, PartialEq, Eq)]
896enum PreparedEnd {
897    MaxDepth,
898    MissingAtRoot,
899    MissingAfterCurrent,
900    MismatchAtCurrentSegment,
901}
902
903#[derive(Clone, Copy, Debug, PartialEq)]
904struct PreparedStep {
905    source: ExistingSource,
906    span: u32,
907    sibling_weight: f64,
908    has_sibling: u8,
909}
910
911#[derive(Clone, Copy, Debug)]
912/// Compact node payload used by [`CtArena`].
913pub struct CtNode {
914    children: [ChildRef; 2],
915    log_prob_kt: f64,
916    log_prob_weighted: f64,
917    symbol_count: [u32; 2],
918}
919
920#[derive(Clone, Copy, Debug)]
921struct CtSegment {
922    tail: ChildRef,
923    log_prob_kt: f64,
924    head_log_prob_weighted: f64,
925    symbol_count: [u32; 2],
926    payload: SegmentPayload,
927}
928
929impl Default for CtSegment {
930    fn default() -> Self {
931        Self {
932            tail: ChildRef::NONE,
933            log_prob_kt: 0.0,
934            head_log_prob_weighted: 0.0,
935            symbol_count: [0, 0],
936            payload: SegmentPayload::default(),
937        }
938    }
939}
940
941impl CtSegment {
942    #[inline(always)]
943    fn len(self) -> u32 {
944        self.payload.len()
945    }
946
947    #[inline(always)]
948    fn set_len(&mut self, len: u32) {
949        self.payload.set_len(len);
950    }
951}
952
953#[inline(always)]
954fn low_bits_mask_u64(len: u32) -> u64 {
955    if len >= 64 {
956        u64::MAX
957    } else {
958        (1u64 << len) - 1
959    }
960}
961
962#[inline(always)]
963fn path_bits_from_history<H: HistoryAccess + ?Sized>(history: &H, depth: usize, len: usize) -> u64 {
964    if let Some(bits) = history.recent_path_bits(depth, len) {
965        return bits;
966    }
967    let history_len = history.len();
968    let available = history_len.saturating_sub(depth).min(len);
969    if available == 0 {
970        return 0;
971    }
972
973    let mut bits = 0u64;
974    let mut hist_idx = history_len - depth - 1;
975    for offset in 0..available {
976        bits |= (history.bit(hist_idx) as u64) << offset;
977        if hist_idx == 0 {
978            break;
979        }
980        hist_idx -= 1;
981    }
982    bits
983}
984
985#[inline(always)]
986fn shift_path_bits(path_bits: u64, consumed: usize) -> u64 {
987    if consumed >= 64 {
988        0
989    } else {
990        path_bits >> consumed
991    }
992}
993
994#[inline(always)]
995fn first_exact_segment_mismatch(
996    exact_bits: u64,
997    path_bits: u64,
998    comparable_len: usize,
999) -> Option<(usize, bool, bool)> {
1000    if comparable_len == 0 {
1001        return None;
1002    }
1003
1004    let diff = (exact_bits ^ path_bits) & low_bits_mask_u64(comparable_len as u32);
1005    if diff == 0 {
1006        None
1007    } else {
1008        let offset = diff.trailing_zeros() as usize;
1009        Some((
1010            offset,
1011            ((path_bits >> offset) & 1) != 0,
1012            ((exact_bits >> offset) & 1) != 0,
1013        ))
1014    }
1015}
1016
1017#[inline(always)]
1018fn predict_ratio_kt(counts: [u32; 2], sym_idx: usize) -> f64 {
1019    let total = (counts[0] + counts[1]) as f64;
1020    let sym_count = counts[sym_idx] as f64;
1021    (sym_count + 0.5) / (total + 1.0)
1022}
1023
1024#[inline(always)]
1025fn predict_ratio_kt_one(counts: [u32; 2]) -> f64 {
1026    let total = (counts[0] + counts[1]) as f64;
1027    let sym_count = counts[1] as f64;
1028    (sym_count + 0.5) / (total + 1.0)
1029}
1030
1031#[inline(always)]
1032fn update_weighted_log_prob_non_leaf(kt_log_prob: f64, log_prob_w0: f64, log_prob_w1: f64) -> f64 {
1033    let child_log_prob = log_prob_w0 + log_prob_w1;
1034    if child_log_prob.to_bits() == kt_log_prob.to_bits() {
1035        return clamp_log_prob(kt_log_prob);
1036    }
1037    let delta = child_log_prob - kt_log_prob;
1038    let log_prob_weighted = if delta >= 0.0 {
1039        child_log_prob + (-delta).exp().ln_1p() - std::f64::consts::LN_2
1040    } else {
1041        kt_log_prob + delta.exp().ln_1p() - std::f64::consts::LN_2
1042    };
1043    clamp_log_prob(log_prob_weighted)
1044}
1045
1046#[inline(always)]
1047fn update_weighted_log_prob(
1048    kt_log_prob: f64,
1049    log_prob_w0: f64,
1050    log_prob_w1: f64,
1051    is_leaf: bool,
1052) -> f64 {
1053    if is_leaf {
1054        clamp_log_prob(kt_log_prob)
1055    } else {
1056        update_weighted_log_prob_non_leaf(kt_log_prob, log_prob_w0, log_prob_w1)
1057    }
1058}
1059
1060#[inline(always)]
1061fn clamp_log_prob(log_prob: f64) -> f64 {
1062    if log_prob > 1.0e-10 { 0.0 } else { log_prob }
1063}
1064
1065#[inline(always)]
1066fn logsumexp_pair(lhs: f64, rhs: f64) -> f64 {
1067    if lhs == f64::NEG_INFINITY {
1068        return rhs;
1069    }
1070    if rhs == f64::NEG_INFINITY {
1071        return lhs;
1072    }
1073    let pivot = lhs.max(rhs);
1074    pivot + ((lhs - pivot).exp() + (rhs - pivot).exp()).ln()
1075}
1076
1077#[inline(always)]
1078fn unary_chain_log_weight(kt_log_prob: f64, continuation_log_prob: f64, len: u32) -> f64 {
1079    debug_assert!(len > 0);
1080    if kt_log_prob.to_bits() == continuation_log_prob.to_bits() {
1081        return kt_log_prob;
1082    }
1083    let log_alpha = -(len as f64) * std::f64::consts::LN_2;
1084    let alpha = log_alpha.exp();
1085    let log_kt_mass = kt_log_prob + (-alpha).ln_1p();
1086    let log_cont_mass = continuation_log_prob + log_alpha;
1087    clamp_log_prob(logsumexp_pair(log_kt_mass, log_cont_mass))
1088}
1089
1090#[inline(always)]
1091fn combined_weight_ratio_internal(
1092    kt_log_prob: f64,
1093    counts: [u32; 2],
1094    path_child_log_prob: f64,
1095    sibling_log_prob: f64,
1096    child_ratio: f64,
1097    sym_idx: usize,
1098) -> (f64, f64) {
1099    let kt_ratio = predict_ratio_kt(counts, sym_idx);
1100    let child_log_prob = path_child_log_prob + sibling_log_prob;
1101    if child_log_prob.to_bits() == kt_log_prob.to_bits() {
1102        return (clamp_log_prob(kt_log_prob), 0.5 * (kt_ratio + child_ratio));
1103    }
1104    let delta = child_log_prob - kt_log_prob;
1105    if delta >= 0.0 {
1106        let x = (-delta).exp();
1107        (
1108            clamp_log_prob(child_log_prob + x.ln_1p() - std::f64::consts::LN_2),
1109            (kt_ratio * x + child_ratio) / (1.0 + x),
1110        )
1111    } else {
1112        let x = delta.exp();
1113        (
1114            clamp_log_prob(kt_log_prob + x.ln_1p() - std::f64::consts::LN_2),
1115            (kt_ratio + x * child_ratio) / (1.0 + x),
1116        )
1117    }
1118}
1119
1120#[inline(always)]
1121fn combined_weight_ratio_internal_one(
1122    kt_log_prob: f64,
1123    counts: [u32; 2],
1124    path_child_log_prob: f64,
1125    sibling_log_prob: f64,
1126    child_ratio: f64,
1127) -> (f64, f64) {
1128    let kt_ratio = predict_ratio_kt_one(counts);
1129    let child_log_prob = path_child_log_prob + sibling_log_prob;
1130    if child_log_prob.to_bits() == kt_log_prob.to_bits() {
1131        return (clamp_log_prob(kt_log_prob), 0.5 * (kt_ratio + child_ratio));
1132    }
1133    let delta = child_log_prob - kt_log_prob;
1134    if delta >= 0.0 {
1135        let x = (-delta).exp();
1136        (
1137            clamp_log_prob(child_log_prob + x.ln_1p() - std::f64::consts::LN_2),
1138            (kt_ratio * x + child_ratio) / (1.0 + x),
1139        )
1140    } else {
1141        let x = delta.exp();
1142        (
1143            clamp_log_prob(kt_log_prob + x.ln_1p() - std::f64::consts::LN_2),
1144            (kt_ratio + x * child_ratio) / (1.0 + x),
1145        )
1146    }
1147}
1148
1149#[inline(always)]
1150fn unary_chain_log_weight_precomputed(
1151    kt_log_prob: f64,
1152    continuation_log_prob: f64,
1153    alpha: f64,
1154    log_alpha: f64,
1155    log_one_minus_alpha: f64,
1156) -> f64 {
1157    if kt_log_prob.to_bits() == continuation_log_prob.to_bits() {
1158        return clamp_log_prob(kt_log_prob);
1159    }
1160
1161    let delta = continuation_log_prob - kt_log_prob;
1162    let log_prob_weighted = if delta >= 0.0 {
1163        let x = ((1.0 - alpha) / alpha) * (-delta).exp();
1164        continuation_log_prob + log_alpha + x.ln_1p()
1165    } else {
1166        let x = (alpha / (1.0 - alpha)) * delta.exp();
1167        kt_log_prob + log_one_minus_alpha + x.ln_1p()
1168    };
1169    clamp_log_prob(log_prob_weighted)
1170}
1171
1172#[inline(always)]
1173// The CTW ratio transform is a hot scalar kernel; grouping these independent
1174// numeric inputs into a temporary struct would add ceremony without clarifying
1175// ownership, invariants, or call-site meaning.
1176#[allow(clippy::too_many_arguments)]
1177fn unary_chain_ratio_transform_precomputed(
1178    kt_log_prob: f64,
1179    counts: [u32; 2],
1180    continuation_log_prob: f64,
1181    continuation_ratio: f64,
1182    alpha: f64,
1183    log_alpha: f64,
1184    log_one_minus_alpha: f64,
1185    sym_idx: usize,
1186) -> (f64, f64) {
1187    let kt_ratio = predict_ratio_kt(counts, sym_idx);
1188    if kt_log_prob.to_bits() == continuation_log_prob.to_bits()
1189        && kt_ratio.to_bits() == continuation_ratio.to_bits()
1190    {
1191        return (clamp_log_prob(kt_log_prob), kt_ratio);
1192    }
1193    if kt_log_prob.to_bits() == continuation_log_prob.to_bits() {
1194        return (
1195            clamp_log_prob(kt_log_prob),
1196            (1.0 - alpha) * kt_ratio + alpha * continuation_ratio,
1197        );
1198    }
1199
1200    let delta = continuation_log_prob - kt_log_prob;
1201    if delta >= 0.0 {
1202        let x = ((1.0 - alpha) / alpha) * (-delta).exp();
1203        (
1204            clamp_log_prob(continuation_log_prob + log_alpha + x.ln_1p()),
1205            (kt_ratio * x + continuation_ratio) / (1.0 + x),
1206        )
1207    } else {
1208        let x = (alpha / (1.0 - alpha)) * delta.exp();
1209        (
1210            clamp_log_prob(kt_log_prob + log_one_minus_alpha + x.ln_1p()),
1211            (kt_ratio + x * continuation_ratio) / (1.0 + x),
1212        )
1213    }
1214}
1215
1216#[inline(always)]
1217fn unary_chain_ratio_transform_precomputed_one(
1218    kt_log_prob: f64,
1219    counts: [u32; 2],
1220    continuation_log_prob: f64,
1221    continuation_ratio: f64,
1222    alpha: f64,
1223    log_alpha: f64,
1224    log_one_minus_alpha: f64,
1225) -> (f64, f64) {
1226    let kt_ratio = predict_ratio_kt_one(counts);
1227    if kt_log_prob.to_bits() == continuation_log_prob.to_bits()
1228        && kt_ratio.to_bits() == continuation_ratio.to_bits()
1229    {
1230        return (clamp_log_prob(kt_log_prob), kt_ratio);
1231    }
1232    if kt_log_prob.to_bits() == continuation_log_prob.to_bits() {
1233        return (
1234            clamp_log_prob(kt_log_prob),
1235            (1.0 - alpha) * kt_ratio + alpha * continuation_ratio,
1236        );
1237    }
1238
1239    let delta = continuation_log_prob - kt_log_prob;
1240    if delta >= 0.0 {
1241        let x = ((1.0 - alpha) / alpha) * (-delta).exp();
1242        (
1243            clamp_log_prob(continuation_log_prob + log_alpha + x.ln_1p()),
1244            (kt_ratio * x + continuation_ratio) / (1.0 + x),
1245        )
1246    } else {
1247        let x = (alpha / (1.0 - alpha)) * delta.exp();
1248        (
1249            clamp_log_prob(kt_log_prob + log_one_minus_alpha + x.ln_1p()),
1250            (kt_ratio + x * continuation_ratio) / (1.0 + x),
1251        )
1252    }
1253}
1254
1255#[inline(always)]
1256fn predict_ratio_internal(
1257    kt_log_prob: f64,
1258    counts: [u32; 2],
1259    path_child_log_prob: f64,
1260    sibling_log_prob: f64,
1261    child_ratio: f64,
1262    sym_idx: usize,
1263) -> f64 {
1264    let kt_ratio = predict_ratio_kt(counts, sym_idx);
1265    let child_log_prob = path_child_log_prob + sibling_log_prob;
1266    if child_log_prob.to_bits() == kt_log_prob.to_bits() {
1267        return 0.5 * (kt_ratio + child_ratio);
1268    }
1269    let delta = child_log_prob - kt_log_prob;
1270    if delta >= 0.0 {
1271        let inv_rho = (-delta).exp();
1272        (kt_ratio * inv_rho + child_ratio) / (1.0 + inv_rho)
1273    } else {
1274        let rho = delta.exp();
1275        (kt_ratio + rho * child_ratio) / (1.0 + rho)
1276    }
1277}
1278
1279#[inline(always)]
1280fn predict_ratio_internal_one(
1281    kt_log_prob: f64,
1282    counts: [u32; 2],
1283    path_child_log_prob: f64,
1284    sibling_log_prob: f64,
1285    child_ratio: f64,
1286) -> f64 {
1287    let kt_ratio = predict_ratio_kt_one(counts);
1288    let child_log_prob = path_child_log_prob + sibling_log_prob;
1289    if child_log_prob.to_bits() == kt_log_prob.to_bits() {
1290        return 0.5 * (kt_ratio + child_ratio);
1291    }
1292    let delta = child_log_prob - kt_log_prob;
1293    if delta >= 0.0 {
1294        let inv_rho = (-delta).exp();
1295        (kt_ratio * inv_rho + child_ratio) / (1.0 + inv_rho)
1296    } else {
1297        let rho = delta.exp();
1298        (kt_ratio + rho * child_ratio) / (1.0 + rho)
1299    }
1300}
1301
1302#[inline(always)]
1303fn path_edge_at_depth<H: HistoryAccess + ?Sized>(
1304    history: &H,
1305    history_len: usize,
1306    depth: usize,
1307) -> bool {
1308    if depth < history_len {
1309        history.bit(history_len - depth - 1)
1310    } else {
1311        false
1312    }
1313}
1314
1315#[inline(always)]
1316fn segment_edge_from_parts<H: HistoryAccess + ?Sized>(
1317    segment: CtSegment,
1318    offset: usize,
1319    history: &H,
1320    history_len: usize,
1321) -> bool {
1322    match segment.payload.mode() {
1323        SEG_MODE_EXACT => ((segment.payload.exact_bits() >> offset) & 1) != 0,
1324        SEG_MODE_HISTORY | SEG_MODE_HISTORY_INVERT => {
1325            if segment.payload.anchor_or_const() as usize >= offset {
1326                let hist_idx = segment.payload.anchor_or_const() as usize - offset;
1327                if hist_idx < history_len {
1328                    let raw = history.bit(hist_idx);
1329                    if segment.payload.mode() == SEG_MODE_HISTORY_INVERT {
1330                        !raw
1331                    } else {
1332                        raw
1333                    }
1334                } else {
1335                    segment.payload.mode() == SEG_MODE_HISTORY_INVERT
1336                }
1337            } else {
1338                segment.payload.mode() == SEG_MODE_HISTORY_INVERT
1339            }
1340        }
1341        SEG_MODE_CONST => segment.payload.const_bit(),
1342        _ => unreachable!("invalid ctw segment payload mode"),
1343    }
1344}
1345
1346#[inline(always)]
1347fn first_segment_mismatch<H: HistoryAccess + ?Sized>(
1348    segment: CtSegment,
1349    depth: usize,
1350    history: &H,
1351    comparable_len: usize,
1352) -> Option<(usize, bool, bool)> {
1353    if comparable_len == 0 {
1354        return None;
1355    }
1356
1357    match segment.payload.mode() {
1358        SEG_MODE_EXACT => first_exact_segment_mismatch(
1359            segment.payload.exact_bits(),
1360            path_bits_from_history(history, depth, comparable_len),
1361            comparable_len,
1362        ),
1363        SEG_MODE_HISTORY | SEG_MODE_HISTORY_INVERT => {
1364            let history_len = history.len() as isize;
1365            let mut path_hist_idx = history_len - depth as isize - 1;
1366            let mut seg_hist_idx = segment.payload.anchor_or_const() as isize;
1367            let invert = segment.payload.mode() == SEG_MODE_HISTORY_INVERT;
1368            for offset in 0..comparable_len {
1369                let path_edge = history_at_or_zero(history, history_len, path_hist_idx);
1370                let existing_raw = history_at_or_zero(history, history_len, seg_hist_idx);
1371                let existing_edge = if invert { !existing_raw } else { existing_raw };
1372                if existing_edge != path_edge {
1373                    return Some((offset, path_edge, existing_edge));
1374                }
1375                path_hist_idx -= 1;
1376                seg_hist_idx -= 1;
1377            }
1378            None
1379        }
1380        SEG_MODE_CONST => {
1381            let history_len = history.len() as isize;
1382            let mut path_hist_idx = history_len - depth as isize - 1;
1383            let existing_edge = segment.payload.const_bit();
1384            for offset in 0..comparable_len {
1385                let path_edge = history_at_or_zero(history, history_len, path_hist_idx);
1386                if existing_edge != path_edge {
1387                    return Some((offset, path_edge, existing_edge));
1388                }
1389                path_hist_idx -= 1;
1390            }
1391            None
1392        }
1393        _ => unreachable!("invalid ctw segment payload mode"),
1394    }
1395}
1396
1397#[inline]
1398fn apply_update_to_state_raw<L: CtLogAccess>(
1399    logs: L,
1400    symbol_count: &mut [u32; 2],
1401    log_prob_kt: &mut f64,
1402    sym_idx: usize,
1403) {
1404    let total_before = (symbol_count[0] + symbol_count[1]) as usize;
1405    let sym_before = symbol_count[sym_idx] as usize;
1406    debug_assert!(sym_before <= total_before);
1407    let log_half_before = logs.log_half(sym_before);
1408    let log_total_after = logs.log_int(total_before + 1);
1409    *log_prob_kt += log_half_before - log_total_after;
1410    if *log_prob_kt > 1.0e-10 {
1411        *log_prob_kt = 0.0;
1412    }
1413    symbol_count[sym_idx] = symbol_count[sym_idx]
1414        .checked_add(1)
1415        .expect("ctw symbol count overflow");
1416}
1417
1418#[inline]
1419fn apply_revert_to_state_raw<L: CtLogAccess>(
1420    logs: L,
1421    symbol_count: &mut [u32; 2],
1422    log_prob_kt: &mut f64,
1423    sym_idx: usize,
1424) {
1425    let total = (symbol_count[0] + symbol_count[1]) as usize;
1426    let sym_count = symbol_count[sym_idx] as usize;
1427    if sym_count > 0 && total > 0 {
1428        let log_half_before = logs.log_half(sym_count - 1);
1429        let log_total = logs.log_int(total);
1430        *log_prob_kt -= log_half_before - log_total;
1431        symbol_count[sym_idx] -= 1;
1432    }
1433    if *log_prob_kt > 1.0e-10 {
1434        *log_prob_kt = 0.0;
1435    }
1436}
1437
1438#[derive(Clone, Debug)]
1439/// Arena allocator and storage for CTW nodes/segments.
1440///
1441/// This structure owns all backing memory for compressed CTW trees and provides
1442/// index-based access used by the update/revert engine.
1443pub struct CtArena {
1444    nodes: Vec<CtNode>,
1445    segments: Vec<CtSegment>,
1446    free_nodes: Vec<NodeIndex>,
1447    free_segments: Vec<SegmentIndex>,
1448}
1449
1450impl CtArena {
1451    /// Create an empty arena with small default capacities.
1452    pub fn new() -> Self {
1453        Self {
1454            nodes: Vec::with_capacity(1024),
1455            segments: Vec::with_capacity(1024),
1456            free_nodes: Vec::new(),
1457            free_segments: Vec::new(),
1458        }
1459    }
1460
1461    /// Create an empty arena with explicit node-oriented capacity hint.
1462    pub fn with_capacity(cap: usize) -> Self {
1463        Self {
1464            nodes: Vec::with_capacity(cap),
1465            segments: Vec::with_capacity(cap / 4 + 1),
1466            free_nodes: Vec::new(),
1467            free_segments: Vec::new(),
1468        }
1469    }
1470
1471    #[inline]
1472    /// Reserve additional capacity for upcoming node/segment allocations.
1473    pub fn reserve_exact(&mut self, additional: usize) {
1474        self.nodes.reserve_exact(additional);
1475        self.segments.reserve_exact(additional / 4 + 1);
1476    }
1477
1478    #[inline(always)]
1479    fn reset_node_slot(&mut self, idx: NodeIndex) {
1480        self.nodes[idx.get()] = CtNode {
1481            children: [ChildRef::NONE, ChildRef::NONE],
1482            log_prob_kt: 0.0,
1483            log_prob_weighted: 0.0,
1484            symbol_count: [0, 0],
1485        };
1486    }
1487
1488    #[inline(always)]
1489    fn reset_segment_slot(&mut self, idx: SegmentIndex) {
1490        self.segments[idx.get()] = CtSegment::default();
1491    }
1492
1493    #[inline(always)]
1494    fn alloc_node(&mut self) -> NodeIndex {
1495        if let Some(idx) = self.free_nodes.pop() {
1496            self.reset_node_slot(idx);
1497            idx
1498        } else {
1499            let idx = NodeIndex::from_usize(self.nodes.len());
1500            self.nodes.push(CtNode {
1501                children: [ChildRef::NONE, ChildRef::NONE],
1502                log_prob_kt: 0.0,
1503                log_prob_weighted: 0.0,
1504                symbol_count: [0, 0],
1505            });
1506            idx
1507        }
1508    }
1509
1510    #[inline(always)]
1511    fn alloc_node_with_state(&mut self, symbol_count: [u32; 2], log_prob_kt: f64) -> NodeIndex {
1512        let idx = self.alloc_node();
1513        self.nodes[idx.get()].symbol_count = symbol_count;
1514        self.nodes[idx.get()].log_prob_kt = log_prob_kt;
1515        idx
1516    }
1517
1518    #[inline(always)]
1519    fn free_node(&mut self, idx: NodeIndex) {
1520        self.free_nodes.push(idx);
1521    }
1522
1523    #[inline(always)]
1524    fn alloc_segment(&mut self) -> SegmentIndex {
1525        if let Some(idx) = self.free_segments.pop() {
1526            self.reset_segment_slot(idx);
1527            idx
1528        } else {
1529            let idx = SegmentIndex::from_usize(self.segments.len());
1530            self.segments.push(CtSegment::default());
1531            idx
1532        }
1533    }
1534
1535    #[inline(always)]
1536    fn free_segment(&mut self, idx: SegmentIndex) {
1537        self.reset_segment_slot(idx);
1538        self.free_segments.push(idx);
1539    }
1540
1541    /// Drop all nodes/segments and free-list bookkeeping.
1542    pub fn clear(&mut self) {
1543        self.nodes.clear();
1544        self.segments.clear();
1545        self.free_nodes.clear();
1546        self.free_segments.clear();
1547    }
1548
1549    #[inline(always)]
1550    fn child(&self, parent_idx: NodeIndex, child_idx: usize) -> ChildRef {
1551        debug_assert!(parent_idx.get() < self.nodes.len());
1552        debug_assert!(child_idx < 2);
1553        unsafe {
1554            *self
1555                .nodes
1556                .get_unchecked(parent_idx.get())
1557                .children
1558                .get_unchecked(child_idx)
1559        }
1560    }
1561
1562    #[inline(always)]
1563    fn set_child(&mut self, parent_idx: NodeIndex, child_idx: usize, child: ChildRef) {
1564        debug_assert!(parent_idx.get() < self.nodes.len());
1565        debug_assert!(child_idx < 2);
1566        unsafe {
1567            *self
1568                .nodes
1569                .get_unchecked_mut(parent_idx.get())
1570                .children
1571                .get_unchecked_mut(child_idx) = child;
1572        }
1573    }
1574
1575    #[inline(always)]
1576    fn set_segment_tail(&mut self, segment_idx: SegmentIndex, child: ChildRef) {
1577        self.segments[segment_idx.get()].tail = child;
1578    }
1579
1580    #[inline(always)]
1581    fn counts(&self, idx: NodeIndex) -> [u32; 2] {
1582        self.nodes[idx.get()].symbol_count
1583    }
1584
1585    #[inline(always)]
1586    fn visits(&self, idx: NodeIndex) -> u32 {
1587        let counts = self.nodes[idx.get()].symbol_count;
1588        counts[0] + counts[1]
1589    }
1590
1591    #[inline(always)]
1592    fn segment_symbol_count(&self, segment_idx: SegmentIndex) -> [u32; 2] {
1593        self.segments[segment_idx.get()].symbol_count
1594    }
1595
1596    #[inline(always)]
1597    fn segment_log_prob_kt(&self, segment_idx: SegmentIndex) -> f64 {
1598        self.segments[segment_idx.get()].log_prob_kt
1599    }
1600
1601    #[inline(always)]
1602    fn segment_len(&self, segment_idx: SegmentIndex) -> u32 {
1603        self.segments[segment_idx.get()].len()
1604    }
1605
1606    #[inline(always)]
1607    fn segment_has_child(&self, segment_idx: SegmentIndex, offset: u32) -> bool {
1608        let segment = self.segments[segment_idx.get()];
1609        offset + 1 < segment.len() || segment.tail.is_some()
1610    }
1611
1612    #[inline(always)]
1613    fn log_prob_weighted(&self, idx: NodeIndex) -> f64 {
1614        self.nodes[idx.get()].log_prob_weighted
1615    }
1616
1617    #[inline(always)]
1618    fn log_prob_kt(&self, idx: NodeIndex) -> f64 {
1619        self.nodes[idx.get()].log_prob_kt
1620    }
1621
1622    #[inline(always)]
1623    unsafe fn child_ref_weighted_unchecked(&self, child: ChildRef) -> f64 {
1624        if child.is_none() {
1625            return 0.0;
1626        }
1627
1628        let raw = child.0;
1629        if (raw & CHILD_SEGMENT_TAG) == 0 {
1630            debug_assert!((raw as usize) < self.nodes.len());
1631            self.nodes.get_unchecked(raw as usize).log_prob_weighted
1632        } else {
1633            let idx = (raw & CHILD_INDEX_MASK) as usize;
1634            debug_assert!(idx < self.segments.len());
1635            self.segments.get_unchecked(idx).head_log_prob_weighted
1636        }
1637    }
1638
1639    #[inline(always)]
1640    fn child_ref_weighted(&self, child: ChildRef) -> f64 {
1641        unsafe { self.child_ref_weighted_unchecked(child) }
1642    }
1643
1644    #[inline(always)]
1645    fn singleton_segment_payload(&self, edge: usize) -> SegmentPayload {
1646        SegmentPayload::exact((edge & 1) as u64, 1)
1647    }
1648
1649    #[inline(always)]
1650    fn segment_edge(
1651        &self,
1652        segment_idx: SegmentIndex,
1653        offset: u32,
1654        history: &(impl HistoryAccess + ?Sized),
1655    ) -> usize {
1656        let segment = self.segments[segment_idx.get()];
1657        segment_edge_from_parts(segment, offset as usize, history, history.len()) as usize
1658    }
1659
1660    fn segment_suffix_weight(&self, segment_idx: SegmentIndex, offset: u32) -> f64 {
1661        let segment = self.segments[segment_idx.get()];
1662        if offset >= segment.len() {
1663            return self.child_ref_weighted(segment.tail);
1664        }
1665        if segment.tail.is_none() {
1666            return segment.log_prob_kt;
1667        }
1668        let remaining = segment.len() - offset;
1669        unary_chain_log_weight(
1670            segment.log_prob_kt,
1671            self.child_ref_weighted(segment.tail),
1672            remaining,
1673        )
1674    }
1675
1676    #[inline(always)]
1677    fn segment_continuation_weight(&self, segment_idx: SegmentIndex, offset: u32) -> f64 {
1678        let segment = self.segments[segment_idx.get()];
1679        if offset + 1 < segment.len() {
1680            self.segment_suffix_weight(segment_idx, offset + 1)
1681        } else {
1682            self.child_ref_weighted(segment.tail)
1683        }
1684    }
1685
1686    fn recompute_segment_head(&mut self, segment_idx: SegmentIndex) {
1687        let segment = self.segments[segment_idx.get()];
1688        let head = if segment.tail.is_some() {
1689            unary_chain_log_weight(
1690                segment.log_prob_kt,
1691                self.child_ref_weighted(segment.tail),
1692                segment.len(),
1693            )
1694        } else {
1695            segment.log_prob_kt
1696        };
1697        self.segments[segment_idx.get()].head_log_prob_weighted = head;
1698    }
1699
1700    fn recompute_node_weight(&mut self, idx: NodeIndex) {
1701        let slot = idx.get();
1702        debug_assert!(slot < self.nodes.len());
1703        let node = unsafe { *self.nodes.get_unchecked(slot) };
1704        let [left, right] = node.children;
1705        let weighted = if left.is_none() && right.is_none() {
1706            clamp_log_prob(node.log_prob_kt)
1707        } else {
1708            let w0 = unsafe { self.child_ref_weighted_unchecked(left) };
1709            let w1 = unsafe { self.child_ref_weighted_unchecked(right) };
1710            update_weighted_log_prob_non_leaf(node.log_prob_kt, w0, w1)
1711        };
1712        unsafe {
1713            self.nodes.get_unchecked_mut(slot).log_prob_weighted = weighted;
1714        }
1715    }
1716
1717    fn alloc_segment_with_parts(
1718        &mut self,
1719        symbol_count: [u32; 2],
1720        log_prob_kt: f64,
1721        tail: ChildRef,
1722        payload: SegmentPayload,
1723    ) -> SegmentIndex {
1724        let segment_idx = self.alloc_segment();
1725        self.segments[segment_idx.get()] = CtSegment {
1726            tail,
1727            log_prob_kt,
1728            head_log_prob_weighted: 0.0,
1729            symbol_count,
1730            payload,
1731        };
1732        if payload.len() == 1 && tail.is_none() {
1733            self.segments[segment_idx.get()].head_log_prob_weighted = log_prob_kt;
1734        } else {
1735            self.recompute_segment_head(segment_idx);
1736        }
1737        segment_idx
1738    }
1739
1740    fn detach_segment_continuation(
1741        &mut self,
1742        segment_idx: SegmentIndex,
1743        offset: u32,
1744        detaches: &mut Vec<Detach>,
1745    ) -> ChildRef {
1746        let segment = self.segments[segment_idx.get()];
1747        if offset + 1 < segment.len() {
1748            let suffix = self.alloc_segment_with_parts(
1749                segment.symbol_count,
1750                segment.log_prob_kt,
1751                segment.tail,
1752                segment.payload.suffix_after(offset + 1),
1753            );
1754            detaches.push(Detach::SegmentNext {
1755                segment: segment_idx,
1756                new_len: offset + 1,
1757            });
1758            ChildRef::from_segment(suffix)
1759        } else {
1760            let tail = segment.tail;
1761            if tail.is_some() {
1762                detaches.push(Detach::SegmentNext {
1763                    segment: segment_idx,
1764                    new_len: segment.len(),
1765                });
1766            }
1767            tail
1768        }
1769    }
1770
1771    // These fields are the exact segment state plus insertion context. Keeping
1772    // them as scalar arguments avoids building a transient descriptor on this
1773    // path-compression hot path.
1774    #[allow(clippy::too_many_arguments)]
1775    fn prepend_or_alloc_segment(
1776        &mut self,
1777        history: &(impl HistoryAccess + ?Sized),
1778        depth: usize,
1779        symbol_count: [u32; 2],
1780        log_prob_kt: f64,
1781        child: ChildRef,
1782        edge: usize,
1783        allow_history_pattern: bool,
1784    ) -> ChildRef {
1785        let singleton_payload = self.singleton_segment_payload(edge);
1786
1787        if let Some(segment_idx) = child.as_segment() {
1788            let segment = self.segments[segment_idx.get()];
1789            let same_state = segment.symbol_count == symbol_count
1790                && segment.log_prob_kt.to_bits() == log_prob_kt.to_bits();
1791            if same_state && segment.tail == child {
1792                let segment = &mut self.segments[segment_idx.get()];
1793                let extended_payload = if segment.payload.is_exact() {
1794                    segment.payload.prepend_exact(edge)
1795                } else if allow_history_pattern {
1796                    let path_payload =
1797                        SegmentPayload::from_path(history, depth, segment.len().saturating_add(1));
1798                    path_payload.filter(|payload| {
1799                        let mut matches = true;
1800                        for offset in 0..segment.len() as usize {
1801                            let seg_edge =
1802                                segment_edge_from_parts(*segment, offset, history, history.len());
1803                            let payload_edge = ((payload.exact_bits() >> (offset + 1)) & 1) != 0;
1804                            if seg_edge != payload_edge {
1805                                matches = false;
1806                                break;
1807                            }
1808                        }
1809                        matches
1810                    })
1811                } else {
1812                    None
1813                };
1814                if let Some(payload) = extended_payload {
1815                    let old_head = segment.head_log_prob_weighted;
1816                    segment.payload = payload;
1817                    segment.head_log_prob_weighted =
1818                        update_weighted_log_prob(log_prob_kt, old_head, 0.0, false);
1819                    return ChildRef::from_segment(segment_idx);
1820                }
1821            }
1822        }
1823
1824        let segment_idx =
1825            self.alloc_segment_with_parts(symbol_count, log_prob_kt, child, singleton_payload);
1826        ChildRef::from_segment(segment_idx)
1827    }
1828
1829    fn free_child_ref(&mut self, child: ChildRef) {
1830        let mut stack = Vec::with_capacity(16);
1831        if child.is_some() {
1832            stack.push(child);
1833        }
1834        while let Some(next) = stack.pop() {
1835            if let Some(node_idx) = next.as_node() {
1836                let children = self.nodes[node_idx.get()].children;
1837                if children[0].is_some() {
1838                    stack.push(children[0]);
1839                }
1840                if children[1].is_some() {
1841                    stack.push(children[1]);
1842                }
1843                self.free_node(node_idx);
1844            } else if let Some(segment_idx) = next.as_segment() {
1845                let tail = self.segments[segment_idx.get()].tail;
1846                if tail.is_some() {
1847                    stack.push(tail);
1848                }
1849                self.free_segment(segment_idx);
1850            }
1851        }
1852    }
1853
1854    /// Approximate heap usage (bytes) for arena-owned storage.
1855    pub fn memory_usage(&self) -> usize {
1856        self.nodes.capacity() * size_of::<CtNode>()
1857            + self.segments.capacity() * size_of::<CtSegment>()
1858            + self.free_nodes.capacity() * size_of::<NodeIndex>()
1859            + self.free_segments.capacity() * size_of::<SegmentIndex>()
1860    }
1861}
1862
1863impl Default for CtArena {
1864    fn default() -> Self {
1865        Self::new()
1866    }
1867}
1868
1869#[derive(Clone)]
1870struct CtEngine {
1871    arena: CtArena,
1872    root: NodeIndex,
1873    max_depth: usize,
1874    segment_alpha: Vec<f64>,
1875    segment_log_alpha: Vec<f64>,
1876    segment_log_one_minus_alpha: Vec<f64>,
1877    levels: Vec<LevelState>,
1878    detaches: Vec<Detach>,
1879    prepared_steps: Vec<PreparedStep>,
1880    prepared_levels: usize,
1881    prepared_end: PreparedEnd,
1882}
1883
1884impl CtEngine {
1885    const RESERVE_MIN_NODES: usize = 4 * 1024;
1886    const RESERVE_MAX_NODES: usize = 1 << 18;
1887
1888    fn new(depth: usize) -> Self {
1889        let mut arena = CtArena::with_capacity(1024.min(1 << depth.min(16)));
1890        let root = arena.alloc_node();
1891        let mut segment_alpha = Vec::with_capacity(depth + 1);
1892        let mut segment_log_alpha = Vec::with_capacity(depth + 1);
1893        let mut segment_log_one_minus_alpha = Vec::with_capacity(depth + 1);
1894        segment_alpha.push(1.0);
1895        segment_log_alpha.push(0.0);
1896        segment_log_one_minus_alpha.push(f64::NEG_INFINITY);
1897        let mut alpha = 1.0f64;
1898        for len in 1..=depth {
1899            alpha *= 0.5;
1900            segment_alpha.push(alpha);
1901            segment_log_alpha.push(-(len as f64) * std::f64::consts::LN_2);
1902            segment_log_one_minus_alpha.push((-alpha).ln_1p());
1903        }
1904        Self {
1905            arena,
1906            root,
1907            max_depth: depth,
1908            segment_alpha,
1909            segment_log_alpha,
1910            segment_log_one_minus_alpha,
1911            levels: vec![LevelState::default(); depth],
1912            detaches: Vec::with_capacity(depth),
1913            prepared_steps: Vec::with_capacity(depth),
1914            prepared_levels: 0,
1915            prepared_end: PreparedEnd::MaxDepth,
1916        }
1917    }
1918
1919    #[inline(always)]
1920    fn root_visits(&self) -> usize {
1921        self.arena.visits(self.root) as usize
1922    }
1923
1924    #[inline(always)]
1925    fn hot_prefix_depth(&self) -> usize {
1926        self.max_depth.min(ctw_hot_prefix_depth_limit())
1927    }
1928
1929    #[inline(always)]
1930    fn push_prepared_segment_step(
1931        &mut self,
1932        segment_idx: SegmentIndex,
1933        offset: usize,
1934        sibling_weight: f64,
1935        has_sibling: u8,
1936    ) {
1937        let span = (offset + 1) as u32;
1938        self.prepared_steps.push(PreparedStep {
1939            source: ExistingSource::Segment(segment_idx, offset as u32),
1940            span,
1941            sibling_weight,
1942            has_sibling,
1943        });
1944        self.prepared_levels += span as usize;
1945    }
1946
1947    #[inline(always)]
1948    fn walk_prepared_exact_segment(
1949        &mut self,
1950        segment_idx: SegmentIndex,
1951        segment: CtSegment,
1952        depth: usize,
1953        path_bits: u64,
1954    ) -> Option<(usize, ExistingSource)> {
1955        let seg_len = segment.len() as usize;
1956        let terminal_offset = self.max_depth.saturating_sub(depth);
1957        let comparable_len = seg_len.min(terminal_offset);
1958        if let Some((offset, _, _)) =
1959            first_exact_segment_mismatch(segment.payload.exact_bits(), path_bits, comparable_len)
1960        {
1961            self.push_prepared_segment_step(
1962                segment_idx,
1963                offset,
1964                self.arena
1965                    .segment_continuation_weight(segment_idx, offset as u32),
1966                1,
1967            );
1968            self.prepared_end = PreparedEnd::MismatchAtCurrentSegment;
1969            return None;
1970        }
1971
1972        if terminal_offset < seg_len {
1973            self.push_prepared_segment_step(segment_idx, terminal_offset, 0.0, 0);
1974            return None;
1975        }
1976
1977        if segment.tail.is_none() {
1978            self.push_prepared_segment_step(segment_idx, seg_len - 1, 0.0, 0);
1979            self.prepared_end = PreparedEnd::MissingAfterCurrent;
1980            return None;
1981        }
1982
1983        self.push_prepared_segment_step(segment_idx, seg_len - 1, 0.0, 0);
1984        let tail = segment.tail;
1985        Some((
1986            depth + seg_len,
1987            Self::child_to_existing_source(tail).unwrap_or(ExistingSource::None),
1988        ))
1989    }
1990
1991    fn clear(&mut self) {
1992        self.arena.clear();
1993        self.root = self.arena.alloc_node();
1994        self.levels.fill(LevelState::default());
1995        self.detaches.clear();
1996        self.prepared_steps.clear();
1997        self.prepared_levels = 0;
1998        self.prepared_end = PreparedEnd::MaxDepth;
1999    }
2000
2001    #[inline]
2002    fn reserve_for_symbols(&mut self, total_symbols: usize) {
2003        if total_symbols == 0 {
2004            return;
2005        }
2006
2007        let depth_scale = self.max_depth.saturating_add(1);
2008        let reserve_nodes = total_symbols
2009            .saturating_div(depth_scale)
2010            .clamp(Self::RESERVE_MIN_NODES, Self::RESERVE_MAX_NODES);
2011        let free_nodes = self
2012            .arena
2013            .nodes
2014            .capacity()
2015            .saturating_sub(self.arena.nodes.len());
2016        if reserve_nodes > free_nodes {
2017            self.arena.reserve_exact(reserve_nodes - free_nodes);
2018        }
2019    }
2020
2021    #[inline]
2022    fn get_log_block_probability(&self) -> f64 {
2023        self.arena.log_prob_weighted(self.root)
2024    }
2025
2026    #[inline]
2027    fn with_cached_logs<R>(
2028        &mut self,
2029        upto: usize,
2030        f: impl FnOnce(&mut Self, CachedLogs<'_>) -> R,
2031    ) -> R {
2032        with_shared_cached_logs(upto, |logs| f(self, logs))
2033    }
2034
2035    #[inline]
2036    fn with_bounded_logs<R>(
2037        &mut self,
2038        upto: usize,
2039        f: impl FnOnce(&mut Self, BoundedLogs<'_>) -> R,
2040    ) -> R {
2041        with_shared_bounded_logs(upto, |logs| f(self, logs))
2042    }
2043
2044    #[inline]
2045    fn log_cache_memory_usage(&self) -> usize {
2046        shared_log_cache_memory_usage()
2047    }
2048
2049    #[inline(always)]
2050    fn segment_constants(&self, len: u32) -> (f64, f64, f64) {
2051        let idx = len as usize;
2052        debug_assert!(idx < self.segment_alpha.len());
2053        debug_assert!(idx < self.segment_log_alpha.len());
2054        debug_assert!(idx < self.segment_log_one_minus_alpha.len());
2055        (
2056            unsafe { *self.segment_alpha.get_unchecked(idx) },
2057            unsafe { *self.segment_log_alpha.get_unchecked(idx) },
2058            unsafe { *self.segment_log_one_minus_alpha.get_unchecked(idx) },
2059        )
2060    }
2061
2062    #[inline(always)]
2063    fn source_counts_and_kt_log_prob(&self, source: ExistingSource) -> ([u32; 2], f64) {
2064        match source {
2065            ExistingSource::Node(node_idx) => {
2066                let slot = node_idx.get();
2067                let node = unsafe { *self.arena.nodes.get_unchecked(slot) };
2068                (node.symbol_count, node.log_prob_kt)
2069            }
2070            ExistingSource::Segment(segment_idx, _) => {
2071                let slot = segment_idx.get();
2072                let segment = unsafe { *self.arena.segments.get_unchecked(slot) };
2073                (segment.symbol_count, segment.log_prob_kt)
2074            }
2075            ExistingSource::None => unreachable!("prepared step should never store None"),
2076        }
2077    }
2078
2079    fn build_missing_segment_path(
2080        &mut self,
2081        depth: usize,
2082        history: &(impl HistoryAccess + ?Sized),
2083        sym_idx: usize,
2084        singleton_log_prob_kt: f64,
2085    ) -> ChildRef {
2086        if depth > self.max_depth {
2087            return ChildRef::NONE;
2088        }
2089
2090        let mut counts = [0u32; 2];
2091        counts[sym_idx] = 1;
2092        let log_prob_kt = singleton_log_prob_kt;
2093        let total_len = self.max_depth - depth + 1;
2094
2095        if let Some(payload) = SegmentPayload::from_path(history, depth, total_len as u32) {
2096            let segment =
2097                self.arena
2098                    .alloc_segment_with_parts(counts, log_prob_kt, ChildRef::NONE, payload);
2099            return ChildRef::from_segment(segment);
2100        }
2101
2102        let history_nodes = if depth < history.len() {
2103            (self.max_depth.min(history.len() - 1) - depth) + 1
2104        } else {
2105            0
2106        };
2107        let const_nodes = total_len - history_nodes;
2108
2109        let mut built = ChildRef::NONE;
2110        if const_nodes > 0 {
2111            let const_segment = self.arena.alloc_segment_with_parts(
2112                counts,
2113                log_prob_kt,
2114                ChildRef::NONE,
2115                SegmentPayload::constant(false, const_nodes as u32),
2116            );
2117            built = ChildRef::from_segment(const_segment);
2118        }
2119        if history_nodes > 0 {
2120            let history_segment = self.arena.alloc_segment_with_parts(
2121                counts,
2122                log_prob_kt,
2123                built,
2124                SegmentPayload::history(
2125                    (history.len() - depth - 1) as u32,
2126                    history_nodes as u32,
2127                    false,
2128                ),
2129            );
2130            built = ChildRef::from_segment(history_segment);
2131        }
2132        built
2133    }
2134
2135    fn build_missing_path(
2136        &mut self,
2137        depth: usize,
2138        history: &(impl HistoryAccess + ?Sized),
2139        sym_idx: usize,
2140        singleton_log_prob_kt: f64,
2141    ) -> ChildRef {
2142        if depth > self.max_depth {
2143            return ChildRef::NONE;
2144        }
2145
2146        let hot_prefix_depth = self.hot_prefix_depth();
2147        if depth > hot_prefix_depth {
2148            return self.build_missing_segment_path(depth, history, sym_idx, singleton_log_prob_kt);
2149        }
2150
2151        let mut counts = [0u32; 2];
2152        counts[sym_idx] = 1;
2153        let mut built = if hot_prefix_depth < self.max_depth {
2154            self.build_missing_segment_path(
2155                hot_prefix_depth + 1,
2156                history,
2157                sym_idx,
2158                singleton_log_prob_kt,
2159            )
2160        } else {
2161            ChildRef::NONE
2162        };
2163
2164        for node_depth in (depth..=hot_prefix_depth).rev() {
2165            let node = self
2166                .arena
2167                .alloc_node_with_state(counts, singleton_log_prob_kt);
2168            if node_depth < self.max_depth {
2169                let edge = history_symbol(history, node_depth) as usize;
2170                self.arena.set_child(node, edge, built);
2171            }
2172            self.arena.recompute_node_weight(node);
2173            built = ChildRef::from_node(node);
2174        }
2175        built
2176    }
2177
2178    #[inline(always)]
2179    fn build_missing_segment_path_exact_bits(
2180        &mut self,
2181        depth: usize,
2182        path_bits: u64,
2183        sym_idx: usize,
2184        singleton_log_prob_kt: f64,
2185    ) -> ChildRef {
2186        debug_assert!(self.max_depth <= SEG_EXACT_MAX_LEN as usize);
2187        if depth > self.max_depth {
2188            return ChildRef::NONE;
2189        }
2190
2191        let mut counts = [0u32; 2];
2192        counts[sym_idx] = 1;
2193        let total_len = self.max_depth - depth + 1;
2194        let payload = SegmentPayload::exact(
2195            path_bits & low_bits_mask_u64(total_len as u32),
2196            total_len as u32,
2197        );
2198        let segment = self.arena.alloc_segment_with_parts(
2199            counts,
2200            singleton_log_prob_kt,
2201            ChildRef::NONE,
2202            payload,
2203        );
2204        ChildRef::from_segment(segment)
2205    }
2206
2207    #[inline(always)]
2208    fn build_missing_path_exact_bits(
2209        &mut self,
2210        depth: usize,
2211        path_bits: u64,
2212        sym_idx: usize,
2213        singleton_log_prob_kt: f64,
2214    ) -> ChildRef {
2215        debug_assert!(self.max_depth <= SEG_EXACT_MAX_LEN as usize);
2216        if depth > self.max_depth {
2217            return ChildRef::NONE;
2218        }
2219
2220        let hot_prefix_depth = self.hot_prefix_depth();
2221        if depth > hot_prefix_depth {
2222            return self.build_missing_segment_path_exact_bits(
2223                depth,
2224                path_bits,
2225                sym_idx,
2226                singleton_log_prob_kt,
2227            );
2228        }
2229
2230        let mut counts = [0u32; 2];
2231        counts[sym_idx] = 1;
2232        let mut built = if hot_prefix_depth < self.max_depth {
2233            self.build_missing_segment_path_exact_bits(
2234                hot_prefix_depth + 1,
2235                shift_path_bits(path_bits, hot_prefix_depth + 1 - depth),
2236                sym_idx,
2237                singleton_log_prob_kt,
2238            )
2239        } else {
2240            ChildRef::NONE
2241        };
2242
2243        for node_depth in (depth..=hot_prefix_depth).rev() {
2244            let node = self
2245                .arena
2246                .alloc_node_with_state(counts, singleton_log_prob_kt);
2247            if node_depth < self.max_depth {
2248                let edge = ((path_bits >> (node_depth - depth)) & 1) as usize;
2249                self.arena.set_child(node, edge, built);
2250            }
2251            self.arena.recompute_node_weight(node);
2252            built = ChildRef::from_node(node);
2253        }
2254        built
2255    }
2256
2257    #[inline(always)]
2258    fn child_to_existing_source(child: ChildRef) -> Option<ExistingSource> {
2259        if let Some(node) = child.as_node() {
2260            Some(ExistingSource::Node(node))
2261        } else {
2262            child
2263                .as_segment()
2264                .map(|segment| ExistingSource::Segment(segment, 0))
2265        }
2266    }
2267
2268    #[inline(always)]
2269    fn update_source_state<L: CtLogAccess>(
2270        &mut self,
2271        logs: L,
2272        source: ExistingSource,
2273        sym_idx: usize,
2274    ) {
2275        match source {
2276            ExistingSource::Node(node_idx) => {
2277                let slot = node_idx.get();
2278                let mut counts = self.arena.nodes[slot].symbol_count;
2279                let mut log_prob_kt = self.arena.nodes[slot].log_prob_kt;
2280                apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
2281                self.arena.nodes[slot].symbol_count = counts;
2282                self.arena.nodes[slot].log_prob_kt = log_prob_kt;
2283            }
2284            ExistingSource::Segment(segment_idx, _) => {
2285                let slot = segment_idx.get();
2286                let mut counts = self.arena.segments[slot].symbol_count;
2287                let mut log_prob_kt = self.arena.segments[slot].log_prob_kt;
2288                apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
2289                self.arena.segments[slot].symbol_count = counts;
2290                self.arena.segments[slot].log_prob_kt = log_prob_kt;
2291            }
2292            ExistingSource::None => unreachable!("prepared update should never visit None"),
2293        }
2294    }
2295
2296    #[inline(always)]
2297    fn recompute_source_weight(&mut self, source: ExistingSource) {
2298        match source {
2299            ExistingSource::Node(node_idx) => self.arena.recompute_node_weight(node_idx),
2300            ExistingSource::Segment(segment_idx, _) => self.recompute_segment_head(segment_idx),
2301            ExistingSource::None => unreachable!("prepared update should never visit None"),
2302        }
2303    }
2304
2305    #[inline(always)]
2306    fn recompute_segment_head(&mut self, segment_idx: SegmentIndex) {
2307        let segment = self.arena.segments[segment_idx.get()];
2308        let head = if segment.tail.is_some() {
2309            let (alpha, log_alpha, log_one_minus_alpha) = self.segment_constants(segment.len());
2310            unary_chain_log_weight_precomputed(
2311                segment.log_prob_kt,
2312                self.arena.child_ref_weighted(segment.tail),
2313                alpha,
2314                log_alpha,
2315                log_one_minus_alpha,
2316            )
2317        } else {
2318            segment.log_prob_kt
2319        };
2320        self.arena.segments[segment_idx.get()].head_log_prob_weighted = head;
2321    }
2322
2323    fn attach_missing_after_prepared_path(
2324        &mut self,
2325        history: &(impl HistoryAccess + ?Sized),
2326        sym_idx: usize,
2327        singleton_log_prob_kt: f64,
2328    ) {
2329        let Some(last_step) = self.prepared_steps.last().copied() else {
2330            return;
2331        };
2332        let depth = self.prepared_levels;
2333        match last_step.source {
2334            ExistingSource::Node(node_idx) => {
2335                debug_assert!(depth < self.max_depth);
2336                let path_edge = history_symbol(history, depth) as usize;
2337                debug_assert!(self.arena.child(node_idx, path_edge).is_none());
2338                let new_child =
2339                    self.build_missing_path(depth + 1, history, sym_idx, singleton_log_prob_kt);
2340                self.arena.set_child(node_idx, path_edge, new_child);
2341            }
2342            ExistingSource::Segment(segment_idx, offset) => {
2343                debug_assert!(depth < self.max_depth);
2344                debug_assert_eq!(offset + 1, self.arena.segment_len(segment_idx));
2345                debug_assert!(self.arena.segments[segment_idx.get()].tail.is_none());
2346                let new_tail =
2347                    self.build_missing_path(depth + 1, history, sym_idx, singleton_log_prob_kt);
2348                self.arena.set_segment_tail(segment_idx, new_tail);
2349            }
2350            ExistingSource::None => unreachable!("prepared path should never end in None source"),
2351        }
2352    }
2353
2354    fn replace_prepared_child(
2355        &mut self,
2356        history: &(impl HistoryAccess + ?Sized),
2357        step_index: usize,
2358        current_start_depth: usize,
2359        new_child: ChildRef,
2360    ) {
2361        if step_index == 0 {
2362            let root_edge = history_symbol(history, 0) as usize;
2363            self.arena.set_child(self.root, root_edge, new_child);
2364            return;
2365        }
2366
2367        match self.prepared_steps[step_index - 1].source {
2368            ExistingSource::Node(node_idx) => {
2369                let edge = history_symbol(history, current_start_depth - 1) as usize;
2370                self.arena.set_child(node_idx, edge, new_child);
2371            }
2372            ExistingSource::Segment(segment_idx, offset) => {
2373                debug_assert_eq!(offset + 1, self.arena.segment_len(segment_idx));
2374                self.arena.set_segment_tail(segment_idx, new_child);
2375            }
2376            ExistingSource::None => unreachable!("prepared path should never parent from None"),
2377        }
2378    }
2379
2380    fn update_prepared_mismatch<L: CtLogAccess>(
2381        &mut self,
2382        logs: L,
2383        history: &(impl HistoryAccess + ?Sized),
2384        sym_idx: usize,
2385        singleton_log_prob_kt: f64,
2386    ) -> ChildRef {
2387        let last_index = self.prepared_steps.len() - 1;
2388        for idx in 0..last_index {
2389            self.update_source_state(logs, self.prepared_steps[idx].source, sym_idx);
2390        }
2391
2392        let last_step = self.prepared_steps[last_index];
2393        let ExistingSource::Segment(segment_idx, offset_u32) = last_step.source else {
2394            unreachable!("prepared segment mismatch must end at a segment");
2395        };
2396
2397        let original = self.arena.segments[segment_idx.get()];
2398        let offset = offset_u32 as usize;
2399        let seg_len = original.len() as usize;
2400        let history_len = history.len();
2401        let current_start_depth = self.prepared_levels - last_step.span as usize + 1;
2402        let node_depth = current_start_depth + offset;
2403        let path_edge = path_edge_at_depth(history, history_len, node_depth);
2404        let existing_edge = segment_edge_from_parts(original, offset, history, history_len);
2405        debug_assert_ne!(path_edge, existing_edge);
2406
2407        let old_continuation = if offset + 1 < seg_len {
2408            if offset == 0 {
2409                let segment = &mut self.arena.segments[segment_idx.get()];
2410                segment.payload = original.payload.suffix_after(1);
2411                segment.tail = original.tail;
2412                segment.symbol_count = original.symbol_count;
2413                segment.log_prob_kt = original.log_prob_kt;
2414                self.recompute_segment_head(segment_idx);
2415                ChildRef::from_segment(segment_idx)
2416            } else {
2417                ChildRef::from_segment(self.arena.alloc_segment_with_parts(
2418                    original.symbol_count,
2419                    original.log_prob_kt,
2420                    original.tail,
2421                    original.payload.suffix_after(offset as u32 + 1),
2422                ))
2423            }
2424        } else {
2425            original.tail
2426        };
2427
2428        let new_tail =
2429            self.build_missing_path(node_depth + 1, history, sym_idx, singleton_log_prob_kt);
2430        let mut updated_counts = original.symbol_count;
2431        let mut updated_log_prob_kt = original.log_prob_kt;
2432        apply_update_to_state_raw(logs, &mut updated_counts, &mut updated_log_prob_kt, sym_idx);
2433
2434        let branch = self
2435            .arena
2436            .alloc_node_with_state(updated_counts, updated_log_prob_kt);
2437        self.arena
2438            .set_child(branch, existing_edge as usize, old_continuation);
2439        self.arena.set_child(branch, path_edge as usize, new_tail);
2440        self.arena.recompute_node_weight(branch);
2441
2442        if offset == 0 {
2443            if seg_len == 1 {
2444                self.arena.free_segment(segment_idx);
2445            }
2446            self.replace_prepared_child(
2447                history,
2448                last_index,
2449                current_start_depth,
2450                ChildRef::from_node(branch),
2451            );
2452        } else {
2453            let segment = &mut self.arena.segments[segment_idx.get()];
2454            segment.payload = original.payload.prefix(offset as u32);
2455            segment.tail = ChildRef::from_node(branch);
2456            segment.symbol_count = updated_counts;
2457            segment.log_prob_kt = updated_log_prob_kt;
2458            self.recompute_segment_head(segment_idx);
2459        }
2460
2461        for idx in (0..last_index).rev() {
2462            self.recompute_source_weight(self.prepared_steps[idx].source);
2463        }
2464
2465        let root_edge = history_symbol(history, 0) as usize;
2466        self.arena.child(self.root, root_edge)
2467    }
2468
2469    fn update_prepared_cached_path<L: CtLogAccess>(
2470        &mut self,
2471        logs: L,
2472        history: &(impl HistoryAccess + ?Sized),
2473        sym_idx: usize,
2474        singleton_log_prob_kt: f64,
2475    ) {
2476        debug_assert!(!self.prepared_steps.is_empty());
2477        debug_assert!(matches!(
2478            self.prepared_end,
2479            PreparedEnd::MaxDepth | PreparedEnd::MissingAfterCurrent
2480        ));
2481
2482        if self.prepared_end == PreparedEnd::MissingAfterCurrent {
2483            self.attach_missing_after_prepared_path(history, sym_idx, singleton_log_prob_kt);
2484        }
2485
2486        let last_index = self.prepared_steps.len() - 1;
2487        let mut child_weight = if self.prepared_end == PreparedEnd::MissingAfterCurrent {
2488            let last_step = self.prepared_steps[last_index];
2489            match last_step.source {
2490                ExistingSource::Node(node_idx) => {
2491                    let depth = self.prepared_levels;
2492                    let edge = history_symbol(history, depth) as usize;
2493                    self.arena
2494                        .child_ref_weighted(self.arena.child(node_idx, edge))
2495                }
2496                ExistingSource::Segment(segment_idx, offset) => {
2497                    debug_assert_eq!(offset + 1, self.arena.segment_len(segment_idx));
2498                    self.arena
2499                        .child_ref_weighted(self.arena.segments[segment_idx.get()].tail)
2500                }
2501                ExistingSource::None => unreachable!("prepared path should never end in None"),
2502            }
2503        } else {
2504            0.0
2505        };
2506
2507        for idx in (0..=last_index).rev() {
2508            let step = self.prepared_steps[idx];
2509            match step.source {
2510                ExistingSource::Node(node_idx) => {
2511                    let (mut counts, mut log_prob_kt) =
2512                        self.source_counts_and_kt_log_prob(step.source);
2513                    apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
2514                    let weighted =
2515                        if idx == last_index && self.prepared_end == PreparedEnd::MaxDepth {
2516                            debug_assert_eq!(step.has_sibling, 0);
2517                            clamp_log_prob(log_prob_kt)
2518                        } else {
2519                            update_weighted_log_prob(
2520                                log_prob_kt,
2521                                child_weight,
2522                                step.sibling_weight,
2523                                false,
2524                            )
2525                        };
2526                    let slot = node_idx.get();
2527                    self.arena.nodes[slot].symbol_count = counts;
2528                    self.arena.nodes[slot].log_prob_kt = log_prob_kt;
2529                    self.arena.nodes[slot].log_prob_weighted = weighted;
2530                    child_weight = weighted;
2531                }
2532                ExistingSource::Segment(segment_idx, offset) => {
2533                    let (mut counts, mut log_prob_kt) =
2534                        self.source_counts_and_kt_log_prob(step.source);
2535                    apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
2536                    let slot = segment_idx.get();
2537                    let weighted =
2538                        if idx == last_index && self.prepared_end == PreparedEnd::MaxDepth {
2539                            debug_assert_eq!(offset + 1, self.arena.segment_len(segment_idx));
2540                            debug_assert!(self.arena.segments[slot].tail.is_none());
2541                            clamp_log_prob(log_prob_kt)
2542                        } else {
2543                            let (alpha, log_alpha, log_one_minus_alpha) =
2544                                self.segment_constants(step.span);
2545                            unary_chain_log_weight_precomputed(
2546                                log_prob_kt,
2547                                child_weight,
2548                                alpha,
2549                                log_alpha,
2550                                log_one_minus_alpha,
2551                            )
2552                        };
2553                    self.arena.segments[slot].symbol_count = counts;
2554                    self.arena.segments[slot].log_prob_kt = log_prob_kt;
2555                    self.arena.segments[slot].head_log_prob_weighted = weighted;
2556                    child_weight = weighted;
2557                }
2558                ExistingSource::None => unreachable!("prepared update should never visit None"),
2559            }
2560        }
2561    }
2562
2563    fn update_child_fast<L: CtLogAccess>(
2564        &mut self,
2565        logs: L,
2566        child: ChildRef,
2567        depth: usize,
2568        history: &(impl HistoryAccess + ?Sized),
2569        sym_idx: usize,
2570        singleton_log_prob_kt: f64,
2571    ) -> ChildRef {
2572        if depth > self.max_depth {
2573            return child;
2574        }
2575        if child.is_none() {
2576            return self.build_missing_path(depth, history, sym_idx, singleton_log_prob_kt);
2577        }
2578
2579        if let Some(node_idx) = child.as_node() {
2580            if depth < self.max_depth {
2581                let path_edge = history_symbol(history, depth) as usize;
2582                let next = self.arena.child(node_idx, path_edge);
2583                let updated = self.update_child_fast(
2584                    logs,
2585                    next,
2586                    depth + 1,
2587                    history,
2588                    sym_idx,
2589                    singleton_log_prob_kt,
2590                );
2591                if updated != next {
2592                    self.arena.set_child(node_idx, path_edge, updated);
2593                }
2594            }
2595            let mut counts = self.arena.nodes[node_idx.get()].symbol_count;
2596            let mut log_prob_kt = self.arena.nodes[node_idx.get()].log_prob_kt;
2597            apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
2598            self.arena.nodes[node_idx.get()].symbol_count = counts;
2599            self.arena.nodes[node_idx.get()].log_prob_kt = log_prob_kt;
2600            self.arena.recompute_node_weight(node_idx);
2601            return ChildRef::from_node(node_idx);
2602        }
2603
2604        let segment_idx = child.as_segment().unwrap();
2605        let original = self.arena.segments[segment_idx.get()];
2606        let seg_len = original.len() as usize;
2607        let mut updated_counts = original.symbol_count;
2608        let mut updated_log_prob_kt = original.log_prob_kt;
2609        apply_update_to_state_raw(logs, &mut updated_counts, &mut updated_log_prob_kt, sym_idx);
2610
2611        let depth_budget = self.max_depth.saturating_sub(depth);
2612        let comparable_len = if original.tail.is_none() {
2613            seg_len.saturating_sub(1)
2614        } else {
2615            seg_len
2616        }
2617        .min(depth_budget);
2618        let mismatch = first_segment_mismatch(original, depth, history, comparable_len).map(
2619            |(offset, path_edge, existing_edge)| (offset, depth + offset, path_edge, existing_edge),
2620        );
2621
2622        if let Some((offset, node_depth, path_edge, existing_edge)) = mismatch {
2623            let old_continuation = if offset + 1 < seg_len {
2624                if offset == 0 {
2625                    let segment = &mut self.arena.segments[segment_idx.get()];
2626                    segment.payload = original.payload.suffix_after(1);
2627                    segment.tail = original.tail;
2628                    segment.symbol_count = original.symbol_count;
2629                    segment.log_prob_kt = original.log_prob_kt;
2630                    self.recompute_segment_head(segment_idx);
2631                    ChildRef::from_segment(segment_idx)
2632                } else {
2633                    ChildRef::from_segment(self.arena.alloc_segment_with_parts(
2634                        original.symbol_count,
2635                        original.log_prob_kt,
2636                        original.tail,
2637                        original.payload.suffix_after(offset as u32 + 1),
2638                    ))
2639                }
2640            } else {
2641                original.tail
2642            };
2643
2644            let new_tail =
2645                self.build_missing_path(node_depth + 1, history, sym_idx, singleton_log_prob_kt);
2646            let branch = self
2647                .arena
2648                .alloc_node_with_state(updated_counts, updated_log_prob_kt);
2649            self.arena
2650                .set_child(branch, existing_edge as usize, old_continuation);
2651            self.arena.set_child(branch, path_edge as usize, new_tail);
2652            self.arena.recompute_node_weight(branch);
2653
2654            if offset == 0 {
2655                if offset + 1 >= seg_len {
2656                    self.arena.free_segment(segment_idx);
2657                }
2658                return ChildRef::from_node(branch);
2659            }
2660
2661            let segment = &mut self.arena.segments[segment_idx.get()];
2662            segment.payload = original.payload.prefix(offset as u32);
2663            segment.tail = ChildRef::from_node(branch);
2664            segment.symbol_count = updated_counts;
2665            segment.log_prob_kt = updated_log_prob_kt;
2666            self.recompute_segment_head(segment_idx);
2667            return ChildRef::from_segment(segment_idx);
2668        }
2669
2670        if depth_budget < seg_len {
2671            self.arena.segments[segment_idx.get()].symbol_count = updated_counts;
2672            self.arena.segments[segment_idx.get()].log_prob_kt = updated_log_prob_kt;
2673            self.recompute_segment_head(segment_idx);
2674            return ChildRef::from_segment(segment_idx);
2675        }
2676
2677        if original.tail.is_none() {
2678            let new_tail =
2679                self.build_missing_path(depth + seg_len, history, sym_idx, singleton_log_prob_kt);
2680            self.arena.segments[segment_idx.get()].tail = new_tail;
2681            self.arena.segments[segment_idx.get()].symbol_count = updated_counts;
2682            self.arena.segments[segment_idx.get()].log_prob_kt = updated_log_prob_kt;
2683            self.recompute_segment_head(segment_idx);
2684            return ChildRef::from_segment(segment_idx);
2685        }
2686
2687        let tail = original.tail;
2688        let updated_tail = self.update_child_fast(
2689            logs,
2690            tail,
2691            depth + seg_len,
2692            history,
2693            sym_idx,
2694            singleton_log_prob_kt,
2695        );
2696        if updated_tail != tail {
2697            self.arena.set_segment_tail(segment_idx, updated_tail);
2698        }
2699        self.arena.segments[segment_idx.get()].symbol_count = updated_counts;
2700        self.arena.segments[segment_idx.get()].log_prob_kt = updated_log_prob_kt;
2701        self.recompute_segment_head(segment_idx);
2702        ChildRef::from_segment(segment_idx)
2703    }
2704
2705    // Exact-mode updates thread together the log table, path state, and KT
2706    // singleton value; a wrapper would only hide the data dependencies in this
2707    // inner CTW update kernel.
2708    #[allow(clippy::too_many_arguments)]
2709    fn update_child_fast_exact<L: CtLogAccess>(
2710        &mut self,
2711        logs: L,
2712        child: ChildRef,
2713        depth: usize,
2714        history: &(impl HistoryAccess + ?Sized),
2715        path_bits: u64,
2716        sym_idx: usize,
2717        singleton_log_prob_kt: f64,
2718    ) -> ChildRef {
2719        debug_assert!(self.max_depth <= SEG_EXACT_MAX_LEN as usize);
2720        if depth > self.max_depth {
2721            return child;
2722        }
2723        if child.is_none() {
2724            return self.build_missing_path_exact_bits(
2725                depth,
2726                path_bits,
2727                sym_idx,
2728                singleton_log_prob_kt,
2729            );
2730        }
2731
2732        if let Some(node_idx) = child.as_node() {
2733            if depth < self.max_depth {
2734                let path_edge = (path_bits & 1) as usize;
2735                let next = self.arena.child(node_idx, path_edge);
2736                let updated = self.update_child_fast_exact(
2737                    logs,
2738                    next,
2739                    depth + 1,
2740                    history,
2741                    shift_path_bits(path_bits, 1),
2742                    sym_idx,
2743                    singleton_log_prob_kt,
2744                );
2745                if updated != next {
2746                    self.arena.set_child(node_idx, path_edge, updated);
2747                }
2748            }
2749            let slot = node_idx.get();
2750            let mut counts = self.arena.nodes[slot].symbol_count;
2751            let mut log_prob_kt = self.arena.nodes[slot].log_prob_kt;
2752            apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
2753            let [left, right] = self.arena.nodes[slot].children;
2754            let weighted = if left.is_none() && right.is_none() {
2755                clamp_log_prob(log_prob_kt)
2756            } else {
2757                // Safety: `left`/`right` come from this node's stored children, so any
2758                // non-none child index is arena-owned and in-bounds for this arena.
2759                let w0 = unsafe { self.arena.child_ref_weighted_unchecked(left) };
2760                let w1 = unsafe { self.arena.child_ref_weighted_unchecked(right) };
2761                update_weighted_log_prob_non_leaf(log_prob_kt, w0, w1)
2762            };
2763            self.arena.nodes[slot].symbol_count = counts;
2764            self.arena.nodes[slot].log_prob_kt = log_prob_kt;
2765            self.arena.nodes[slot].log_prob_weighted = weighted;
2766            return ChildRef::from_node(node_idx);
2767        }
2768
2769        let segment_idx = child.as_segment().unwrap();
2770        let original = self.arena.segments[segment_idx.get()];
2771        if !original.payload.is_exact() {
2772            return self.update_child_fast(
2773                logs,
2774                child,
2775                depth,
2776                history,
2777                sym_idx,
2778                singleton_log_prob_kt,
2779            );
2780        }
2781
2782        let seg_len = original.len() as usize;
2783        let mut updated_counts = original.symbol_count;
2784        let mut updated_log_prob_kt = original.log_prob_kt;
2785        apply_update_to_state_raw(logs, &mut updated_counts, &mut updated_log_prob_kt, sym_idx);
2786
2787        let depth_budget = self.max_depth.saturating_sub(depth);
2788        let comparable_len = if original.tail.is_none() {
2789            seg_len.saturating_sub(1)
2790        } else {
2791            seg_len
2792        }
2793        .min(depth_budget);
2794        let mismatch =
2795            first_exact_segment_mismatch(original.payload.exact_bits(), path_bits, comparable_len)
2796                .map(|(offset, path_edge, existing_edge)| {
2797                    (offset, depth + offset, path_edge, existing_edge)
2798                });
2799
2800        if let Some((offset, node_depth, path_edge, existing_edge)) = mismatch {
2801            let old_continuation = if offset + 1 < seg_len {
2802                if offset == 0 {
2803                    let segment = &mut self.arena.segments[segment_idx.get()];
2804                    segment.payload = original.payload.suffix_after(1);
2805                    segment.tail = original.tail;
2806                    segment.symbol_count = original.symbol_count;
2807                    segment.log_prob_kt = original.log_prob_kt;
2808                    self.recompute_segment_head(segment_idx);
2809                    ChildRef::from_segment(segment_idx)
2810                } else {
2811                    ChildRef::from_segment(self.arena.alloc_segment_with_parts(
2812                        original.symbol_count,
2813                        original.log_prob_kt,
2814                        original.tail,
2815                        original.payload.suffix_after(offset as u32 + 1),
2816                    ))
2817                }
2818            } else {
2819                original.tail
2820            };
2821
2822            let new_tail = self.build_missing_path_exact_bits(
2823                node_depth + 1,
2824                shift_path_bits(path_bits, offset + 1),
2825                sym_idx,
2826                singleton_log_prob_kt,
2827            );
2828            let branch = self
2829                .arena
2830                .alloc_node_with_state(updated_counts, updated_log_prob_kt);
2831            self.arena
2832                .set_child(branch, existing_edge as usize, old_continuation);
2833            self.arena.set_child(branch, path_edge as usize, new_tail);
2834            self.arena.recompute_node_weight(branch);
2835
2836            if offset == 0 {
2837                if offset + 1 >= seg_len {
2838                    self.arena.free_segment(segment_idx);
2839                }
2840                return ChildRef::from_node(branch);
2841            }
2842
2843            let segment = &mut self.arena.segments[segment_idx.get()];
2844            segment.payload = original.payload.prefix(offset as u32);
2845            segment.tail = ChildRef::from_node(branch);
2846            segment.symbol_count = updated_counts;
2847            segment.log_prob_kt = updated_log_prob_kt;
2848            self.recompute_segment_head(segment_idx);
2849            return ChildRef::from_segment(segment_idx);
2850        }
2851
2852        if depth_budget < seg_len {
2853            self.arena.segments[segment_idx.get()].symbol_count = updated_counts;
2854            self.arena.segments[segment_idx.get()].log_prob_kt = updated_log_prob_kt;
2855            self.recompute_segment_head(segment_idx);
2856            return ChildRef::from_segment(segment_idx);
2857        }
2858
2859        if original.tail.is_none() {
2860            let new_tail = self.build_missing_path_exact_bits(
2861                depth + seg_len,
2862                shift_path_bits(path_bits, seg_len),
2863                sym_idx,
2864                singleton_log_prob_kt,
2865            );
2866            self.arena.segments[segment_idx.get()].tail = new_tail;
2867            self.arena.segments[segment_idx.get()].symbol_count = updated_counts;
2868            self.arena.segments[segment_idx.get()].log_prob_kt = updated_log_prob_kt;
2869            self.recompute_segment_head(segment_idx);
2870            return ChildRef::from_segment(segment_idx);
2871        }
2872
2873        let tail = original.tail;
2874        let updated_tail = self.update_child_fast_exact(
2875            logs,
2876            tail,
2877            depth + seg_len,
2878            history,
2879            shift_path_bits(path_bits, seg_len),
2880            sym_idx,
2881            singleton_log_prob_kt,
2882        );
2883        if updated_tail != tail {
2884            self.arena.set_segment_tail(segment_idx, updated_tail);
2885        }
2886        self.arena.segments[segment_idx.get()].symbol_count = updated_counts;
2887        self.arena.segments[segment_idx.get()].log_prob_kt = updated_log_prob_kt;
2888        self.recompute_segment_head(segment_idx);
2889        ChildRef::from_segment(segment_idx)
2890    }
2891
2892    #[inline(always)]
2893    fn update_root_child<L: CtLogAccess>(
2894        &mut self,
2895        logs: L,
2896        child: ChildRef,
2897        history: &(impl HistoryAccess + ?Sized),
2898        sym_idx: usize,
2899        singleton_log_prob_kt: f64,
2900    ) -> ChildRef {
2901        if self.max_depth <= SEG_EXACT_MAX_LEN as usize {
2902            let path_bits = path_bits_from_history(history, 1, self.max_depth);
2903            self.update_child_fast_exact(
2904                logs,
2905                child,
2906                1,
2907                history,
2908                path_bits,
2909                sym_idx,
2910                singleton_log_prob_kt,
2911            )
2912        } else {
2913            self.update_child_fast(logs, child, 1, history, sym_idx, singleton_log_prob_kt)
2914        }
2915    }
2916
2917    fn collect_existing_levels(&mut self, history: &(impl HistoryAccess + ?Sized)) -> ChildRef {
2918        if self.max_depth == 0 {
2919            self.detaches.clear();
2920            return ChildRef::NONE;
2921        }
2922
2923        self.detaches.clear();
2924        self.levels.fill(LevelState::default());
2925
2926        let root_edge = history_symbol(history, 0) as usize;
2927        let old_child = self.arena.child(self.root, root_edge);
2928        let mut source = if let Some(node) = old_child.as_node() {
2929            ExistingSource::Node(node)
2930        } else if let Some(segment) = old_child.as_segment() {
2931            ExistingSource::Segment(segment, 0)
2932        } else {
2933            ExistingSource::None
2934        };
2935
2936        for depth in 1..=self.max_depth {
2937            let slot = depth - 1;
2938            self.levels[slot] = LevelState::default();
2939
2940            match source {
2941                ExistingSource::None => {}
2942                ExistingSource::Node(node_idx) => {
2943                    self.levels[slot].symbol_count = self.arena.counts(node_idx);
2944                    self.levels[slot].log_prob_kt = self.arena.log_prob_kt(node_idx);
2945                    if depth < self.max_depth {
2946                        let path_edge = history_symbol(history, depth) as usize;
2947                        let sibling_edge = path_edge ^ 1;
2948                        let sibling = self.arena.child(node_idx, sibling_edge);
2949                        self.levels[slot].sibling = sibling;
2950                        if sibling.is_some() {
2951                            self.detaches.push(Detach::NodeChild {
2952                                node: node_idx,
2953                                edge: sibling_edge,
2954                            });
2955                        }
2956                        let next = self.arena.child(node_idx, path_edge);
2957                        source = if let Some(next_node) = next.as_node() {
2958                            ExistingSource::Node(next_node)
2959                        } else if let Some(next_segment) = next.as_segment() {
2960                            ExistingSource::Segment(next_segment, 0)
2961                        } else {
2962                            ExistingSource::None
2963                        };
2964                    }
2965                }
2966                ExistingSource::Segment(segment_idx, offset) => {
2967                    self.levels[slot].symbol_count = self.arena.segment_symbol_count(segment_idx);
2968                    self.levels[slot].log_prob_kt = self.arena.segment_log_prob_kt(segment_idx);
2969                    if depth < self.max_depth {
2970                        let path_edge = history_symbol(history, depth) as usize;
2971                        if self.arena.segment_has_child(segment_idx, offset) {
2972                            let existing_edge =
2973                                self.arena.segment_edge(segment_idx, offset, history);
2974                            if path_edge == existing_edge {
2975                                let seg_len = self.arena.segment_len(segment_idx);
2976                                if offset + 1 < seg_len {
2977                                    source = ExistingSource::Segment(segment_idx, offset + 1);
2978                                } else {
2979                                    let tail = self.arena.segments[segment_idx.get()].tail;
2980                                    source = if let Some(next_node) = tail.as_node() {
2981                                        ExistingSource::Node(next_node)
2982                                    } else if let Some(next_segment) = tail.as_segment() {
2983                                        ExistingSource::Segment(next_segment, 0)
2984                                    } else {
2985                                        ExistingSource::None
2986                                    };
2987                                }
2988                            } else {
2989                                let continuation = self.arena.detach_segment_continuation(
2990                                    segment_idx,
2991                                    offset,
2992                                    &mut self.detaches,
2993                                );
2994                                self.levels[slot].sibling = continuation;
2995                                source = ExistingSource::None;
2996                            }
2997                        } else {
2998                            source = ExistingSource::None;
2999                        }
3000                    }
3001                }
3002            }
3003        }
3004
3005        old_child
3006    }
3007
3008    fn rebuild_path_subtree(&mut self, history: &(impl HistoryAccess + ?Sized)) -> ChildRef {
3009        let mut built = ChildRef::NONE;
3010
3011        for depth in (1..=self.max_depth).rev() {
3012            let level = self.levels[depth - 1];
3013            let visits = level.symbol_count[0] + level.symbol_count[1];
3014            if visits == 0 {
3015                built = ChildRef::NONE;
3016                continue;
3017            }
3018
3019            let path_edge = if depth < self.max_depth {
3020                history_symbol(history, depth) as usize
3021            } else {
3022                0
3023            };
3024            let has_path_child = built.is_some();
3025            let has_sibling = level.sibling.is_some();
3026            let force_node = depth <= self.hot_prefix_depth();
3027
3028            if force_node || (has_path_child && has_sibling) {
3029                let node = self
3030                    .arena
3031                    .alloc_node_with_state(level.symbol_count, level.log_prob_kt);
3032                if has_path_child {
3033                    self.arena.set_child(node, path_edge, built);
3034                }
3035                if has_sibling {
3036                    self.arena.set_child(node, path_edge ^ 1, level.sibling);
3037                }
3038                self.arena.recompute_node_weight(node);
3039                built = ChildRef::from_node(node);
3040            } else {
3041                let (edge, child) = if has_path_child {
3042                    (path_edge, built)
3043                } else if has_sibling {
3044                    (path_edge ^ 1, level.sibling)
3045                } else {
3046                    (path_edge, ChildRef::NONE)
3047                };
3048                built = self.arena.prepend_or_alloc_segment(
3049                    history,
3050                    depth,
3051                    level.symbol_count,
3052                    level.log_prob_kt,
3053                    child,
3054                    edge,
3055                    false,
3056                );
3057            }
3058        }
3059
3060        built
3061    }
3062
3063    fn apply_detaches(&mut self) {
3064        for detach in self.detaches.drain(..) {
3065            match detach {
3066                Detach::NodeChild { node, edge } => {
3067                    self.arena.set_child(node, edge, ChildRef::NONE);
3068                }
3069                Detach::SegmentNext { segment, new_len } => {
3070                    self.arena.segments[segment.get()].set_len(new_len);
3071                    self.arena.set_segment_tail(segment, ChildRef::NONE);
3072                }
3073            }
3074        }
3075    }
3076
3077    fn update_with_logs<L: CtLogAccess>(
3078        &mut self,
3079        logs: L,
3080        sym: Symbol,
3081        history: &(impl HistoryAccess + ?Sized),
3082    ) {
3083        let sym_idx = sym as usize;
3084        let singleton_log_prob_kt = logs.log_half(0) - logs.log_int(1);
3085        {
3086            let slot = self.root.get();
3087            let mut counts = self.arena.nodes[slot].symbol_count;
3088            let mut log_prob_kt = self.arena.nodes[slot].log_prob_kt;
3089            apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
3090            self.arena.nodes[slot].symbol_count = counts;
3091            self.arena.nodes[slot].log_prob_kt = log_prob_kt;
3092        }
3093
3094        if self.max_depth > 0 {
3095            let root_edge = history_symbol(history, 0) as usize;
3096            let old_child = self.arena.child(self.root, root_edge);
3097            let new_child =
3098                self.update_root_child(logs, old_child, history, sym_idx, singleton_log_prob_kt);
3099            self.arena.set_child(self.root, root_edge, new_child);
3100        }
3101
3102        self.arena.recompute_node_weight(self.root);
3103    }
3104
3105    fn update(&mut self, sym: Symbol, history: &(impl HistoryAccess + ?Sized)) {
3106        let upto = self.root_visits() + 1;
3107        if upto <= ctw_log_cache_limit() {
3108            self.with_cached_logs(upto, |this, logs| {
3109                this.update_with_logs(logs, sym, history);
3110            });
3111        } else {
3112            self.with_bounded_logs(upto, |this, logs| {
3113                this.update_with_logs(logs, sym, history);
3114            });
3115        }
3116    }
3117
3118    fn update_prepared_with_logs<L: CtLogAccess>(
3119        &mut self,
3120        logs: L,
3121        history: &(impl HistoryAccess + ?Sized),
3122        sym_idx: usize,
3123        use_prepared: bool,
3124    ) {
3125        let singleton_log_prob_kt = logs.log_half(0) - logs.log_int(1);
3126        {
3127            let slot = self.root.get();
3128            let mut counts = self.arena.nodes[slot].symbol_count;
3129            let mut log_prob_kt = self.arena.nodes[slot].log_prob_kt;
3130            apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
3131            self.arena.nodes[slot].symbol_count = counts;
3132            self.arena.nodes[slot].log_prob_kt = log_prob_kt;
3133        }
3134
3135        if self.max_depth > 0 {
3136            let root_edge = history_symbol(history, 0) as usize;
3137            let old_child = self.arena.child(self.root, root_edge);
3138            let new_child = if use_prepared {
3139                match self.prepared_end {
3140                    PreparedEnd::MissingAtRoot => {
3141                        self.build_missing_path(1, history, sym_idx, singleton_log_prob_kt)
3142                    }
3143                    PreparedEnd::MaxDepth | PreparedEnd::MissingAfterCurrent => {
3144                        if !self.prepared_steps.is_empty() {
3145                            self.update_prepared_cached_path(
3146                                logs,
3147                                history,
3148                                sym_idx,
3149                                singleton_log_prob_kt,
3150                            );
3151                        }
3152                        old_child
3153                    }
3154                    PreparedEnd::MismatchAtCurrentSegment => {
3155                        self.update_prepared_mismatch(logs, history, sym_idx, singleton_log_prob_kt)
3156                    }
3157                }
3158            } else {
3159                self.update_root_child(logs, old_child, history, sym_idx, singleton_log_prob_kt)
3160            };
3161            self.arena.set_child(self.root, root_edge, new_child);
3162        }
3163
3164        self.arena.recompute_node_weight(self.root);
3165    }
3166
3167    fn update_prepared(
3168        &mut self,
3169        sym: Symbol,
3170        history: &(impl HistoryAccess + ?Sized),
3171        use_prepared: bool,
3172    ) {
3173        let upto = self.root_visits() + 1;
3174        let sym_idx = sym as usize;
3175        if upto <= ctw_log_cache_limit() {
3176            self.with_cached_logs(upto, |this, logs| {
3177                this.update_prepared_with_logs(logs, history, sym_idx, use_prepared);
3178            });
3179        } else {
3180            self.with_bounded_logs(upto, |this, logs| {
3181                this.update_prepared_with_logs(logs, history, sym_idx, use_prepared);
3182            });
3183        }
3184    }
3185
3186    fn revert_with_logs<L: CtLogAccess>(
3187        &mut self,
3188        logs: L,
3189        history: &(impl HistoryAccess + ?Sized),
3190        sym_idx: usize,
3191    ) {
3192        let old_child = self.collect_existing_levels(history);
3193
3194        {
3195            let slot = self.root.get();
3196            let mut counts = self.arena.nodes[slot].symbol_count;
3197            let mut log_prob_kt = self.arena.nodes[slot].log_prob_kt;
3198            apply_revert_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
3199            self.arena.nodes[slot].symbol_count = counts;
3200            self.arena.nodes[slot].log_prob_kt = log_prob_kt;
3201        }
3202
3203        for level in &mut self.levels {
3204            let mut counts = level.symbol_count;
3205            let mut log_prob_kt = level.log_prob_kt;
3206            apply_revert_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
3207            level.symbol_count = counts;
3208            level.log_prob_kt = log_prob_kt;
3209        }
3210
3211        if self.max_depth > 0 {
3212            let new_child = self.rebuild_path_subtree(history);
3213            let root_edge = history_symbol(history, 0) as usize;
3214            self.apply_detaches();
3215            self.arena.free_child_ref(old_child);
3216            self.arena.set_child(self.root, root_edge, new_child);
3217        }
3218
3219        self.arena.recompute_node_weight(self.root);
3220    }
3221
3222    fn revert(&mut self, sym: Symbol, history: &(impl HistoryAccess + ?Sized)) {
3223        let upto = self.root_visits();
3224        let sym_idx = sym as usize;
3225        if upto <= ctw_log_cache_limit() {
3226            self.with_cached_logs(upto, |this, logs| {
3227                this.revert_with_logs(logs, history, sym_idx);
3228            });
3229        } else {
3230            self.with_bounded_logs(upto, |this, logs| {
3231                this.revert_with_logs(logs, history, sym_idx);
3232            });
3233        }
3234    }
3235
3236    fn predict(&mut self, sym: Symbol, history: &(impl HistoryAccess + ?Sized)) -> f64 {
3237        self.prepared_steps.clear();
3238        self.prepared_levels = 0;
3239        self.prepared_end = PreparedEnd::MaxDepth;
3240
3241        let (root_sibling, root_has_sibling, mut source) = if self.max_depth > 0 {
3242            let root_edge = history_symbol(history, 0) as usize;
3243            let path_child = self.arena.child(self.root, root_edge);
3244            let sibling = self.arena.child(self.root, root_edge ^ 1);
3245            (
3246                self.arena.child_ref_weighted(sibling),
3247                sibling.is_some() as u8,
3248                Self::child_to_existing_source(path_child).unwrap_or(ExistingSource::None),
3249            )
3250        } else {
3251            (0.0, 0, ExistingSource::None)
3252        };
3253        if self.max_depth > 0 && matches!(source, ExistingSource::None) {
3254            self.prepared_end = PreparedEnd::MissingAtRoot;
3255        }
3256
3257        let history_len = history.len();
3258        let mut depth = 1usize;
3259        'walk: while depth <= self.max_depth {
3260            match source {
3261                ExistingSource::None => break,
3262                ExistingSource::Node(node_idx) => {
3263                    if depth == self.max_depth {
3264                        self.prepared_steps.push(PreparedStep {
3265                            source: ExistingSource::Node(node_idx),
3266                            span: 1,
3267                            sibling_weight: 0.0,
3268                            has_sibling: 0,
3269                        });
3270                        self.prepared_levels += 1;
3271                        break;
3272                    }
3273                    let path_edge = history_symbol(history, depth) as usize;
3274                    let sibling = self.arena.child(node_idx, path_edge ^ 1);
3275                    self.prepared_steps.push(PreparedStep {
3276                        source: ExistingSource::Node(node_idx),
3277                        span: 1,
3278                        sibling_weight: self.arena.child_ref_weighted(sibling),
3279                        has_sibling: sibling.is_some() as u8,
3280                    });
3281                    self.prepared_levels += 1;
3282                    let next = self.arena.child(node_idx, path_edge);
3283                    source = Self::child_to_existing_source(next).unwrap_or(ExistingSource::None);
3284                    if matches!(source, ExistingSource::None) {
3285                        self.prepared_end = PreparedEnd::MissingAfterCurrent;
3286                        break;
3287                    }
3288                    depth += 1;
3289                }
3290                ExistingSource::Segment(segment_idx, _) => {
3291                    let segment = self.arena.segments[segment_idx.get()];
3292                    if segment.payload.is_exact() {
3293                        let path_bits = path_bits_from_history(
3294                            history,
3295                            depth,
3296                            self.max_depth.saturating_sub(depth),
3297                        );
3298                        if let Some((next_depth, next_source)) =
3299                            self.walk_prepared_exact_segment(segment_idx, segment, depth, path_bits)
3300                        {
3301                            source = next_source;
3302                            if matches!(source, ExistingSource::None) {
3303                                self.prepared_end = PreparedEnd::MissingAfterCurrent;
3304                                break 'walk;
3305                            }
3306                            depth = next_depth;
3307                            continue 'walk;
3308                        }
3309                        break 'walk;
3310                    }
3311                    let seg_len = segment.len() as usize;
3312                    for offset in 0..seg_len {
3313                        let node_depth = depth + offset;
3314                        let span = (offset + 1) as u32;
3315
3316                        if node_depth == self.max_depth {
3317                            self.prepared_steps.push(PreparedStep {
3318                                source: ExistingSource::Segment(segment_idx, offset as u32),
3319                                span,
3320                                sibling_weight: 0.0,
3321                                has_sibling: 0,
3322                            });
3323                            self.prepared_levels += span as usize;
3324                            break 'walk;
3325                        }
3326
3327                        if offset + 1 >= seg_len && segment.tail.is_none() {
3328                            self.prepared_steps.push(PreparedStep {
3329                                source: ExistingSource::Segment(segment_idx, offset as u32),
3330                                span,
3331                                sibling_weight: 0.0,
3332                                has_sibling: 0,
3333                            });
3334                            self.prepared_levels += span as usize;
3335                            self.prepared_end = PreparedEnd::MissingAfterCurrent;
3336                            break 'walk;
3337                        }
3338
3339                        let path_edge = path_edge_at_depth(history, history_len, node_depth);
3340                        let existing_edge =
3341                            segment_edge_from_parts(segment, offset, history, history_len);
3342                        if path_edge != existing_edge {
3343                            self.prepared_steps.push(PreparedStep {
3344                                source: ExistingSource::Segment(segment_idx, offset as u32),
3345                                span,
3346                                sibling_weight: self
3347                                    .arena
3348                                    .segment_continuation_weight(segment_idx, offset as u32),
3349                                has_sibling: 1,
3350                            });
3351                            self.prepared_levels += span as usize;
3352                            self.prepared_end = PreparedEnd::MismatchAtCurrentSegment;
3353                            break 'walk;
3354                        }
3355
3356                        if offset + 1 < seg_len {
3357                            continue;
3358                        }
3359
3360                        self.prepared_steps.push(PreparedStep {
3361                            source: ExistingSource::Segment(segment_idx, offset as u32),
3362                            span,
3363                            sibling_weight: 0.0,
3364                            has_sibling: 0,
3365                        });
3366                        self.prepared_levels += span as usize;
3367                        let tail = segment.tail;
3368                        source =
3369                            Self::child_to_existing_source(tail).unwrap_or(ExistingSource::None);
3370                        if matches!(source, ExistingSource::None) {
3371                            self.prepared_end = PreparedEnd::MissingAfterCurrent;
3372                            break 'walk;
3373                        }
3374                        depth = node_depth + 1;
3375                        continue 'walk;
3376                    }
3377                }
3378            }
3379        }
3380
3381        let sym_idx = sym as usize;
3382        if self.prepared_levels == 0 {
3383            let counts = self.arena.counts(self.root);
3384            let kt_log_prob = self.arena.log_prob_kt(self.root);
3385            return if self.prepared_end == PreparedEnd::MaxDepth || root_has_sibling == 0 {
3386                predict_ratio_kt(counts, sym_idx)
3387            } else {
3388                predict_ratio_internal(kt_log_prob, counts, 0.0, root_sibling, 0.5, sym_idx)
3389            };
3390        }
3391
3392        let last_step = *self.prepared_steps.last().unwrap();
3393        let (last_counts, last_kt_log_prob) = self.source_counts_and_kt_log_prob(last_step.source);
3394        let (mut child_weight, mut ratio) = if (self.prepared_end == PreparedEnd::MaxDepth
3395            && self.prepared_levels == self.max_depth)
3396            || last_step.has_sibling == 0
3397        {
3398            (last_kt_log_prob, predict_ratio_kt(last_counts, sym_idx))
3399        } else {
3400            combined_weight_ratio_internal(
3401                last_kt_log_prob,
3402                last_counts,
3403                0.0,
3404                last_step.sibling_weight,
3405                0.5,
3406                sym_idx,
3407            )
3408        };
3409
3410        if matches!(last_step.source, ExistingSource::Segment(_, _)) && last_step.span > 1 {
3411            let (alpha, log_alpha, log_one_minus_alpha) =
3412                self.segment_constants(last_step.span - 1);
3413            (child_weight, ratio) = unary_chain_ratio_transform_precomputed(
3414                last_kt_log_prob,
3415                last_counts,
3416                child_weight,
3417                ratio,
3418                alpha,
3419                log_alpha,
3420                log_one_minus_alpha,
3421                sym_idx,
3422            );
3423        }
3424
3425        for idx in (0..self.prepared_steps.len() - 1).rev() {
3426            let step = self.prepared_steps[idx];
3427            let (step_counts, step_kt_log_prob) = self.source_counts_and_kt_log_prob(step.source);
3428            match step.source {
3429                ExistingSource::Node(_) => {
3430                    (child_weight, ratio) = combined_weight_ratio_internal(
3431                        step_kt_log_prob,
3432                        step_counts,
3433                        child_weight,
3434                        step.sibling_weight,
3435                        ratio,
3436                        sym_idx,
3437                    );
3438                }
3439                ExistingSource::Segment(_, _) => {
3440                    let (alpha, log_alpha, log_one_minus_alpha) = self.segment_constants(step.span);
3441                    (child_weight, ratio) = unary_chain_ratio_transform_precomputed(
3442                        step_kt_log_prob,
3443                        step_counts,
3444                        child_weight,
3445                        ratio,
3446                        alpha,
3447                        log_alpha,
3448                        log_one_minus_alpha,
3449                        sym_idx,
3450                    );
3451                }
3452                ExistingSource::None => unreachable!("prepared step should never store None"),
3453            }
3454        }
3455
3456        let root_counts = self.arena.counts(self.root);
3457        let root_kt_log_prob = self.arena.log_prob_kt(self.root);
3458        predict_ratio_internal(
3459            root_kt_log_prob,
3460            root_counts,
3461            child_weight,
3462            root_sibling,
3463            ratio,
3464            sym_idx,
3465        )
3466    }
3467
3468    fn predict_one(&mut self, history: &(impl HistoryAccess + ?Sized)) -> f64 {
3469        self.prepared_steps.clear();
3470        self.prepared_levels = 0;
3471        self.prepared_end = PreparedEnd::MaxDepth;
3472
3473        let (root_sibling, root_has_sibling, mut source) = if self.max_depth > 0 {
3474            let root_edge = history_symbol(history, 0) as usize;
3475            let path_child = self.arena.child(self.root, root_edge);
3476            let sibling = self.arena.child(self.root, root_edge ^ 1);
3477            (
3478                self.arena.child_ref_weighted(sibling),
3479                sibling.is_some() as u8,
3480                Self::child_to_existing_source(path_child).unwrap_or(ExistingSource::None),
3481            )
3482        } else {
3483            (0.0, 0, ExistingSource::None)
3484        };
3485        if self.max_depth > 0 && matches!(source, ExistingSource::None) {
3486            self.prepared_end = PreparedEnd::MissingAtRoot;
3487        }
3488
3489        let history_len = history.len();
3490        let mut depth = 1usize;
3491        'walk: while depth <= self.max_depth {
3492            match source {
3493                ExistingSource::None => break,
3494                ExistingSource::Node(node_idx) => {
3495                    if depth == self.max_depth {
3496                        self.prepared_steps.push(PreparedStep {
3497                            source: ExistingSource::Node(node_idx),
3498                            span: 1,
3499                            sibling_weight: 0.0,
3500                            has_sibling: 0,
3501                        });
3502                        self.prepared_levels += 1;
3503                        break;
3504                    }
3505                    let path_edge = history_symbol(history, depth) as usize;
3506                    let sibling = self.arena.child(node_idx, path_edge ^ 1);
3507                    self.prepared_steps.push(PreparedStep {
3508                        source: ExistingSource::Node(node_idx),
3509                        span: 1,
3510                        sibling_weight: self.arena.child_ref_weighted(sibling),
3511                        has_sibling: sibling.is_some() as u8,
3512                    });
3513                    self.prepared_levels += 1;
3514                    let next = self.arena.child(node_idx, path_edge);
3515                    source = Self::child_to_existing_source(next).unwrap_or(ExistingSource::None);
3516                    if matches!(source, ExistingSource::None) {
3517                        self.prepared_end = PreparedEnd::MissingAfterCurrent;
3518                        break;
3519                    }
3520                    depth += 1;
3521                }
3522                ExistingSource::Segment(segment_idx, _) => {
3523                    let segment = self.arena.segments[segment_idx.get()];
3524                    if segment.payload.is_exact() {
3525                        let path_bits = path_bits_from_history(
3526                            history,
3527                            depth,
3528                            self.max_depth.saturating_sub(depth),
3529                        );
3530                        if let Some((next_depth, next_source)) =
3531                            self.walk_prepared_exact_segment(segment_idx, segment, depth, path_bits)
3532                        {
3533                            source = next_source;
3534                            if matches!(source, ExistingSource::None) {
3535                                self.prepared_end = PreparedEnd::MissingAfterCurrent;
3536                                break 'walk;
3537                            }
3538                            depth = next_depth;
3539                            continue 'walk;
3540                        }
3541                        break 'walk;
3542                    }
3543                    let seg_len = segment.len() as usize;
3544                    for offset in 0..seg_len {
3545                        let node_depth = depth + offset;
3546                        let span = (offset + 1) as u32;
3547
3548                        if node_depth == self.max_depth {
3549                            self.prepared_steps.push(PreparedStep {
3550                                source: ExistingSource::Segment(segment_idx, offset as u32),
3551                                span,
3552                                sibling_weight: 0.0,
3553                                has_sibling: 0,
3554                            });
3555                            self.prepared_levels += span as usize;
3556                            break 'walk;
3557                        }
3558
3559                        if offset + 1 >= seg_len && segment.tail.is_none() {
3560                            self.prepared_steps.push(PreparedStep {
3561                                source: ExistingSource::Segment(segment_idx, offset as u32),
3562                                span,
3563                                sibling_weight: 0.0,
3564                                has_sibling: 0,
3565                            });
3566                            self.prepared_levels += span as usize;
3567                            self.prepared_end = PreparedEnd::MissingAfterCurrent;
3568                            break 'walk;
3569                        }
3570
3571                        let path_edge = path_edge_at_depth(history, history_len, node_depth);
3572                        let existing_edge =
3573                            segment_edge_from_parts(segment, offset, history, history_len);
3574                        if path_edge != existing_edge {
3575                            self.prepared_steps.push(PreparedStep {
3576                                source: ExistingSource::Segment(segment_idx, offset as u32),
3577                                span,
3578                                sibling_weight: self
3579                                    .arena
3580                                    .segment_continuation_weight(segment_idx, offset as u32),
3581                                has_sibling: 1,
3582                            });
3583                            self.prepared_levels += span as usize;
3584                            self.prepared_end = PreparedEnd::MismatchAtCurrentSegment;
3585                            break 'walk;
3586                        }
3587
3588                        if offset + 1 < seg_len {
3589                            continue;
3590                        }
3591
3592                        self.prepared_steps.push(PreparedStep {
3593                            source: ExistingSource::Segment(segment_idx, offset as u32),
3594                            span,
3595                            sibling_weight: 0.0,
3596                            has_sibling: 0,
3597                        });
3598                        self.prepared_levels += span as usize;
3599                        let tail = segment.tail;
3600                        source =
3601                            Self::child_to_existing_source(tail).unwrap_or(ExistingSource::None);
3602                        if matches!(source, ExistingSource::None) {
3603                            self.prepared_end = PreparedEnd::MissingAfterCurrent;
3604                            break 'walk;
3605                        }
3606                        depth = node_depth + 1;
3607                        continue 'walk;
3608                    }
3609                }
3610            }
3611        }
3612
3613        if self.prepared_levels == 0 {
3614            let counts = self.arena.counts(self.root);
3615            let kt_log_prob = self.arena.log_prob_kt(self.root);
3616            return if self.prepared_end == PreparedEnd::MaxDepth || root_has_sibling == 0 {
3617                predict_ratio_kt_one(counts)
3618            } else {
3619                predict_ratio_internal_one(kt_log_prob, counts, 0.0, root_sibling, 0.5)
3620            };
3621        }
3622
3623        let last_step = *self.prepared_steps.last().unwrap();
3624        let (last_counts, last_kt_log_prob) = self.source_counts_and_kt_log_prob(last_step.source);
3625        let (mut child_weight, mut ratio) = if (self.prepared_end == PreparedEnd::MaxDepth
3626            && self.prepared_levels == self.max_depth)
3627            || last_step.has_sibling == 0
3628        {
3629            (last_kt_log_prob, predict_ratio_kt_one(last_counts))
3630        } else {
3631            combined_weight_ratio_internal_one(
3632                last_kt_log_prob,
3633                last_counts,
3634                0.0,
3635                last_step.sibling_weight,
3636                0.5,
3637            )
3638        };
3639
3640        if let ExistingSource::Segment(_, _) = last_step.source
3641            && last_step.span > 1
3642        {
3643            let (alpha, log_alpha, log_one_minus_alpha) =
3644                self.segment_constants(last_step.span - 1);
3645            (child_weight, ratio) = unary_chain_ratio_transform_precomputed_one(
3646                last_kt_log_prob,
3647                last_counts,
3648                child_weight,
3649                ratio,
3650                alpha,
3651                log_alpha,
3652                log_one_minus_alpha,
3653            );
3654        }
3655
3656        for idx in (0..self.prepared_steps.len() - 1).rev() {
3657            let step = self.prepared_steps[idx];
3658            let (step_counts, step_kt_log_prob) = self.source_counts_and_kt_log_prob(step.source);
3659            match step.source {
3660                ExistingSource::Node(_) => {
3661                    (child_weight, ratio) = combined_weight_ratio_internal_one(
3662                        step_kt_log_prob,
3663                        step_counts,
3664                        child_weight,
3665                        step.sibling_weight,
3666                        ratio,
3667                    );
3668                }
3669                ExistingSource::Segment(_, _) => {
3670                    let (alpha, log_alpha, log_one_minus_alpha) = self.segment_constants(step.span);
3671                    (child_weight, ratio) = unary_chain_ratio_transform_precomputed_one(
3672                        step_kt_log_prob,
3673                        step_counts,
3674                        child_weight,
3675                        ratio,
3676                        alpha,
3677                        log_alpha,
3678                        log_one_minus_alpha,
3679                    );
3680                }
3681                ExistingSource::None => unreachable!("prepared step should never store None"),
3682            }
3683        }
3684
3685        let root_counts = self.arena.counts(self.root);
3686        let root_kt_log_prob = self.arena.log_prob_kt(self.root);
3687        predict_ratio_internal_one(
3688            root_kt_log_prob,
3689            root_counts,
3690            child_weight,
3691            root_sibling,
3692            ratio,
3693        )
3694    }
3695
3696    fn memory_usage(&self) -> usize {
3697        self.arena.memory_usage()
3698            + self.segment_alpha.capacity() * size_of::<f64>()
3699            + self.segment_log_alpha.capacity() * size_of::<f64>()
3700            + self.segment_log_one_minus_alpha.capacity() * size_of::<f64>()
3701            + self.levels.capacity() * size_of::<LevelState>()
3702            + self.detaches.capacity() * size_of::<Detach>()
3703            + self.prepared_steps.capacity() * size_of::<PreparedStep>()
3704    }
3705
3706    #[cfg(any(test, feature = "research-tooling"))]
3707    fn scratch_memory_usage(&self) -> usize {
3708        self.segment_alpha.capacity() * size_of::<f64>()
3709            + self.segment_log_alpha.capacity() * size_of::<f64>()
3710            + self.segment_log_one_minus_alpha.capacity() * size_of::<f64>()
3711            + self.levels.capacity() * size_of::<LevelState>()
3712            + self.detaches.capacity() * size_of::<Detach>()
3713            + self.prepared_steps.capacity() * size_of::<PreparedStep>()
3714    }
3715
3716    #[cfg(any(test, feature = "research-tooling"))]
3717    fn telemetry(&self, bit_index: usize) -> FacContextTreeTreeTelemetry {
3718        let mut exact_segments: usize = 0;
3719        let mut history_segments: usize = 0;
3720        let mut history_invert_segments: usize = 0;
3721        let mut const_segments: usize = 0;
3722        let mut segment_bits: u64 = 0;
3723        let mut max_segment_len: u32 = 0;
3724        for segment in &self.arena.segments {
3725            let len = segment.len();
3726            segment_bits = segment_bits.saturating_add(len as u64);
3727            max_segment_len = max_segment_len.max(len);
3728            match segment.payload.mode() {
3729                SEG_MODE_EXACT => exact_segments = exact_segments.saturating_add(1),
3730                SEG_MODE_HISTORY => history_segments = history_segments.saturating_add(1),
3731                SEG_MODE_HISTORY_INVERT => {
3732                    history_invert_segments = history_invert_segments.saturating_add(1);
3733                }
3734                SEG_MODE_CONST => const_segments = const_segments.saturating_add(1),
3735                _ => unreachable!("invalid ctw segment payload mode"),
3736            }
3737        }
3738
3739        let node_payload_bytes = self.arena.nodes.len() * size_of::<CtNode>();
3740        let node_bytes = self.arena.nodes.capacity() * size_of::<CtNode>();
3741        let segment_payload_bytes = self.arena.segments.len() * size_of::<CtSegment>();
3742        let segment_bytes = self.arena.segments.capacity() * size_of::<CtSegment>();
3743        let free_list_bytes = self.arena.free_nodes.capacity() * size_of::<NodeIndex>()
3744            + self.arena.free_segments.capacity() * size_of::<SegmentIndex>();
3745        let scratch_bytes = self.scratch_memory_usage();
3746        let arena_slack_bytes = node_bytes
3747            .saturating_sub(node_payload_bytes)
3748            .saturating_add(segment_bytes.saturating_sub(segment_payload_bytes));
3749
3750        FacContextTreeTreeTelemetry {
3751            bit_index,
3752            max_depth: self.max_depth,
3753            root_visits: self.root_visits(),
3754            nodes_len: self.arena.nodes.len(),
3755            nodes_capacity: self.arena.nodes.capacity(),
3756            segments_len: self.arena.segments.len(),
3757            segments_capacity: self.arena.segments.capacity(),
3758            free_nodes_len: self.arena.free_nodes.len(),
3759            free_nodes_capacity: self.arena.free_nodes.capacity(),
3760            free_segments_len: self.arena.free_segments.len(),
3761            free_segments_capacity: self.arena.free_segments.capacity(),
3762            node_bytes,
3763            node_payload_bytes,
3764            segment_bytes,
3765            segment_payload_bytes,
3766            free_list_bytes,
3767            scratch_bytes,
3768            total_bytes: node_bytes
3769                .saturating_add(segment_bytes)
3770                .saturating_add(free_list_bytes)
3771                .saturating_add(scratch_bytes),
3772            arena_slack_bytes,
3773            exact_segments,
3774            history_segments,
3775            history_invert_segments,
3776            const_segments,
3777            segment_bits,
3778            max_segment_len,
3779        }
3780    }
3781}
3782
3783/// A Context Tree for binary sequence prediction.
3784#[derive(Clone)]
3785pub struct ContextTree {
3786    engine: CtEngine,
3787    history: BitHistory,
3788    history_version: u64,
3789    prepared_valid: bool,
3790    prepared_history_len: usize,
3791    prepared_history_version: u64,
3792}
3793
3794#[derive(Clone)]
3795pub(crate) struct ContextTreeLifecycleSnapshot {
3796    history: BitHistory,
3797    history_version: u64,
3798}
3799
3800impl ContextTree {
3801    /// Construct a binary CTW predictor with maximum context depth `depth`.
3802    pub fn new(depth: usize) -> Self {
3803        Self {
3804            engine: CtEngine::new(depth),
3805            history: BitHistory::default(),
3806            history_version: 0,
3807            prepared_valid: false,
3808            prepared_history_len: 0,
3809            prepared_history_version: 0,
3810        }
3811    }
3812
3813    #[inline]
3814    fn bump_history_version(&mut self) {
3815        self.history_version = self.history_version.wrapping_add(1);
3816    }
3817
3818    #[inline]
3819    fn clear_prepared_prediction(&mut self) {
3820        self.prepared_valid = false;
3821    }
3822
3823    #[inline]
3824    fn prepared_prediction_matches_history(&self) -> bool {
3825        self.prepared_valid
3826            && self.prepared_history_len == self.history.len()
3827            && self.prepared_history_version == self.history_version
3828    }
3829
3830    /// Reset tree parameters and clear conditioning history.
3831    pub fn clear(&mut self) {
3832        self.history.clear();
3833        self.engine.clear();
3834        self.history_version = 0;
3835        self.clear_prepared_prediction();
3836        self.prepared_history_len = 0;
3837        self.prepared_history_version = 0;
3838    }
3839
3840    #[inline]
3841    pub(crate) fn reserve_for_symbols(&mut self, total_symbols: usize) {
3842        if total_symbols == 0 {
3843            return;
3844        }
3845        self.engine.reserve_for_symbols(total_symbols);
3846        self.history.reserve_exact(total_symbols);
3847    }
3848
3849    #[inline]
3850    /// Observe one binary symbol and update the model.
3851    pub fn update(&mut self, sym: Symbol) {
3852        let use_prepared = self.prepared_prediction_matches_history();
3853        self.clear_prepared_prediction();
3854        if use_prepared {
3855            self.engine.update_prepared(sym, &self.history, true);
3856        } else {
3857            self.engine.update(sym, &self.history);
3858        }
3859        self.history.push(sym);
3860        self.bump_history_version();
3861    }
3862
3863    #[inline]
3864    /// Revert the last symbol update if history is non-empty.
3865    pub fn revert(&mut self) {
3866        let Some(last_sym) = self.history.pop() else {
3867            return;
3868        };
3869        self.clear_prepared_prediction();
3870        self.engine.revert(last_sym, &self.history);
3871        self.bump_history_version();
3872    }
3873
3874    #[inline]
3875    /// Append external symbols to history without touching model state.
3876    pub fn update_history(&mut self, symbols: &[Symbol]) {
3877        if symbols.is_empty() {
3878            return;
3879        }
3880        self.clear_prepared_prediction();
3881        self.history.extend_from_slice(symbols);
3882        self.bump_history_version();
3883    }
3884
3885    #[inline]
3886    /// Remove one history symbol without reverting model statistics.
3887    pub fn revert_history(&mut self) {
3888        if self.history.pop().is_some() {
3889            self.clear_prepared_prediction();
3890            self.bump_history_version();
3891        }
3892    }
3893
3894    /// Truncate the stored history to `new_size` symbols.
3895    pub fn truncate_history(&mut self, new_size: usize) {
3896        if new_size < self.history.len() {
3897            self.clear_prepared_prediction();
3898            self.history.truncate(new_size);
3899            self.bump_history_version();
3900        }
3901    }
3902
3903    /// Capture rollback state for stream lifecycle transactions.
3904    ///
3905    /// This clones the full conditioning history, so the allocation and copy are
3906    /// O(history length). It is intended for stream lifecycle boundaries rather
3907    /// than per-symbol speculative prediction.
3908    pub(crate) fn lifecycle_snapshot(&self) -> ContextTreeLifecycleSnapshot {
3909        ContextTreeLifecycleSnapshot {
3910            history: self.history.clone(),
3911            history_version: self.history_version,
3912        }
3913    }
3914
3915    pub(crate) fn restore_lifecycle_snapshot(&mut self, snapshot: ContextTreeLifecycleSnapshot) {
3916        self.history = snapshot.history;
3917        self.history_version = snapshot.history_version;
3918        self.clear_prepared_prediction();
3919    }
3920
3921    #[inline]
3922    /// Predict `P(sym | history)` under current weighted CTW model.
3923    pub fn predict(&mut self, sym: Symbol) -> f64 {
3924        let prob = self.engine.predict(sym, &self.history);
3925        self.prepared_valid = true;
3926        self.prepared_history_len = self.history.len();
3927        self.prepared_history_version = self.history_version;
3928        prob
3929    }
3930
3931    #[inline]
3932    pub(crate) fn predict_one(&mut self) -> f64 {
3933        let prob = self.engine.predict_one(&self.history);
3934        self.prepared_valid = true;
3935        self.prepared_history_len = self.history.len();
3936        self.prepared_history_version = self.history_version;
3937        prob
3938    }
3939
3940    #[inline]
3941    /// Predict probability of symbol `true`.
3942    pub fn predict_sym_prob(&mut self) -> f64 {
3943        self.predict_one()
3944    }
3945
3946    #[inline]
3947    /// Return log block probability accumulated by the root model.
3948    pub fn get_log_block_probability(&self) -> f64 {
3949        self.engine.get_log_block_probability()
3950    }
3951
3952    #[inline]
3953    /// Maximum context depth configured for this tree.
3954    pub fn depth(&self) -> usize {
3955        self.engine.max_depth
3956    }
3957
3958    #[inline]
3959    /// Current number of stored history symbols.
3960    pub fn history_size(&self) -> usize {
3961        self.history.len()
3962    }
3963}
3964
3965#[derive(Clone)]
3966struct ContextTreeCore {
3967    engine: CtEngine,
3968    prepared_valid: bool,
3969    prepared_history_len: usize,
3970    prepared_history_version: u64,
3971}
3972
3973#[derive(Clone, Copy)]
3974struct ContextTreeCorePreparedSnapshot {
3975    prepared_valid: bool,
3976    prepared_history_len: usize,
3977    prepared_history_version: u64,
3978}
3979
3980impl ContextTreeCore {
3981    fn new(depth: usize) -> Self {
3982        Self {
3983            engine: CtEngine::new(depth),
3984            prepared_valid: false,
3985            prepared_history_len: 0,
3986            prepared_history_version: 0,
3987        }
3988    }
3989
3990    fn clear(&mut self) {
3991        self.engine.clear();
3992        self.prepared_valid = false;
3993        self.prepared_history_len = 0;
3994        self.prepared_history_version = 0;
3995    }
3996
3997    #[inline]
3998    fn reserve_for_symbols(&mut self, total_symbols: usize) {
3999        self.engine.reserve_for_symbols(total_symbols);
4000    }
4001
4002    #[inline]
4003    fn update_predicted(
4004        &mut self,
4005        sym: Symbol,
4006        shared_history: &(impl HistoryAccess + ?Sized),
4007        history_version: u64,
4008    ) {
4009        let use_prepared = self.prepared_valid
4010            && self.prepared_history_len == shared_history.len()
4011            && self.prepared_history_version == history_version;
4012        self.prepared_valid = false;
4013        self.engine
4014            .update_prepared(sym, shared_history, use_prepared);
4015    }
4016
4017    #[inline]
4018    fn update_predicted_with_logs<L: CtLogAccess>(
4019        &mut self,
4020        logs: L,
4021        sym: Symbol,
4022        shared_history: &(impl HistoryAccess + ?Sized),
4023        history_version: u64,
4024    ) {
4025        let use_prepared = self.prepared_valid
4026            && self.prepared_history_len == shared_history.len()
4027            && self.prepared_history_version == history_version;
4028        self.prepared_valid = false;
4029        self.engine
4030            .update_prepared_with_logs(logs, shared_history, sym as usize, use_prepared);
4031    }
4032
4033    #[inline]
4034    fn revert(&mut self, last_sym: Symbol, shared_history: &(impl HistoryAccess + ?Sized)) {
4035        self.prepared_valid = false;
4036        self.engine.revert(last_sym, shared_history);
4037    }
4038
4039    #[inline]
4040    fn predict(
4041        &mut self,
4042        sym: Symbol,
4043        shared_history: &(impl HistoryAccess + ?Sized),
4044        history_version: u64,
4045    ) -> f64 {
4046        let prob = self.engine.predict(sym, shared_history);
4047        self.prepared_valid = true;
4048        self.prepared_history_len = shared_history.len();
4049        self.prepared_history_version = history_version;
4050        prob
4051    }
4052
4053    #[inline]
4054    fn predict_one(
4055        &mut self,
4056        shared_history: &(impl HistoryAccess + ?Sized),
4057        history_version: u64,
4058    ) -> f64 {
4059        let prob = self.engine.predict_one(shared_history);
4060        self.prepared_valid = true;
4061        self.prepared_history_len = shared_history.len();
4062        self.prepared_history_version = history_version;
4063        prob
4064    }
4065
4066    #[inline]
4067    fn get_log_block_probability(&self) -> f64 {
4068        self.engine.get_log_block_probability()
4069    }
4070
4071    #[inline]
4072    fn prepared_snapshot(&self) -> ContextTreeCorePreparedSnapshot {
4073        ContextTreeCorePreparedSnapshot {
4074            prepared_valid: self.prepared_valid,
4075            prepared_history_len: self.prepared_history_len,
4076            prepared_history_version: self.prepared_history_version,
4077        }
4078    }
4079
4080    #[inline]
4081    fn restore_prepared_snapshot(&mut self, snapshot: ContextTreeCorePreparedSnapshot) {
4082        self.prepared_valid = snapshot.prepared_valid;
4083        self.prepared_history_len = snapshot.prepared_history_len;
4084        self.prepared_history_version = snapshot.prepared_history_version;
4085    }
4086}
4087
4088/// Factorized Action-Conditional Context Tree Weighting.
4089#[derive(Clone)]
4090pub struct FacContextTree {
4091    trees: Vec<ContextTreeCore>,
4092    shared_history: BitHistory,
4093    base_depth: usize,
4094    num_bits: usize,
4095    shared_history_version: u64,
4096}
4097
4098#[derive(Clone)]
4099pub(crate) struct FacContextTreeLifecycleSnapshot {
4100    shared_history: BitHistory,
4101    shared_history_version: u64,
4102    prepared: Vec<ContextTreeCorePreparedSnapshot>,
4103}
4104
4105/// Approximate heap-memory breakdown for a [`FacContextTree`].
4106#[cfg(any(test, feature = "research-tooling"))]
4107#[derive(Clone, Copy, Debug, PartialEq, Eq)]
4108pub struct FacContextTreeMemoryUsage {
4109    /// Bytes owned by per-bit CTW tree engines, including arenas and scratch buffers.
4110    pub tree_bytes: usize,
4111    /// Bytes held by the thread-local shared CTW logarithm cache.
4112    pub shared_log_cache_bytes: usize,
4113    /// Bytes reserved for the factorized shared history buffer.
4114    pub shared_history_bytes: usize,
4115}
4116
4117/// Per-tree CTW arena and scratch telemetry.
4118#[cfg(any(test, feature = "research-tooling"))]
4119#[derive(Clone, Debug, PartialEq, Eq)]
4120#[non_exhaustive]
4121pub struct FacContextTreeTreeTelemetry {
4122    /// Factorized bit position for this tree.
4123    pub bit_index: usize,
4124    /// Maximum context depth for this tree.
4125    pub max_depth: usize,
4126    /// Number of symbols observed by this tree root.
4127    pub root_visits: usize,
4128    /// Number of allocated explicit nodes.
4129    pub nodes_len: usize,
4130    /// Reserved explicit-node capacity.
4131    pub nodes_capacity: usize,
4132    /// Number of allocated unary path segments.
4133    pub segments_len: usize,
4134    /// Reserved unary-segment capacity.
4135    pub segments_capacity: usize,
4136    /// Number of node slots currently on the free list.
4137    pub free_nodes_len: usize,
4138    /// Reserved free-node list capacity.
4139    pub free_nodes_capacity: usize,
4140    /// Number of segment slots currently on the free list.
4141    pub free_segments_len: usize,
4142    /// Reserved free-segment list capacity.
4143    pub free_segments_capacity: usize,
4144    /// Reserved explicit-node bytes.
4145    pub node_bytes: usize,
4146    /// Explicit-node payload bytes at current length.
4147    pub node_payload_bytes: usize,
4148    /// Reserved segment bytes.
4149    pub segment_bytes: usize,
4150    /// Unary-segment payload bytes at current length.
4151    pub segment_payload_bytes: usize,
4152    /// Reserved free-list bytes.
4153    pub free_list_bytes: usize,
4154    /// Reserved engine scratch bytes.
4155    pub scratch_bytes: usize,
4156    /// Total reserved tree bytes for this tree.
4157    pub total_bytes: usize,
4158    /// Reserved tree bytes currently unused by arena capacity.
4159    pub arena_slack_bytes: usize,
4160    /// Number of exact-bit segment payloads.
4161    pub exact_segments: usize,
4162    /// Number of history-anchor segment payloads.
4163    pub history_segments: usize,
4164    /// Number of inverted history-anchor segment payloads.
4165    pub history_invert_segments: usize,
4166    /// Number of constant-bit segment payloads.
4167    pub const_segments: usize,
4168    /// Sum of segment lengths in bits.
4169    pub segment_bits: u64,
4170    /// Maximum segment length in bits.
4171    pub max_segment_len: u32,
4172}
4173
4174/// Detailed FAC-CTW memory telemetry.
4175#[cfg(any(test, feature = "research-tooling"))]
4176#[derive(Clone, Debug, PartialEq, Eq)]
4177#[non_exhaustive]
4178pub struct FacContextTreeTelemetry {
4179    /// Base depth used to construct tree 0.
4180    pub base_depth: usize,
4181    /// Number of factorized bit trees.
4182    pub num_bits: usize,
4183    /// Shared history length in bits.
4184    pub shared_history_len_bits: usize,
4185    /// Shared history reserved capacity in bits.
4186    pub shared_history_capacity_bits: usize,
4187    /// Reserved bytes for shared history.
4188    pub shared_history_bytes: usize,
4189    /// Shared-history payload bytes at current length.
4190    pub shared_history_payload_bytes: usize,
4191    /// Shared-history reserved slack bytes.
4192    pub shared_history_slack_bytes: usize,
4193    /// Reserved bytes for shared log caches.
4194    pub shared_log_cache_bytes: usize,
4195    /// Sum of per-tree reserved bytes.
4196    pub tree_bytes: usize,
4197    /// Sum of live per-tree node and segment payload bytes.
4198    pub tree_payload_bytes: usize,
4199    /// Sum of per-tree arena slack bytes.
4200    pub tree_arena_slack_bytes: usize,
4201    /// Total reserved bytes reported by FAC memory accounting.
4202    pub total_bytes: usize,
4203    /// Total reserved slack bytes attributable to allocator headroom.
4204    pub total_slack_bytes: usize,
4205    /// Sum of allocated explicit nodes across trees.
4206    pub nodes_len: usize,
4207    /// Sum of explicit-node capacity across trees.
4208    pub nodes_capacity: usize,
4209    /// Sum of allocated unary path segments across trees.
4210    pub segments_len: usize,
4211    /// Sum of unary-segment capacity across trees.
4212    pub segments_capacity: usize,
4213    /// Sum of node slots currently on free lists.
4214    pub free_nodes_len: usize,
4215    /// Sum of segment slots currently on free lists.
4216    pub free_segments_len: usize,
4217    /// Sum of exact-bit segment payloads across trees.
4218    pub exact_segments: usize,
4219    /// Sum of history-anchor segment payloads across trees.
4220    pub history_segments: usize,
4221    /// Sum of inverted history-anchor segment payloads across trees.
4222    pub history_invert_segments: usize,
4223    /// Sum of constant-bit segment payloads across trees.
4224    pub const_segments: usize,
4225    /// Sum of segment lengths in bits across trees.
4226    pub segment_bits: u64,
4227    /// Maximum segment length observed across trees.
4228    pub max_segment_len: u32,
4229    /// Per-tree telemetry.
4230    pub trees: Vec<FacContextTreeTreeTelemetry>,
4231}
4232
4233#[cfg(any(test, feature = "research-tooling"))]
4234impl FacContextTreeMemoryUsage {
4235    /// Total approximate heap memory in bytes.
4236    #[inline]
4237    pub fn total_bytes(self) -> usize {
4238        self.tree_bytes
4239            .saturating_add(self.shared_log_cache_bytes)
4240            .saturating_add(self.shared_history_bytes)
4241    }
4242}
4243
4244impl FacContextTree {
4245    /// Create a factorized CTW stack over `num_percept_bits` bit positions.
4246    ///
4247    /// Tree `i` uses depth `base_depth + i`, matching FAC-CTW's increasing context.
4248    pub fn new(base_depth: usize, num_percept_bits: usize) -> Self {
4249        let trees = (0..num_percept_bits)
4250            .map(|i| ContextTreeCore::new(base_depth + i))
4251            .collect();
4252        Self {
4253            trees,
4254            shared_history: BitHistory::default(),
4255            base_depth,
4256            num_bits: num_percept_bits,
4257            shared_history_version: 0,
4258        }
4259    }
4260
4261    #[inline(always)]
4262    fn bump_shared_history_version(&mut self) {
4263        self.shared_history_version = self.shared_history_version.wrapping_add(1);
4264    }
4265
4266    #[inline]
4267    /// Reserve history/tree capacity for approximately `total_symbols` updates.
4268    pub fn reserve_for_symbols(&mut self, total_symbols: usize) {
4269        if total_symbols == 0 {
4270            return;
4271        }
4272        self.shared_history
4273            .reserve_exact(total_symbols.saturating_mul(self.num_bits));
4274        for tree in &mut self.trees {
4275            tree.reserve_for_symbols(total_symbols);
4276        }
4277    }
4278
4279    #[inline]
4280    /// Number of bit positions currently modeled per symbol.
4281    pub fn num_bits(&self) -> usize {
4282        self.num_bits
4283    }
4284
4285    #[inline]
4286    /// Base depth used to construct the first factorized tree.
4287    pub fn base_depth(&self) -> usize {
4288        self.base_depth
4289    }
4290
4291    #[inline]
4292    /// Update one bit position with a binary symbol.
4293    pub fn update(&mut self, sym: Symbol, bit_index: usize) {
4294        debug_assert!(bit_index < self.num_bits);
4295        self.trees[bit_index].update_predicted(
4296            sym,
4297            &self.shared_history,
4298            self.shared_history_version,
4299        );
4300        self.shared_history.push(sym);
4301        self.bump_shared_history_version();
4302    }
4303
4304    #[inline]
4305    /// Update all bit positions from one byte, most-significant bit first.
4306    pub fn update_byte_msb(&mut self, byte: u8) {
4307        if self.num_bits != 8 {
4308            for bit_idx in 0..self.num_bits {
4309                let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
4310                self.update(bit, bit_idx);
4311            }
4312            return;
4313        }
4314
4315        let upto = self.trees[0].engine.root_visits() + 1;
4316        debug_assert!(
4317            self.trees
4318                .iter()
4319                .all(|tree| tree.engine.root_visits() + 1 == upto)
4320        );
4321        if upto <= ctw_log_cache_limit() {
4322            with_shared_cached_logs(upto, |logs| {
4323                for bit_idx in 0..8usize {
4324                    let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
4325                    let tree = &mut self.trees[bit_idx];
4326                    tree.prepared_valid = false;
4327                    tree.engine
4328                        .update_with_logs(logs, bit, &self.shared_history);
4329                    self.shared_history.push(bit);
4330                }
4331            });
4332        } else {
4333            with_shared_bounded_logs(upto, |logs| {
4334                for bit_idx in 0..8usize {
4335                    let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
4336                    let tree = &mut self.trees[bit_idx];
4337                    tree.prepared_valid = false;
4338                    tree.engine
4339                        .update_with_logs(logs, bit, &self.shared_history);
4340                    self.shared_history.push(bit);
4341                }
4342            });
4343        }
4344        self.bump_shared_history_version();
4345    }
4346
4347    #[inline]
4348    fn log_prob_update_byte_msb_with_logs<L: CtLogAccess>(&mut self, logs: L, byte: u8) -> f64 {
4349        let mut logp = 0.0;
4350        for bit_idx in 0..8usize {
4351            let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
4352            let p =
4353                self.trees[bit_idx].predict(bit, &self.shared_history, self.shared_history_version);
4354            if p.is_finite() && p > 0.0 {
4355                logp += p.ln();
4356            } else {
4357                logp = f64::NEG_INFINITY;
4358            }
4359            self.trees[bit_idx].update_predicted_with_logs(
4360                logs,
4361                bit,
4362                &self.shared_history,
4363                self.shared_history_version,
4364            );
4365            self.shared_history.push(bit);
4366            self.bump_shared_history_version();
4367        }
4368        logp
4369    }
4370
4371    #[inline]
4372    /// Return the MSB-first log probability of `byte` and then update the model.
4373    pub fn log_prob_update_byte_msb(&mut self, byte: u8) -> f64 {
4374        debug_assert_eq!(self.num_bits, 8);
4375        let upto = self.trees[0].engine.root_visits() + 1;
4376        debug_assert!(
4377            self.trees
4378                .iter()
4379                .all(|tree| tree.engine.root_visits() + 1 == upto)
4380        );
4381        if upto <= ctw_log_cache_limit() {
4382            with_shared_cached_logs(upto, |logs| {
4383                self.log_prob_update_byte_msb_with_logs(logs, byte)
4384            })
4385        } else {
4386            with_shared_bounded_logs(upto, |logs| {
4387                self.log_prob_update_byte_msb_with_logs(logs, byte)
4388            })
4389        }
4390    }
4391
4392    #[inline]
4393    /// Update all active bit positions from one byte, least-significant bit first.
4394    pub fn update_byte_lsb(&mut self, byte: u8) {
4395        let bits = self.num_bits.clamp(1, 8);
4396        let upto = self.trees[0].engine.root_visits() + 1;
4397        debug_assert!(
4398            self.trees
4399                .iter()
4400                .take(bits)
4401                .all(|tree| tree.engine.root_visits() + 1 == upto)
4402        );
4403        if upto <= ctw_log_cache_limit() {
4404            with_shared_cached_logs(upto, |logs| {
4405                for bit_idx in 0..bits {
4406                    let bit = ((byte >> bit_idx) & 1) == 1;
4407                    let tree = &mut self.trees[bit_idx];
4408                    tree.prepared_valid = false;
4409                    tree.engine
4410                        .update_with_logs(logs, bit, &self.shared_history);
4411                    self.shared_history.push(bit);
4412                }
4413            });
4414        } else {
4415            with_shared_bounded_logs(upto, |logs| {
4416                for bit_idx in 0..bits {
4417                    let bit = ((byte >> bit_idx) & 1) == 1;
4418                    let tree = &mut self.trees[bit_idx];
4419                    tree.prepared_valid = false;
4420                    tree.engine
4421                        .update_with_logs(logs, bit, &self.shared_history);
4422                    self.shared_history.push(bit);
4423                }
4424            });
4425        }
4426        self.bump_shared_history_version();
4427    }
4428
4429    #[inline]
4430    /// Commit an update using the fast-path prepared by a prior prediction call.
4431    pub fn update_predicted(&mut self, sym: Symbol, bit_index: usize) {
4432        debug_assert!(bit_index < self.num_bits);
4433        self.trees[bit_index].update_predicted(
4434            sym,
4435            &self.shared_history,
4436            self.shared_history_version,
4437        );
4438        self.shared_history.push(sym);
4439        self.bump_shared_history_version();
4440    }
4441
4442    #[inline]
4443    /// Predict a single bit probability for `bit_index`.
4444    pub fn predict(&mut self, sym: Symbol, bit_index: usize) -> f64 {
4445        debug_assert!(bit_index < self.num_bits);
4446        self.trees[bit_index].predict(sym, &self.shared_history, self.shared_history_version)
4447    }
4448
4449    #[inline]
4450    pub(crate) fn predict_one(&mut self, bit_index: usize) -> f64 {
4451        debug_assert!(bit_index < self.num_bits);
4452        self.trees[bit_index].predict_one(&self.shared_history, self.shared_history_version)
4453    }
4454
4455    #[inline]
4456    /// Revert the most recent bit update for `bit_index`.
4457    pub fn revert(&mut self, bit_index: usize) {
4458        debug_assert!(bit_index < self.num_bits);
4459        let Some(last_sym) = self.shared_history.pop() else {
4460            return;
4461        };
4462        self.trees[bit_index].revert(last_sym, &self.shared_history);
4463        self.bump_shared_history_version();
4464    }
4465
4466    #[inline]
4467    /// Append raw shared-history symbols without model updates.
4468    pub fn update_history(&mut self, symbols: &[Symbol]) {
4469        if symbols.is_empty() {
4470            return;
4471        }
4472        self.shared_history.extend_from_slice(symbols);
4473        self.bump_shared_history_version();
4474    }
4475
4476    #[inline]
4477    /// Drop the last `count` symbols from shared history.
4478    pub fn revert_history(&mut self, count: usize) {
4479        let old_len = self.shared_history.len();
4480        let new_len = self.shared_history.len().saturating_sub(count);
4481        if new_len == old_len {
4482            return;
4483        }
4484        self.shared_history.truncate(new_len);
4485        self.bump_shared_history_version();
4486    }
4487
4488    #[inline]
4489    /// Clear shared history while preserving learned tree parameters.
4490    pub fn reset_history_only(&mut self) {
4491        if self.shared_history.is_empty() {
4492            return;
4493        }
4494        self.shared_history.clear();
4495        self.bump_shared_history_version();
4496    }
4497
4498    /// Capture rollback state for stream lifecycle transactions.
4499    ///
4500    /// This clones the shared conditioning history and per-tree prepared-prefix
4501    /// state, so the allocation and copy are O(shared history length + number of
4502    /// FAC component trees). It is intended for stream lifecycle boundaries
4503    /// rather than per-symbol speculative prediction.
4504    pub(crate) fn lifecycle_snapshot(&self) -> FacContextTreeLifecycleSnapshot {
4505        FacContextTreeLifecycleSnapshot {
4506            shared_history: self.shared_history.clone(),
4507            shared_history_version: self.shared_history_version,
4508            prepared: self
4509                .trees
4510                .iter()
4511                .map(ContextTreeCore::prepared_snapshot)
4512                .collect(),
4513        }
4514    }
4515
4516    pub(crate) fn restore_lifecycle_snapshot(&mut self, snapshot: FacContextTreeLifecycleSnapshot) {
4517        debug_assert_eq!(self.trees.len(), snapshot.prepared.len());
4518        self.shared_history = snapshot.shared_history;
4519        self.shared_history_version = snapshot.shared_history_version;
4520        for (tree, prepared) in self.trees.iter_mut().zip(snapshot.prepared) {
4521            tree.restore_prepared_snapshot(prepared);
4522        }
4523    }
4524
4525    #[inline]
4526    /// Sum of per-tree log block probabilities.
4527    pub fn get_log_block_probability(&self) -> f64 {
4528        self.trees
4529            .iter()
4530            .map(|t| t.get_log_block_probability())
4531            .sum()
4532    }
4533
4534    /// Clear all trees and shared history.
4535    pub fn clear(&mut self) {
4536        for tree in &mut self.trees {
4537            tree.clear();
4538        }
4539        self.shared_history.clear();
4540        self.shared_history_version = 0;
4541    }
4542
4543    /// Approximate heap-memory usage broken down by CTW component.
4544    #[cfg(any(test, feature = "research-tooling"))]
4545    pub fn memory_usage_breakdown(&self) -> FacContextTreeMemoryUsage {
4546        let tree_mem: usize = self.trees.iter().map(|t| t.engine.memory_usage()).sum();
4547        let log_cache_mem = self
4548            .trees
4549            .first()
4550            .map(|t| t.engine.log_cache_memory_usage())
4551            .unwrap_or(0);
4552        let history_mem = self.shared_history.memory_usage();
4553        FacContextTreeMemoryUsage {
4554            tree_bytes: tree_mem,
4555            shared_log_cache_bytes: log_cache_mem,
4556            shared_history_bytes: history_mem,
4557        }
4558    }
4559
4560    /// Detailed CTW arena, segment, scratch, history, and log-cache telemetry.
4561    #[cfg(any(test, feature = "research-tooling"))]
4562    pub fn telemetry(&self) -> FacContextTreeTelemetry {
4563        let usage = self.memory_usage_breakdown();
4564        let trees: Vec<FacContextTreeTreeTelemetry> = self
4565            .trees
4566            .iter()
4567            .enumerate()
4568            .map(|(bit_index, tree)| tree.engine.telemetry(bit_index))
4569            .collect();
4570        let shared_history_payload_bytes =
4571            self.shared_history.len().div_ceil(HISTORY_WORD_BITS) * size_of::<u64>();
4572        let shared_history_slack_bytes = usage
4573            .shared_history_bytes
4574            .saturating_sub(shared_history_payload_bytes);
4575        let nodes_len = trees.iter().map(|tree| tree.nodes_len).sum();
4576        let nodes_capacity = trees.iter().map(|tree| tree.nodes_capacity).sum();
4577        let segments_len = trees.iter().map(|tree| tree.segments_len).sum();
4578        let segments_capacity = trees.iter().map(|tree| tree.segments_capacity).sum();
4579        let free_nodes_len = trees.iter().map(|tree| tree.free_nodes_len).sum();
4580        let free_segments_len = trees.iter().map(|tree| tree.free_segments_len).sum();
4581        let exact_segments = trees.iter().map(|tree| tree.exact_segments).sum();
4582        let history_segments = trees.iter().map(|tree| tree.history_segments).sum();
4583        let history_invert_segments = trees.iter().map(|tree| tree.history_invert_segments).sum();
4584        let const_segments = trees.iter().map(|tree| tree.const_segments).sum();
4585        let segment_bits = trees.iter().map(|tree| tree.segment_bits).sum();
4586        let max_segment_len = trees
4587            .iter()
4588            .map(|tree| tree.max_segment_len)
4589            .max()
4590            .unwrap_or(0);
4591        let tree_payload_bytes = trees
4592            .iter()
4593            .map(|tree| {
4594                tree.node_payload_bytes
4595                    .saturating_add(tree.segment_payload_bytes)
4596            })
4597            .sum();
4598        let tree_arena_slack_bytes: usize = trees.iter().map(|tree| tree.arena_slack_bytes).sum();
4599        let total_slack_bytes = tree_arena_slack_bytes.saturating_add(shared_history_slack_bytes);
4600
4601        FacContextTreeTelemetry {
4602            base_depth: self.base_depth,
4603            num_bits: self.num_bits,
4604            shared_history_len_bits: self.shared_history.len(),
4605            shared_history_capacity_bits: self.shared_history.words.capacity() * HISTORY_WORD_BITS,
4606            shared_history_bytes: usage.shared_history_bytes,
4607            shared_history_payload_bytes,
4608            shared_history_slack_bytes,
4609            shared_log_cache_bytes: usage.shared_log_cache_bytes,
4610            tree_bytes: usage.tree_bytes,
4611            tree_payload_bytes,
4612            tree_arena_slack_bytes,
4613            total_bytes: usage.total_bytes(),
4614            total_slack_bytes,
4615            nodes_len,
4616            nodes_capacity,
4617            segments_len,
4618            segments_capacity,
4619            free_nodes_len,
4620            free_segments_len,
4621            exact_segments,
4622            history_segments,
4623            history_invert_segments,
4624            const_segments,
4625            segment_bits,
4626            max_segment_len,
4627            trees,
4628        }
4629    }
4630
4631    /// Approximate heap memory usage in bytes.
4632    pub fn memory_usage(&self) -> usize {
4633        let tree_mem: usize = self.trees.iter().map(|t| t.engine.memory_usage()).sum();
4634        let log_cache_mem = self
4635            .trees
4636            .first()
4637            .map(|t| t.engine.log_cache_memory_usage())
4638            .unwrap_or(0);
4639        let history_mem = self.shared_history.memory_usage();
4640        tree_mem
4641            .saturating_add(log_cache_mem)
4642            .saturating_add(history_mem)
4643    }
4644}
4645
4646#[inline]
4647fn compact_symbol_msb_shift(bits_per_symbol: usize, bit_idx: usize) -> usize {
4648    bits_per_symbol.saturating_sub(1).saturating_sub(bit_idx)
4649}
4650
4651#[inline]
4652pub(crate) fn ctw_symbol_bit_msb(symbol: u8, bits_per_symbol: usize, bit_idx: usize) -> bool {
4653    let bits = bits_per_symbol.clamp(1, 8);
4654    let shift = compact_symbol_msb_shift(bits, bit_idx);
4655    ((symbol >> shift) & 1) == 1
4656}
4657
4658#[inline]
4659pub(crate) fn ctw_log_prob_msb(
4660    tree: &mut ContextTree,
4661    symbol: u8,
4662    bits_per_symbol: usize,
4663    min_prob: f64,
4664) -> f64 {
4665    let bits = bits_per_symbol.clamp(1, 8);
4666    let mut logp = 0.0;
4667    for bit_idx in 0..bits {
4668        let bit = ctw_symbol_bit_msb(symbol, bits, bit_idx);
4669        let p = tree.predict(bit);
4670        if p.is_finite() && p > 0.0 {
4671            logp += p.ln();
4672        } else {
4673            logp = f64::NEG_INFINITY;
4674        }
4675        tree.update(bit);
4676    }
4677    for _ in 0..bits {
4678        tree.revert();
4679    }
4680    if logp.is_finite() {
4681        logp.max(min_prob.ln())
4682    } else {
4683        min_prob.ln()
4684    }
4685}
4686
4687#[inline]
4688pub(crate) fn ctw_log_prob_update_msb(
4689    tree: &mut ContextTree,
4690    symbol: u8,
4691    bits_per_symbol: usize,
4692    min_prob: f64,
4693) -> f64 {
4694    let bits = bits_per_symbol.clamp(1, 8);
4695    let mut logp = 0.0;
4696    for bit_idx in 0..bits {
4697        let bit = ctw_symbol_bit_msb(symbol, bits, bit_idx);
4698        let p = tree.predict(bit);
4699        if p.is_finite() && p > 0.0 {
4700            logp += p.ln();
4701        } else {
4702            logp = f64::NEG_INFINITY;
4703        }
4704        tree.update(bit);
4705    }
4706    if logp.is_finite() {
4707        logp.max(min_prob.ln())
4708    } else {
4709        min_prob.ln()
4710    }
4711}
4712
4713#[inline]
4714pub(crate) fn ctw_log_prob_update_lsb(
4715    tree: &mut FacContextTree,
4716    symbol: u8,
4717    bits_per_symbol: usize,
4718    min_prob: f64,
4719) -> f64 {
4720    let mut logp = 0.0;
4721    for bit_idx in 0..bits_per_symbol {
4722        let bit = ((symbol >> bit_idx) & 1) == 1;
4723        let p = tree.predict(bit, bit_idx);
4724        if p.is_finite() && p > 0.0 {
4725            logp += p.ln();
4726        } else {
4727            logp = f64::NEG_INFINITY;
4728        }
4729        tree.update_predicted(bit, bit_idx);
4730    }
4731    if logp.is_finite() {
4732        logp.max(min_prob.ln())
4733    } else {
4734        min_prob.ln()
4735    }
4736}
4737
4738pub(crate) fn fill_ctw_tree_log_probs(
4739    tree: &mut ContextTree,
4740    bits_per_symbol: usize,
4741    min_logp: f64,
4742    out: &mut [f64; 256],
4743) {
4744    let bits = bits_per_symbol.clamp(1, 8);
4745    let patterns = 1usize << bits;
4746    let mut pattern_logps = [f64::NEG_INFINITY; 256];
4747    let log_before = tree.get_log_block_probability();
4748
4749    fn rec(
4750        tree: &mut ContextTree,
4751        depth: usize,
4752        bits: usize,
4753        log_before: f64,
4754        min_logp: f64,
4755        symbol_acc: u8,
4756        pattern_logps: &mut [f64; 256],
4757    ) {
4758        if depth == bits {
4759            let pat = symbol_acc as usize;
4760            let logp = (tree.get_log_block_probability() - log_before).max(min_logp);
4761            pattern_logps[pat] = logp;
4762            return;
4763        }
4764
4765        for bit in [false, true] {
4766            tree.update(bit);
4767            let shift = compact_symbol_msb_shift(bits, depth);
4768            let next_symbol = if bit {
4769                symbol_acc | (1u8 << shift)
4770            } else {
4771                symbol_acc
4772            };
4773            rec(
4774                tree,
4775                depth + 1,
4776                bits,
4777                log_before,
4778                min_logp,
4779                next_symbol,
4780                pattern_logps,
4781            );
4782            tree.revert();
4783        }
4784    }
4785
4786    rec(tree, 0, bits, log_before, min_logp, 0, &mut pattern_logps);
4787
4788    if bits == 8 {
4789        out.copy_from_slice(&pattern_logps);
4790    } else {
4791        let aliases = 1usize << (8 - bits);
4792        let alias_ln = (aliases as f64).ln();
4793        let mask = patterns - 1;
4794        for byte in 0..256usize {
4795            out[byte] = pattern_logps[byte & mask] - alias_ln;
4796        }
4797    }
4798}
4799
4800pub(crate) fn fill_fac_tree_log_probs(
4801    tree: &mut FacContextTree,
4802    bits_per_symbol: usize,
4803    msb_first: bool,
4804    min_logp: f64,
4805    out: &mut [f64; 256],
4806) {
4807    struct RecParams {
4808        bits: usize,
4809        msb_first: bool,
4810        log_before: f64,
4811        min_logp: f64,
4812    }
4813
4814    let bits = bits_per_symbol.clamp(1, 8);
4815    let patterns = 1usize << bits;
4816    let mut pattern_logps = [f64::NEG_INFINITY; 256];
4817    let params = RecParams {
4818        bits,
4819        msb_first,
4820        log_before: tree.get_log_block_probability(),
4821        min_logp,
4822    };
4823
4824    fn rec(
4825        tree: &mut FacContextTree,
4826        depth: usize,
4827        params: &RecParams,
4828        symbol_acc: u8,
4829        pattern_logps: &mut [f64; 256],
4830    ) {
4831        if depth == params.bits {
4832            let pat = symbol_acc as usize;
4833            let logp = (tree.get_log_block_probability() - params.log_before).max(params.min_logp);
4834            pattern_logps[pat] = logp;
4835            return;
4836        }
4837
4838        for bit in [false, true] {
4839            tree.update(bit, depth);
4840            let mut next_symbol = symbol_acc;
4841            if params.msb_first {
4842                // Keep sub-byte MSB symbols packed into the low pattern range so the
4843                // alias expansion below can index them via `byte & mask`.
4844                let shift = compact_symbol_msb_shift(params.bits, depth);
4845                if bit {
4846                    next_symbol |= 1u8 << shift;
4847                }
4848            } else if bit {
4849                next_symbol |= 1u8 << depth;
4850            }
4851            rec(tree, depth + 1, params, next_symbol, pattern_logps);
4852            tree.revert(depth);
4853        }
4854    }
4855
4856    rec(tree, 0, &params, 0, &mut pattern_logps);
4857
4858    if bits == 8 {
4859        out.copy_from_slice(&pattern_logps);
4860    } else {
4861        let aliases = 1usize << (8 - bits);
4862        let alias_ln = (aliases as f64).ln();
4863        let mask = patterns - 1;
4864        for byte in 0..256usize {
4865            out[byte] = pattern_logps[byte & mask] - alias_ln;
4866        }
4867    }
4868}
4869
4870#[cfg(test)]
4871mod tests {
4872    use super::*;
4873
4874    struct ScopedLogCacheLimit {
4875        previous: Option<usize>,
4876    }
4877
4878    struct ScopedLogOverflowCacheSlots {
4879        previous: Option<usize>,
4880    }
4881
4882    impl ScopedLogCacheLimit {
4883        fn set(limit: usize) -> Self {
4884            CTW_TEST_LOG_CACHE_LIMIT.with(|cell| {
4885                let previous = cell.replace(Some(limit));
4886                Self { previous }
4887            })
4888        }
4889    }
4890
4891    impl ScopedLogOverflowCacheSlots {
4892        fn set(slots: usize) -> Self {
4893            CTW_TEST_LOG_OVERFLOW_CACHE_SLOTS.with(|cell| {
4894                let previous = cell.replace(Some(slots));
4895                Self { previous }
4896            })
4897        }
4898    }
4899
4900    impl Drop for ScopedLogCacheLimit {
4901        fn drop(&mut self) {
4902            CTW_TEST_LOG_CACHE_LIMIT.with(|cell| {
4903                cell.replace(self.previous);
4904            });
4905        }
4906    }
4907
4908    impl Drop for ScopedLogOverflowCacheSlots {
4909        fn drop(&mut self) {
4910            CTW_TEST_LOG_OVERFLOW_CACHE_SLOTS.with(|cell| {
4911                cell.replace(self.previous);
4912            });
4913        }
4914    }
4915
4916    #[derive(Clone)]
4917    struct RefNode {
4918        children: [Option<Box<RefNode>>; 2],
4919        log_prob_kt: f64,
4920        log_prob_weighted: f64,
4921        symbol_count: [u32; 2],
4922    }
4923
4924    impl Default for RefNode {
4925        fn default() -> Self {
4926            Self {
4927                children: [None, None],
4928                log_prob_kt: 0.0,
4929                log_prob_weighted: 0.0,
4930                symbol_count: [0, 0],
4931            }
4932        }
4933    }
4934
4935    #[derive(Clone)]
4936    struct RefContextTree {
4937        root: RefNode,
4938        history: Vec<Symbol>,
4939        max_depth: usize,
4940        log_int: Vec<f64>,
4941        log_half: Vec<f64>,
4942    }
4943
4944    impl RefContextTree {
4945        fn new(depth: usize) -> Self {
4946            Self {
4947                root: RefNode::default(),
4948                history: Vec::new(),
4949                max_depth: depth,
4950                log_int: vec![f64::NEG_INFINITY],
4951                log_half: vec![(0.5f64).ln()],
4952            }
4953        }
4954
4955        fn root_visits(&self) -> usize {
4956            (self.root.symbol_count[0] + self.root.symbol_count[1]) as usize
4957        }
4958
4959        fn recompute(node: &mut RefNode) {
4960            let w0 = node.children[0]
4961                .as_ref()
4962                .map(|c| c.log_prob_weighted)
4963                .unwrap_or(0.0);
4964            let w1 = node.children[1]
4965                .as_ref()
4966                .map(|c| c.log_prob_weighted)
4967                .unwrap_or(0.0);
4968            let is_leaf = node.children[0].is_none() && node.children[1].is_none();
4969            node.log_prob_weighted = update_weighted_log_prob(node.log_prob_kt, w0, w1, is_leaf);
4970        }
4971
4972        fn update(&mut self, sym: Symbol) {
4973            let upto = self.root_visits() + 1;
4974            ensure_log_caches(&mut self.log_int, &mut self.log_half, upto);
4975            let sym_idx = sym as usize;
4976            Self::update_node(
4977                &mut self.root,
4978                0,
4979                self.max_depth,
4980                &self.history,
4981                sym_idx,
4982                CachedLogs::new(&self.log_int, &self.log_half),
4983            );
4984            self.history.push(sym);
4985        }
4986
4987        fn revert(&mut self) {
4988            let Some(last_sym) = self.history.pop() else {
4989                return;
4990            };
4991            let upto = self.root_visits();
4992            ensure_log_caches(&mut self.log_int, &mut self.log_half, upto);
4993            let sym_idx = last_sym as usize;
4994            let _ = Self::revert_node(
4995                &mut self.root,
4996                0,
4997                self.max_depth,
4998                &self.history,
4999                sym_idx,
5000                CachedLogs::new(&self.log_int, &self.log_half),
5001            );
5002        }
5003
5004        fn predict(&mut self, sym: Symbol) -> f64 {
5005            let sym_idx = sym as usize;
5006            let mut entries = Vec::with_capacity(self.max_depth + 1);
5007            let reached_max_depth = Self::collect_predict_entries(
5008                &self.root,
5009                0,
5010                self.max_depth,
5011                &self.history,
5012                &mut entries,
5013            );
5014
5015            let deepest = entries.len() - 1;
5016            let mut ratio = if reached_max_depth && deepest == self.max_depth {
5017                predict_ratio_kt(entries[deepest].symbol_count, sym_idx)
5018            } else {
5019                0.5
5020            };
5021            for idx in (0..=deepest).rev() {
5022                if reached_max_depth && idx == deepest {
5023                    continue;
5024                }
5025                let child_weight = if idx < deepest {
5026                    entries[idx + 1].log_prob_weighted
5027                } else {
5028                    0.0
5029                };
5030                ratio = predict_ratio_internal(
5031                    entries[idx].log_prob_kt,
5032                    entries[idx].symbol_count,
5033                    child_weight,
5034                    entries[idx].sibling_weight,
5035                    ratio,
5036                    sym_idx,
5037                );
5038            }
5039            ratio
5040        }
5041
5042        fn get_log_block_probability(&self) -> f64 {
5043            self.root.log_prob_weighted
5044        }
5045
5046        fn update_node<L: CtLogAccess>(
5047            node: &mut RefNode,
5048            depth: usize,
5049            max_depth: usize,
5050            history: &(impl HistoryAccess + ?Sized),
5051            sym_idx: usize,
5052            logs: L,
5053        ) {
5054            if depth < max_depth {
5055                let edge = history_symbol(history, depth) as usize;
5056                if node.children[edge].is_none() {
5057                    node.children[edge] = Some(Box::new(RefNode::default()));
5058                }
5059                Self::update_node(
5060                    node.children[edge].as_deref_mut().unwrap(),
5061                    depth + 1,
5062                    max_depth,
5063                    history,
5064                    sym_idx,
5065                    logs,
5066                );
5067            }
5068            apply_update_to_state_raw(logs, &mut node.symbol_count, &mut node.log_prob_kt, sym_idx);
5069            Self::recompute(node);
5070        }
5071
5072        fn revert_node<L: CtLogAccess>(
5073            node: &mut RefNode,
5074            depth: usize,
5075            max_depth: usize,
5076            history: &(impl HistoryAccess + ?Sized),
5077            sym_idx: usize,
5078            logs: L,
5079        ) -> bool {
5080            if depth < max_depth {
5081                let edge = history_symbol(history, depth) as usize;
5082                let remove_child = if let Some(child) = node.children[edge].as_deref_mut() {
5083                    Self::revert_node(child, depth + 1, max_depth, history, sym_idx, logs)
5084                } else {
5085                    false
5086                };
5087                if remove_child {
5088                    node.children[edge] = None;
5089                }
5090            }
5091            apply_revert_to_state_raw(logs, &mut node.symbol_count, &mut node.log_prob_kt, sym_idx);
5092            Self::recompute(node);
5093            node.symbol_count[0] + node.symbol_count[1] == 0
5094        }
5095
5096        fn collect_predict_entries(
5097            node: &RefNode,
5098            depth: usize,
5099            max_depth: usize,
5100            history: &(impl HistoryAccess + ?Sized),
5101            entries: &mut Vec<PredictEntry>,
5102        ) -> bool {
5103            let sibling_weight = if depth < max_depth {
5104                let path_edge = history_symbol(history, depth) as usize;
5105                node.children[path_edge ^ 1]
5106                    .as_ref()
5107                    .map(|c| c.log_prob_weighted)
5108                    .unwrap_or(0.0)
5109            } else {
5110                0.0
5111            };
5112            entries.push(PredictEntry {
5113                symbol_count: node.symbol_count,
5114                log_prob_kt: node.log_prob_kt,
5115                log_prob_weighted: node.log_prob_weighted,
5116                sibling_weight,
5117                has_sibling: depth < max_depth
5118                    && node.children[(history_symbol(history, depth) as usize) ^ 1].is_some(),
5119            });
5120            if depth == max_depth {
5121                return true;
5122            }
5123            let edge = history_symbol(history, depth) as usize;
5124            let Some(child) = node.children[edge].as_ref() else {
5125                return false;
5126            };
5127            Self::collect_predict_entries(child, depth + 1, max_depth, history, entries)
5128        }
5129    }
5130
5131    #[derive(Clone)]
5132    struct RefFacContextTree {
5133        trees: Vec<RefContextTree>,
5134        history: Vec<Symbol>,
5135    }
5136
5137    impl RefFacContextTree {
5138        fn new(base_depth: usize, num_bits: usize) -> Self {
5139            Self {
5140                trees: (0..num_bits)
5141                    .map(|i| RefContextTree::new(base_depth + i))
5142                    .collect(),
5143                history: Vec::new(),
5144            }
5145        }
5146
5147        fn update(&mut self, sym: Symbol, bit_index: usize) {
5148            let tree = &mut self.trees[bit_index];
5149            tree.history = self.history.clone();
5150            tree.update(sym);
5151            self.history.push(sym);
5152        }
5153
5154        fn predict(&mut self, sym: Symbol, bit_index: usize) -> f64 {
5155            let tree = &mut self.trees[bit_index];
5156            tree.history = self.history.clone();
5157            tree.predict(sym)
5158        }
5159
5160        fn revert(&mut self, bit_index: usize) {
5161            let Some(last_sym) = self.history.pop() else {
5162                return;
5163            };
5164            let tree = &mut self.trees[bit_index];
5165            tree.history = self.history.clone();
5166            tree.history.push(last_sym);
5167            tree.revert();
5168        }
5169
5170        fn get_log_block_probability(&self) -> f64 {
5171            self.trees
5172                .iter()
5173                .map(RefContextTree::get_log_block_probability)
5174                .sum()
5175        }
5176    }
5177
5178    fn assert_close(a: f64, b: f64) {
5179        let diff = (a - b).abs();
5180        let scale = a.abs().max(b.abs()).max(1.0);
5181        assert!(diff <= 1e-12 * scale, "a={a} b={b} diff={diff}");
5182    }
5183
5184    fn child_after_hot_prefix(
5185        tree: &ContextTree,
5186        history_before_update: &(impl HistoryAccess + ?Sized),
5187    ) -> ChildRef {
5188        let hot_prefix_depth = tree.engine.hot_prefix_depth();
5189        if hot_prefix_depth == 0 {
5190            return ChildRef::NONE;
5191        }
5192
5193        let root_edge = history_symbol(history_before_update, 0) as usize;
5194        let mut current = tree
5195            .engine
5196            .arena
5197            .child(tree.engine.root, root_edge)
5198            .as_node()
5199            .expect("hot-prefix node");
5200        for node_depth in 1..hot_prefix_depth {
5201            let edge = history_symbol(history_before_update, node_depth) as usize;
5202            current = tree
5203                .engine
5204                .arena
5205                .child(current, edge)
5206                .as_node()
5207                .expect("next hot-prefix node");
5208        }
5209        let tail_edge = history_symbol(history_before_update, hot_prefix_depth) as usize;
5210        tree.engine.arena.child(current, tail_edge)
5211    }
5212
5213    #[test]
5214    #[should_panic(expected = "ctw node index overflow")]
5215    fn node_index_from_usize_rejects_overflow() {
5216        let _ = NodeIndex::from_usize(INDEX_LIMIT);
5217    }
5218
5219    #[test]
5220    #[should_panic(expected = "ctw node index overflow")]
5221    fn node_index_from_usize_rejects_large_values() {
5222        let _ = NodeIndex::from_usize(u32::MAX as usize);
5223    }
5224
5225    #[test]
5226    fn ctw_count_lane_stays_packed() {
5227        assert_eq!(std::mem::size_of::<CtNode>(), 32);
5228    }
5229
5230    #[test]
5231    fn ctw_segment_payload_stays_packed() {
5232        assert_eq!(std::mem::size_of::<CtSegment>(), 40);
5233    }
5234
5235    #[test]
5236    fn log_lookup_matches_direct_log_formulas_past_cache_limit() {
5237        let log_int = vec![f64::NEG_INFINITY, 0.0, (2.0f64).ln()];
5238        let log_half = vec![(0.5f64).ln(), (1.5f64).ln(), (2.5f64).ln()];
5239        let lookup = BoundedLogs::new(&log_int, &log_half);
5240
5241        for n in 0..16usize {
5242            let expected_int = if n == 0 {
5243                f64::NEG_INFINITY
5244            } else {
5245                (n as f64).ln()
5246            };
5247            assert_eq!(lookup.log_int(n).to_bits(), expected_int.to_bits());
5248            assert_eq!(
5249                lookup.log_half(n).to_bits(),
5250                (n as f64 + 0.5).ln().to_bits()
5251            );
5252        }
5253    }
5254
5255    #[test]
5256    fn log_lookup_overflow_cache_reuses_hot_exact_values() {
5257        let slots = [LogCacheSlot::empty(), LogCacheSlot::empty()];
5258        let misses = Cell::new(0usize);
5259        let first = BoundedLogs::lookup_overflow(&slots, 17, || {
5260            misses.set(misses.get() + 1);
5261            (17.0f64).ln()
5262        });
5263        let second = BoundedLogs::lookup_overflow(&slots, 17, || {
5264            misses.set(misses.get() + 1);
5265            f64::NAN
5266        });
5267
5268        assert_eq!(first.to_bits(), (17.0f64).ln().to_bits());
5269        assert_eq!(second.to_bits(), first.to_bits());
5270        assert_eq!(misses.get(), 1);
5271    }
5272
5273    #[test]
5274    fn shared_log_cache_respects_test_limit() {
5275        reset_shared_log_cache_for_test();
5276        let _limit = ScopedLogCacheLimit::set(3);
5277        let _overflow_slots = ScopedLogOverflowCacheSlots::set(4);
5278        with_shared_bounded_logs(32, |lookup| {
5279            assert_eq!(lookup.log_int(32).to_bits(), (32.0f64).ln().to_bits());
5280            assert_eq!(lookup.log_half(32).to_bits(), (32.5f64).ln().to_bits());
5281        });
5282
5283        let (log_int_len, log_half_len) = shared_log_cache_lens();
5284        let (overflow_log_int_len, overflow_log_half_len) = shared_log_overflow_cache_lens();
5285        assert!(log_int_len <= 4, "log_int_len={log_int_len}");
5286        assert!(log_half_len <= 4, "log_half_len={log_half_len}");
5287        assert_eq!(overflow_log_int_len, 4);
5288        assert_eq!(overflow_log_half_len, 4);
5289    }
5290
5291    #[test]
5292    fn fac_ctw_memory_usage_breakdown_sums_to_existing_total() {
5293        let mut fac = FacContextTree::new(7, 8);
5294        let payload = b"ctw memory usage breakdown payload";
5295        for &byte in payload {
5296            fac.update_byte_msb(byte);
5297        }
5298
5299        let usage = fac.memory_usage_breakdown();
5300        assert_eq!(fac.memory_usage(), usage.total_bytes());
5301        assert!(usage.tree_bytes > 0);
5302        assert!(usage.shared_log_cache_bytes > 0);
5303        assert!(usage.shared_history_bytes > 0);
5304
5305        let telemetry = fac.telemetry();
5306        assert_eq!(telemetry.total_bytes, usage.total_bytes());
5307        assert_eq!(telemetry.tree_bytes, usage.tree_bytes);
5308        assert_eq!(
5309            telemetry.shared_log_cache_bytes,
5310            usage.shared_log_cache_bytes
5311        );
5312        assert_eq!(telemetry.shared_history_bytes, usage.shared_history_bytes);
5313        assert_eq!(telemetry.shared_history_len_bits, payload.len() * 8);
5314        assert_eq!(telemetry.trees.len(), 8);
5315        assert_eq!(
5316            telemetry.nodes_len,
5317            telemetry
5318                .trees
5319                .iter()
5320                .map(|tree| tree.nodes_len)
5321                .sum::<usize>()
5322        );
5323        assert_eq!(
5324            telemetry.segments_len,
5325            telemetry
5326                .trees
5327                .iter()
5328                .map(|tree| tree.segments_len)
5329                .sum::<usize>()
5330        );
5331        assert_eq!(
5332            telemetry.segments_len,
5333            telemetry.exact_segments
5334                + telemetry.history_segments
5335                + telemetry.history_invert_segments
5336                + telemetry.const_segments
5337        );
5338    }
5339
5340    #[test]
5341    fn bit_history_preserves_logical_bits_and_packs_memory() {
5342        let mut history = BitHistory::default();
5343        let bits: Vec<Symbol> = (0..130usize).map(|idx| (idx * 17 + 5) % 7 < 3).collect();
5344        history.extend_from_slice(&bits);
5345        assert_eq!(history.len(), bits.len());
5346        assert_eq!(history.to_vec(), bits);
5347        assert_eq!(history.memory_usage(), 3 * size_of::<u64>());
5348        for depth in 0..160usize {
5349            assert_eq!(
5350                history_symbol(&history, depth),
5351                history_symbol(&bits, depth)
5352            );
5353            for len in [0usize, 1, 2, 7, 31, 32, 33, 63, 64] {
5354                assert_eq!(
5355                    path_bits_from_history(&history, depth, len),
5356                    path_bits_from_history(&bits, depth, len),
5357                    "depth={depth} len={len}",
5358                );
5359            }
5360        }
5361
5362        let last = history.pop();
5363        assert_eq!(last, bits.last().copied());
5364        assert_eq!(history.len(), bits.len() - 1);
5365        let popped_bits = &bits[..bits.len() - 1];
5366        for depth in 0..160usize {
5367            assert_eq!(
5368                history_symbol(&history, depth),
5369                history_symbol(popped_bits, depth)
5370            );
5371            for len in [0usize, 1, 2, 7, 31, 32, 33, 63, 64] {
5372                assert_eq!(
5373                    path_bits_from_history(&history, depth, len),
5374                    path_bits_from_history(popped_bits, depth, len),
5375                    "after pop depth={depth} len={len}",
5376                );
5377            }
5378        }
5379
5380        history.truncate(65);
5381        assert_eq!(history.to_vec(), bits[..65].to_vec());
5382        assert_eq!(history.memory_usage(), 3 * size_of::<u64>());
5383
5384        history.clear();
5385        assert!(history.is_empty());
5386        assert_eq!(history.memory_usage(), 3 * size_of::<u64>());
5387    }
5388
5389    #[test]
5390    fn fac_ctw_history_memory_is_bit_packed() {
5391        let mut fac = FacContextTree::new(4, 8);
5392        fac.reserve_for_symbols(1_000);
5393        let usage = fac.memory_usage_breakdown();
5394        assert_eq!(
5395            usage.shared_history_bytes,
5396            history_word_len(8_000) * size_of::<u64>()
5397        );
5398    }
5399
5400    #[test]
5401    fn fac_ctw_memory_usage_breakdown_reports_bounded_log_cache_component() {
5402        reset_shared_log_cache_for_test();
5403        let _limit = ScopedLogCacheLimit::set(2);
5404        let _overflow_slots = ScopedLogOverflowCacheSlots::set(4);
5405        let mut fac = FacContextTree::new(7, 8);
5406        for &byte in b"bounded log cache memory component payload" {
5407            fac.update_byte_msb(byte);
5408        }
5409
5410        let usage = fac.memory_usage_breakdown();
5411        let (log_int_len, log_half_len) = shared_log_cache_lens();
5412        let (overflow_log_int_len, overflow_log_half_len) = shared_log_overflow_cache_lens();
5413        assert!(log_int_len <= 3, "log_int_len={log_int_len}");
5414        assert!(log_half_len <= 3, "log_half_len={log_half_len}");
5415        assert_eq!(overflow_log_int_len, 4);
5416        assert_eq!(overflow_log_half_len, 4);
5417        assert!(
5418            usage.shared_log_cache_bytes <= 16 * size_of::<f64>() + 32 * size_of::<LogCacheSlot>()
5419        );
5420        assert_eq!(fac.memory_usage(), usage.total_bytes());
5421    }
5422
5423    #[test]
5424    fn bounded_log_lookup_preserves_context_tree_updates_and_reverts() {
5425        let mut bounded = ContextTree::new(9);
5426        let mut unbounded = bounded.clone();
5427        let stream = b"bounded exact log lookup context-tree parity payload";
5428
5429        {
5430            let _limit = ScopedLogCacheLimit::set(3);
5431            for &byte in stream {
5432                for bit_idx in 0..8usize {
5433                    let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
5434                    bounded.update(bit);
5435                }
5436            }
5437        }
5438
5439        for &byte in stream {
5440            for bit_idx in 0..8usize {
5441                let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
5442                unbounded.update(bit);
5443            }
5444        }
5445
5446        assert_eq!(
5447            bounded.get_log_block_probability().to_bits(),
5448            unbounded.get_log_block_probability().to_bits()
5449        );
5450        for &sym in &[false, true] {
5451            assert_eq!(
5452                bounded.predict(sym).to_bits(),
5453                unbounded.predict(sym).to_bits()
5454            );
5455        }
5456
5457        {
5458            let _limit = ScopedLogCacheLimit::set(3);
5459            for _ in 0..16usize {
5460                bounded.revert();
5461            }
5462        }
5463        for _ in 0..16usize {
5464            unbounded.revert();
5465        }
5466
5467        assert_eq!(
5468            bounded.get_log_block_probability().to_bits(),
5469            unbounded.get_log_block_probability().to_bits()
5470        );
5471        for &sym in &[false, true] {
5472            assert_eq!(
5473                bounded.predict(sym).to_bits(),
5474                unbounded.predict(sym).to_bits()
5475            );
5476        }
5477    }
5478
5479    #[test]
5480    fn bounded_log_lookup_preserves_fac_byte_fast_paths() {
5481        let mut bounded_msb = FacContextTree::new(7, 8);
5482        let mut unbounded_msb = bounded_msb.clone();
5483        let mut bounded_lsb = FacContextTree::new(7, 5);
5484        let mut unbounded_lsb = bounded_lsb.clone();
5485        let stream = b"bounded exact log lookup fac byte parity payload";
5486
5487        {
5488            let _limit = ScopedLogCacheLimit::set(2);
5489            for &byte in stream {
5490                bounded_msb.update_byte_msb(byte);
5491                bounded_lsb.update_byte_lsb(byte);
5492            }
5493        }
5494        for &byte in stream {
5495            unbounded_msb.update_byte_msb(byte);
5496            unbounded_lsb.update_byte_lsb(byte);
5497        }
5498
5499        assert_eq!(
5500            bounded_msb.get_log_block_probability().to_bits(),
5501            unbounded_msb.get_log_block_probability().to_bits()
5502        );
5503        assert_eq!(
5504            bounded_lsb.get_log_block_probability().to_bits(),
5505            unbounded_lsb.get_log_block_probability().to_bits()
5506        );
5507        for bit_idx in 0..bounded_msb.num_bits() {
5508            assert_eq!(
5509                bounded_msb.predict(false, bit_idx).to_bits(),
5510                unbounded_msb.predict(false, bit_idx).to_bits()
5511            );
5512            assert_eq!(
5513                bounded_msb.predict_one(bit_idx).to_bits(),
5514                unbounded_msb.predict_one(bit_idx).to_bits()
5515            );
5516        }
5517        for bit_idx in 0..bounded_lsb.num_bits() {
5518            assert_eq!(
5519                bounded_lsb.predict(false, bit_idx).to_bits(),
5520                unbounded_lsb.predict(false, bit_idx).to_bits()
5521            );
5522            assert_eq!(
5523                bounded_lsb.predict_one(bit_idx).to_bits(),
5524                unbounded_lsb.predict_one(bit_idx).to_bits()
5525            );
5526        }
5527    }
5528
5529    #[test]
5530    fn context_tree_singleton_paths_use_hot_prefix_nodes() {
5531        let mut tree = ContextTree::new(13);
5532        tree.update(false);
5533
5534        let hot_prefix_depth = tree.engine.hot_prefix_depth();
5535        let child = tree.engine.arena.child(tree.engine.root, 0);
5536        let mut current = child.as_node().expect("hot-prefix node");
5537        let mut visited_hot_prefix_nodes = 1usize;
5538        for depth in 1..hot_prefix_depth {
5539            let next = tree.engine.arena.child(current, 0);
5540            current = next.as_node().expect("next hot-prefix node");
5541            visited_hot_prefix_nodes += 1;
5542            assert!(depth < hot_prefix_depth);
5543        }
5544        assert_eq!(visited_hot_prefix_nodes, hot_prefix_depth);
5545        let segment = tree
5546            .engine
5547            .arena
5548            .child(current, 0)
5549            .as_segment()
5550            .expect("segment tail");
5551        assert!(tree.engine.arena.child(current, 1).is_none());
5552        assert!(tree.engine.arena.segments[segment.get()].tail.is_none());
5553        assert_close(
5554            tree.engine.arena.segments[segment.get()].head_log_prob_weighted,
5555            -std::f64::consts::LN_2,
5556        );
5557        assert_close(tree.get_log_block_probability(), -std::f64::consts::LN_2);
5558    }
5559
5560    #[test]
5561    fn context_tree_missing_path_tail_uses_exact_segment_payloads() {
5562        let mut tree = ContextTree::new(13);
5563        tree.update(true);
5564        let child = tree.engine.arena.child(tree.engine.root, 0);
5565        let mut current = child.as_node().expect("hot-prefix node");
5566        for _ in 1..tree.engine.hot_prefix_depth() {
5567            current = tree
5568                .engine
5569                .arena
5570                .child(current, 0)
5571                .as_node()
5572                .expect("next hot-prefix node");
5573        }
5574        let segment = tree
5575            .engine
5576            .arena
5577            .child(current, 0)
5578            .as_segment()
5579            .expect("segment tail");
5580        let payload = tree.engine.arena.segments[segment.get()].payload;
5581        assert!(payload.is_exact());
5582        assert_eq!(
5583            payload.len() as usize,
5584            tree.engine.max_depth - tree.engine.hot_prefix_depth()
5585        );
5586        assert_eq!(payload.exact_bits() & low_bits_mask_u64(payload.len()), 0);
5587    }
5588
5589    #[test]
5590    fn context_tree_missing_path_tail_uses_const_payload_beyond_exact_limit() {
5591        let mut tree = ContextTree::new(80);
5592        let history_before = tree.history.clone();
5593        tree.update(false);
5594
5595        let segment = child_after_hot_prefix(&tree, &history_before)
5596            .as_segment()
5597            .expect("segment tail");
5598        let segment = tree.engine.arena.segments[segment.get()];
5599        assert_eq!(segment.payload.mode(), SEG_MODE_CONST);
5600        assert_eq!(
5601            segment.payload.len() as usize,
5602            tree.engine.max_depth - tree.engine.hot_prefix_depth()
5603        );
5604        assert!(!segment.payload.const_bit());
5605        assert!(segment.tail.is_none());
5606    }
5607
5608    #[test]
5609    fn context_tree_missing_path_tail_uses_history_and_const_payloads_beyond_exact_limit() {
5610        let mut tree = ContextTree::new(80);
5611        let seeded_history: Vec<Symbol> = (0..80).map(|i| (i & 1) == 1).collect();
5612        tree.update_history(&seeded_history);
5613        let history_before = tree.history.clone();
5614        tree.update(false);
5615
5616        let first_segment = child_after_hot_prefix(&tree, &history_before)
5617            .as_segment()
5618            .expect("history-backed segment tail");
5619        let first_segment = tree.engine.arena.segments[first_segment.get()];
5620        assert_eq!(first_segment.payload.mode(), SEG_MODE_HISTORY);
5621        assert_eq!(
5622            first_segment.payload.len() as usize,
5623            tree.engine.max_depth - tree.engine.hot_prefix_depth() - 1
5624        );
5625        for offset in [0usize, 1, 7, 31, 66] {
5626            assert_eq!(
5627                segment_edge_from_parts(
5628                    first_segment,
5629                    offset,
5630                    &history_before,
5631                    history_before.len()
5632                ),
5633                history_symbol(&history_before, tree.engine.hot_prefix_depth() + 1 + offset)
5634            );
5635        }
5636
5637        let tail_segment = first_segment
5638            .tail
5639            .as_segment()
5640            .expect("constant fallback tail");
5641        let tail_segment = tree.engine.arena.segments[tail_segment.get()];
5642        assert_eq!(tail_segment.payload.mode(), SEG_MODE_CONST);
5643        assert_eq!(tail_segment.payload.len(), 1);
5644        assert!(!tail_segment.payload.const_bit());
5645        assert!(tail_segment.tail.is_none());
5646    }
5647
5648    #[test]
5649    fn context_tree_matches_reference_on_short_sequences() {
5650        for depth in 0..=6usize {
5651            for len in 0..=6usize {
5652                for mask in 0..(1usize << len) {
5653                    let mut prod = ContextTree::new(depth);
5654                    let mut reference = RefContextTree::new(depth);
5655                    for step in 0..len {
5656                        let p_prod_0 = prod.predict(false);
5657                        let p_ref_0 = reference.predict(false);
5658                        assert!(
5659                            (p_prod_0 - p_ref_0).abs()
5660                                <= 1e-12 * p_prod_0.abs().max(p_ref_0.abs()).max(1.0),
5661                            "predict0 mismatch depth={depth} len={len} mask={mask} step={step} prod={p_prod_0} ref={p_ref_0} history={:?}",
5662                            prod.history
5663                        );
5664                        let p_prod_1 = prod.predict(true);
5665                        let p_ref_1 = reference.predict(true);
5666                        assert!(
5667                            (p_prod_1 - p_ref_1).abs()
5668                                <= 1e-12 * p_prod_1.abs().max(p_ref_1.abs()).max(1.0),
5669                            "predict1 mismatch depth={depth} len={len} mask={mask} step={step} prod={p_prod_1} ref={p_ref_1} history={:?}",
5670                            prod.history
5671                        );
5672                        let log_prod = prod.get_log_block_probability();
5673                        let log_ref = reference.get_log_block_probability();
5674                        assert!(
5675                            (log_prod - log_ref).abs()
5676                                <= 1e-12 * log_prod.abs().max(log_ref.abs()).max(1.0),
5677                            "log mismatch before update depth={depth} len={len} mask={mask} step={step} prod={log_prod} ref={log_ref} history={:?}",
5678                            prod.history
5679                        );
5680                        let bit = ((mask >> step) & 1) == 1;
5681                        prod.update(bit);
5682                        reference.update(bit);
5683                        let log_prod = prod.get_log_block_probability();
5684                        let log_ref = reference.get_log_block_probability();
5685                        assert!(
5686                            (log_prod - log_ref).abs()
5687                                <= 1e-12 * log_prod.abs().max(log_ref.abs()).max(1.0),
5688                            "log mismatch after update depth={depth} len={len} mask={mask} step={step} bit={bit} prod={log_prod} ref={log_ref} history={:?}",
5689                            prod.history
5690                        );
5691                    }
5692                    while prod.history_size() > 0 {
5693                        let p_prod_0 = prod.predict(false);
5694                        let p_ref_0 = reference.predict(false);
5695                        assert!(
5696                            (p_prod_0 - p_ref_0).abs()
5697                                <= 1e-12 * p_prod_0.abs().max(p_ref_0.abs()).max(1.0),
5698                            "revert predict0 mismatch depth={depth} len={len} mask={mask} prod={p_prod_0} ref={p_ref_0} history={:?}",
5699                            prod.history
5700                        );
5701                        let p_prod_1 = prod.predict(true);
5702                        let p_ref_1 = reference.predict(true);
5703                        assert!(
5704                            (p_prod_1 - p_ref_1).abs()
5705                                <= 1e-12 * p_prod_1.abs().max(p_ref_1.abs()).max(1.0),
5706                            "revert predict1 mismatch depth={depth} len={len} mask={mask} prod={p_prod_1} ref={p_ref_1} history={:?}",
5707                            prod.history
5708                        );
5709                        prod.revert();
5710                        reference.revert();
5711                        let log_prod = prod.get_log_block_probability();
5712                        let log_ref = reference.get_log_block_probability();
5713                        assert!(
5714                            (log_prod - log_ref).abs()
5715                                <= 1e-12 * log_prod.abs().max(log_ref.abs()).max(1.0),
5716                            "revert log mismatch depth={depth} len={len} mask={mask} prod={log_prod} ref={log_ref} history={:?}",
5717                            prod.history
5718                        );
5719                    }
5720                }
5721            }
5722        }
5723    }
5724
5725    #[test]
5726    fn context_tree_long_depth_matches_reference_on_short_sequences() {
5727        for &depth in &[65usize, 80usize] {
5728            for len in 0..=6usize {
5729                for mask in 0..(1usize << len) {
5730                    let mut prod = ContextTree::new(depth);
5731                    let mut reference = RefContextTree::new(depth);
5732                    for step in 0..len {
5733                        let p_prod_0 = prod.predict(false);
5734                        let p_ref_0 = reference.predict(false);
5735                        assert!(
5736                            (p_prod_0 - p_ref_0).abs()
5737                                <= 1e-12 * p_prod_0.abs().max(p_ref_0.abs()).max(1.0),
5738                            "long-depth predict0 mismatch depth={depth} len={len} mask={mask} step={step} prod={p_prod_0} ref={p_ref_0} history={:?}",
5739                            prod.history
5740                        );
5741                        let p_prod_1 = prod.predict(true);
5742                        let p_ref_1 = reference.predict(true);
5743                        assert!(
5744                            (p_prod_1 - p_ref_1).abs()
5745                                <= 1e-12 * p_prod_1.abs().max(p_ref_1.abs()).max(1.0),
5746                            "long-depth predict1 mismatch depth={depth} len={len} mask={mask} step={step} prod={p_prod_1} ref={p_ref_1} history={:?}",
5747                            prod.history
5748                        );
5749                        assert_close(
5750                            prod.get_log_block_probability(),
5751                            reference.get_log_block_probability(),
5752                        );
5753                        let bit = ((mask >> step) & 1) == 1;
5754                        prod.update(bit);
5755                        reference.update(bit);
5756                        assert_close(
5757                            prod.get_log_block_probability(),
5758                            reference.get_log_block_probability(),
5759                        );
5760                    }
5761
5762                    while prod.history_size() > 0 {
5763                        assert_close(prod.predict(false), reference.predict(false));
5764                        assert_close(prod.predict(true), reference.predict(true));
5765                        prod.revert();
5766                        reference.revert();
5767                        assert_close(
5768                            prod.get_log_block_probability(),
5769                            reference.get_log_block_probability(),
5770                        );
5771                    }
5772                }
5773            }
5774        }
5775    }
5776
5777    #[test]
5778    fn fac_ctw_matches_reference_on_short_sequences() {
5779        let mut fac = FacContextTree::new(4, 4);
5780        let mut reference = RefFacContextTree::new(4, 4);
5781        let stream = [
5782            (true, 0usize),
5783            (false, 1usize),
5784            (true, 2usize),
5785            (true, 3usize),
5786            (false, 0usize),
5787            (false, 1usize),
5788            (true, 2usize),
5789            (false, 3usize),
5790        ];
5791
5792        for &(bit, idx) in &stream {
5793            assert_close(fac.predict(false, idx), reference.predict(false, idx));
5794            assert_close(fac.predict(true, idx), reference.predict(true, idx));
5795            fac.update(bit, idx);
5796            reference.update(bit, idx);
5797            assert_close(
5798                fac.get_log_block_probability(),
5799                reference.get_log_block_probability(),
5800            );
5801        }
5802
5803        for &(_, idx) in stream.iter().rev() {
5804            fac.revert(idx);
5805            reference.revert(idx);
5806            assert_close(
5807                fac.get_log_block_probability(),
5808                reference.get_log_block_probability(),
5809            );
5810        }
5811    }
5812
5813    #[test]
5814    fn fac_ctw_history_consistency() {
5815        let mut fac = FacContextTree::new(4, 4);
5816
5817        fac.update_history(&[true, false, true]);
5818        assert_eq!(fac.shared_history.len(), 3);
5819
5820        fac.update(true, 0);
5821        fac.update(false, 1);
5822        assert_eq!(fac.shared_history.len(), 5);
5823
5824        fac.revert(1);
5825        assert_eq!(fac.shared_history.len(), 4);
5826
5827        fac.revert(0);
5828        assert_eq!(fac.shared_history.len(), 3);
5829    }
5830
5831    #[test]
5832    fn fac_ctw_predict_one_matches_predict_true() {
5833        let mut fac = FacContextTree::new(6, 8);
5834        for &byte in b"predict-one exactness regression payload" {
5835            for bit_idx in 0..8usize {
5836                let p_generic = fac.predict(true, bit_idx);
5837                let p_one = fac.predict_one(bit_idx);
5838                assert_close(p_generic, p_one);
5839                let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
5840                fac.update_predicted(bit, bit_idx);
5841            }
5842        }
5843    }
5844
5845    #[test]
5846    fn fac_ctw_long_depth_predict_one_matches_predict_true() {
5847        let mut fac = FacContextTree::new(78, 4);
5848        for step in 0..24usize {
5849            for bit_idx in 0..fac.num_bits() {
5850                let p_generic = fac.predict(true, bit_idx);
5851                let p_one = fac.predict_one(bit_idx);
5852                assert_close(p_generic, p_one);
5853                let bit = ((step * 5 + bit_idx * 3) & 1) == 1;
5854                fac.update_predicted(bit, bit_idx);
5855            }
5856        }
5857    }
5858
5859    #[test]
5860    fn fac_ctw_update_byte_msb_matches_bit_updates() {
5861        let mut by_byte = FacContextTree::new(6, 8);
5862        let mut by_bits = FacContextTree::new(6, 8);
5863        for &byte in b"byte update msb regression payload" {
5864            by_byte.update_byte_msb(byte);
5865            for bit_idx in 0..8usize {
5866                let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
5867                by_bits.update(bit, bit_idx);
5868            }
5869            assert_close(
5870                by_byte.get_log_block_probability(),
5871                by_bits.get_log_block_probability(),
5872            );
5873            assert_eq!(by_byte.shared_history, by_bits.shared_history);
5874        }
5875    }
5876
5877    #[test]
5878    fn fac_ctw_log_prob_update_byte_msb_matches_manual_fast_path() {
5879        let mut batched = FacContextTree::new(6, 8);
5880        for &byte in b"log prob update byte msb regression payload" {
5881            let mut manual = batched.clone();
5882            let observed = batched.log_prob_update_byte_msb(byte);
5883            let mut expected = 0.0;
5884            for bit_idx in 0..8usize {
5885                let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
5886                let p = manual.predict(bit, bit_idx);
5887                if p.is_finite() && p > 0.0 {
5888                    expected += p.ln();
5889                } else {
5890                    expected = f64::NEG_INFINITY;
5891                }
5892                manual.update_predicted(bit, bit_idx);
5893            }
5894            assert_eq!(observed.to_bits(), expected.to_bits());
5895            assert_eq!(batched.shared_history, manual.shared_history);
5896            assert_eq!(
5897                batched.get_log_block_probability().to_bits(),
5898                manual.get_log_block_probability().to_bits(),
5899            );
5900        }
5901    }
5902
5903    #[test]
5904    fn fac_ctw_update_byte_lsb_matches_bit_updates() {
5905        let mut by_byte = FacContextTree::new(6, 5);
5906        let mut by_bits = FacContextTree::new(6, 5);
5907        for &byte in b"byte update lsb regression payload" {
5908            by_byte.update_byte_lsb(byte);
5909            for bit_idx in 0..5usize {
5910                let bit = ((byte >> bit_idx) & 1) == 1;
5911                by_bits.update(bit, bit_idx);
5912            }
5913            assert_close(
5914                by_byte.get_log_block_probability(),
5915                by_bits.get_log_block_probability(),
5916            );
5917            assert_eq!(by_byte.shared_history, by_bits.shared_history);
5918        }
5919    }
5920
5921    #[test]
5922    fn fac_ctw_log_cache_tracks_tree_visits_not_shared_history() {
5923        let mut fac = FacContextTree::new(8, 8);
5924        let updates_per_tree = 512usize;
5925        let (log_int_before, log_half_before) = shared_log_cache_lens();
5926
5927        for step in 0..updates_per_tree {
5928            let bit = (step & 1) == 1;
5929            for bit_idx in 0..8usize {
5930                fac.update(bit, bit_idx);
5931            }
5932        }
5933
5934        assert_eq!(fac.shared_history.len(), updates_per_tree * 8);
5935        for tree in &fac.trees {
5936            let visits = tree.engine.arena.visits(tree.engine.root) as usize;
5937            assert_eq!(visits, updates_per_tree);
5938        }
5939
5940        let (log_int_after, log_half_after) = shared_log_cache_lens();
5941        let expected_len = updates_per_tree + 1;
5942        assert!(
5943            log_int_after <= log_int_before.max(expected_len),
5944            "log_int grew to {log_int_after} (before={log_int_before}, expected_len={expected_len})"
5945        );
5946        assert!(
5947            log_half_after <= log_half_before.max(expected_len),
5948            "log_half grew to {log_half_after} (before={log_half_before}, expected_len={expected_len})"
5949        );
5950    }
5951
5952    fn seed_fac_cache_regression_state(fac: &mut FacContextTree) {
5953        for step in 0..24usize {
5954            for bit_idx in 0..fac.num_bits() {
5955                let bit = ((step * 3 + bit_idx) & 1) == 1;
5956                fac.update(bit, bit_idx);
5957            }
5958        }
5959    }
5960
5961    fn assert_update_predicted_matches_fresh_after_history_rewrite<F>(mut rewrite: F)
5962    where
5963        F: FnMut(&mut FacContextTree),
5964    {
5965        let mut predicted = FacContextTree::new(6, 4);
5966        seed_fac_cache_regression_state(&mut predicted);
5967        let mut fresh = predicted.clone();
5968        let original_history = predicted.shared_history.clone();
5969        let target_bit = 2usize;
5970
5971        let _ = predicted.predict(true, target_bit);
5972        rewrite(&mut predicted);
5973        rewrite(&mut fresh);
5974
5975        assert_eq!(predicted.shared_history.len(), original_history.len());
5976        assert_ne!(predicted.shared_history, original_history);
5977        assert_eq!(predicted.shared_history, fresh.shared_history);
5978
5979        predicted.update_predicted(false, target_bit);
5980        fresh.update(false, target_bit);
5981
5982        assert_eq!(predicted.shared_history, fresh.shared_history);
5983        assert_close(
5984            predicted.get_log_block_probability(),
5985            fresh.get_log_block_probability(),
5986        );
5987        for bit_idx in 0..predicted.num_bits() {
5988            assert_close(
5989                predicted.predict(false, bit_idx),
5990                fresh.predict(false, bit_idx),
5991            );
5992            assert_close(
5993                predicted.predict(true, bit_idx),
5994                fresh.predict(true, bit_idx),
5995            );
5996        }
5997    }
5998
5999    #[test]
6000    fn fac_ctw_update_predicted_ignores_stale_cache_after_reset_and_rewrite() {
6001        assert_update_predicted_matches_fresh_after_history_rewrite(|fac| {
6002            let mut rewritten = fac.shared_history.to_vec();
6003            for bit in &mut rewritten {
6004                *bit = !*bit;
6005            }
6006            fac.reset_history_only();
6007            fac.update_history(&rewritten);
6008        });
6009    }
6010
6011    #[test]
6012    fn fac_ctw_update_predicted_ignores_stale_cache_after_revert_and_rewrite() {
6013        assert_update_predicted_matches_fresh_after_history_rewrite(|fac| {
6014            let original = fac.shared_history.to_vec();
6015            let keep = original.len() / 3;
6016            let remove = original.len() - keep;
6017            let mut rewritten_suffix = original[keep..].to_vec();
6018            for bit in &mut rewritten_suffix {
6019                *bit = !*bit;
6020            }
6021            fac.revert_history(remove);
6022            fac.update_history(&rewritten_suffix);
6023        });
6024    }
6025
6026    #[test]
6027    fn fac_ctw_shared_history_version_tracks_mutations() {
6028        let mut fac = FacContextTree::new(4, 2);
6029        let mut version = fac.shared_history_version;
6030
6031        fac.update_history(&[]);
6032        assert_eq!(fac.shared_history_version, version);
6033
6034        fac.update_history(&[true, false]);
6035        assert_ne!(fac.shared_history_version, version);
6036        version = fac.shared_history_version;
6037
6038        fac.revert_history(0);
6039        assert_eq!(fac.shared_history_version, version);
6040
6041        fac.revert_history(1);
6042        assert_ne!(fac.shared_history_version, version);
6043        version = fac.shared_history_version;
6044
6045        let _ = fac.predict(true, 0);
6046        assert_eq!(fac.shared_history_version, version);
6047
6048        fac.update_predicted(true, 0);
6049        assert_ne!(fac.shared_history_version, version);
6050        version = fac.shared_history_version;
6051
6052        fac.reset_history_only();
6053        assert_ne!(fac.shared_history_version, version);
6054    }
6055
6056    #[test]
6057    fn context_tree_predict_preserves_state() {
6058        let mut tree = ContextTree::new(6);
6059        for &bit in &[true, false, true, true, false, false, true, false] {
6060            tree.update(bit);
6061        }
6062        let p0_before = tree.predict(false);
6063        let p1_before = tree.predict(true);
6064        let log_before = tree.get_log_block_probability();
6065        let history_before = tree.history.clone();
6066        let _ = tree.predict(true);
6067
6068        assert_eq!(tree.history, history_before);
6069        assert_close(tree.get_log_block_probability(), log_before);
6070        assert_close(tree.predict(false), p0_before);
6071        assert_close(tree.predict(true), p1_before);
6072    }
6073
6074    #[test]
6075    fn context_tree_predict_matches_update_ratio() {
6076        let mut tree = ContextTree::new(7);
6077        for &bit in &[true, false, true, false, true, true, false, true, false] {
6078            tree.update(bit);
6079        }
6080        for &sym in &[false, true] {
6081            let predicted = tree.predict(sym);
6082            let mut reference = tree.clone();
6083            let before = reference.get_log_block_probability();
6084            reference.update(sym);
6085            let after = reference.get_log_block_probability();
6086            assert_close(predicted, (after - before).exp());
6087        }
6088    }
6089
6090    #[test]
6091    fn fac_ctw_predict_preserves_state() {
6092        let mut fac = FacContextTree::new(5, 8);
6093        for &byte in b"fac ctw state preservation" {
6094            for bit_idx in 0..8usize {
6095                let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
6096                fac.update(bit, bit_idx);
6097            }
6098        }
6099        let p0_before = fac.predict(false, 3);
6100        let p1_before = fac.predict(true, 3);
6101        let log_before = fac.get_log_block_probability();
6102        let history_before = fac.shared_history.clone();
6103        let _ = fac.predict(true, 3);
6104
6105        assert_eq!(fac.shared_history, history_before);
6106        assert_close(fac.get_log_block_probability(), log_before);
6107        assert_close(fac.predict(false, 3), p0_before);
6108        assert_close(fac.predict(true, 3), p1_before);
6109    }
6110
6111    #[test]
6112    fn ctw_inline_node_weight_recompute_matches_recompute_node_weight() {
6113        fn inline_weight(arena: &CtArena, idx: NodeIndex) -> f64 {
6114            let slot = idx.get();
6115            let node = arena.nodes[slot];
6116            let [left, right] = node.children;
6117            if left.is_none() && right.is_none() {
6118                clamp_log_prob(node.log_prob_kt)
6119            } else {
6120                let w0 = arena.child_ref_weighted(left);
6121                let w1 = arena.child_ref_weighted(right);
6122                update_weighted_log_prob_non_leaf(node.log_prob_kt, w0, w1)
6123            }
6124        }
6125
6126        let mut arena = CtArena::new();
6127        let parent = arena.alloc_node_with_state([3, 5], -1.75);
6128        let left_node = arena.alloc_node_with_state([2, 1], -0.25);
6129        let right_node = arena.alloc_node_with_state([1, 2], -0.50);
6130        arena.nodes[left_node.get()].log_prob_weighted = -0.333_333_333_f64;
6131        arena.nodes[right_node.get()].log_prob_weighted = -0.777_777_777_f64;
6132
6133        let seg = arena.alloc_segment();
6134        arena.segments[seg.get()].head_log_prob_weighted = -0.125_f64;
6135
6136        let cases = [
6137            (ChildRef::NONE, ChildRef::NONE, -2.0_f64),
6138            (
6139                ChildRef::from_node(left_node),
6140                ChildRef::from_node(right_node),
6141                -1.0_f64,
6142            ),
6143            (
6144                ChildRef::from_node(left_node),
6145                ChildRef::from_segment(seg),
6146                -0.625_f64,
6147            ),
6148            (
6149                ChildRef::from_segment(seg),
6150                ChildRef::from_node(right_node),
6151                -0.3125_f64,
6152            ),
6153        ];
6154
6155        for (left, right, kt) in cases {
6156            let slot = parent.get();
6157            arena.nodes[slot].children = [left, right];
6158            arena.nodes[slot].log_prob_kt = kt;
6159            arena.nodes[slot].log_prob_weighted = f64::NAN;
6160
6161            let expected = inline_weight(&arena, parent);
6162            arena.recompute_node_weight(parent);
6163            let actual = arena.nodes[slot].log_prob_weighted;
6164            assert_eq!(
6165                actual.to_bits(),
6166                expected.to_bits(),
6167                "node-weight recompute mismatch for children=({left:?},{right:?}) kt={kt}"
6168            );
6169        }
6170    }
6171
6172    #[test]
6173    fn fac_ctw_predict_matches_update_ratio() {
6174        let mut fac = FacContextTree::new(6, 8);
6175        for &byte in b"fac ctw exact predictive ratio" {
6176            for bit_idx in 0..8usize {
6177                let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
6178                fac.update(bit, bit_idx);
6179            }
6180        }
6181        for &sym in &[false, true] {
6182            let predicted = fac.predict(sym, 4);
6183            let mut reference = fac.clone();
6184            let before = reference.get_log_block_probability();
6185            reference.update(sym, 4);
6186            let after = reference.get_log_block_probability();
6187            assert_close(predicted, (after - before).exp());
6188        }
6189    }
6190
6191    #[test]
6192    fn fac_ctw_update_predicted_matches_fresh_update_on_byte_stream() {
6193        let mut predicted = FacContextTree::new(6, 8);
6194        let mut fresh = predicted.clone();
6195        let stream = b"fac-ctw prepared update exactness regression";
6196
6197        for (byte_pos, &byte) in stream.iter().enumerate() {
6198            for bit_idx in 0..8usize {
6199                let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
6200                let _ = predicted.predict(true, bit_idx);
6201                predicted.update_predicted(bit, bit_idx);
6202                fresh.update(bit, bit_idx);
6203
6204                let predicted_log = predicted.get_log_block_probability();
6205                let fresh_log = fresh.get_log_block_probability();
6206                assert!(
6207                    (predicted_log - fresh_log).abs()
6208                        <= 1e-12 * predicted_log.abs().max(fresh_log.abs()).max(1.0),
6209                    "log mismatch byte_pos={byte_pos} bit_idx={bit_idx} bit={bit} predicted={predicted_log} fresh={fresh_log}\nshared_history={:?}\npredicted_arena={:#?}\nfresh_arena={:#?}\npredicted_steps={:?}\nfresh_steps={:?}",
6210                    predicted.shared_history,
6211                    predicted.trees[bit_idx].engine.arena,
6212                    fresh.trees[bit_idx].engine.arena,
6213                    predicted.trees[bit_idx].engine.prepared_steps,
6214                    fresh.trees[bit_idx].engine.prepared_steps,
6215                );
6216                for probe_idx in 0..8usize {
6217                    let p_pred_0 = predicted.predict(false, probe_idx);
6218                    let p_fresh_0 = fresh.predict(false, probe_idx);
6219                    assert!(
6220                        (p_pred_0 - p_fresh_0).abs()
6221                            <= 1e-12 * p_pred_0.abs().max(p_fresh_0.abs()).max(1.0),
6222                        "predict0 mismatch byte_pos={byte_pos} bit_idx={bit_idx} probe_idx={probe_idx} predicted={p_pred_0} fresh={p_fresh_0}",
6223                    );
6224                    let p_pred_1 = predicted.predict(true, probe_idx);
6225                    let p_fresh_1 = fresh.predict(true, probe_idx);
6226                    assert!(
6227                        (p_pred_1 - p_fresh_1).abs()
6228                            <= 1e-12 * p_pred_1.abs().max(p_fresh_1.abs()).max(1.0),
6229                        "predict1 mismatch byte_pos={byte_pos} bit_idx={bit_idx} probe_idx={probe_idx} predicted={p_pred_1} fresh={p_fresh_1}",
6230                    );
6231                }
6232            }
6233        }
6234    }
6235
6236    #[test]
6237    fn fac_ctw_long_depth_update_predicted_matches_fresh_update_on_bit_stream() {
6238        let mut predicted = FacContextTree::new(78, 4);
6239        let mut fresh = predicted.clone();
6240
6241        for step in 0..20usize {
6242            for bit_idx in 0..predicted.num_bits() {
6243                let bit = ((step * 7 + bit_idx * 11) & 1) == 1;
6244                let _ = predicted.predict(true, bit_idx);
6245                predicted.update_predicted(bit, bit_idx);
6246                fresh.update(bit, bit_idx);
6247                assert_eq!(predicted.shared_history, fresh.shared_history);
6248                assert_close(
6249                    predicted.get_log_block_probability(),
6250                    fresh.get_log_block_probability(),
6251                );
6252            }
6253        }
6254
6255        for bit_idx in 0..predicted.num_bits() {
6256            assert_close(
6257                predicted.predict(false, bit_idx),
6258                fresh.predict(false, bit_idx),
6259            );
6260            assert_close(
6261                predicted.predict(true, bit_idx),
6262                fresh.predict(true, bit_idx),
6263            );
6264        }
6265    }
6266
6267    fn scan_symbol_space(tree: &mut FacContextTree, bits: usize) {
6268        fn rec(tree: &mut FacContextTree, bits: usize, depth: usize) {
6269            if depth == bits {
6270                return;
6271            }
6272            for bit in [false, true] {
6273                let bit_idx = depth;
6274                tree.update(bit, bit_idx);
6275                rec(tree, bits, depth + 1);
6276                tree.revert(bit_idx);
6277            }
6278        }
6279        rec(tree, bits, 0);
6280    }
6281
6282    fn fac_symbol_bit(symbol: u8, msb_first: bool, bits: usize, bit_idx: usize) -> bool {
6283        if msb_first {
6284            ctw_symbol_bit_msb(symbol, bits, bit_idx)
6285        } else {
6286            ((symbol >> bit_idx) & 1) == 1
6287        }
6288    }
6289
6290    fn symbol_log_prob(tree: &mut FacContextTree, symbol: u8, msb_first: bool, bits: usize) -> f64 {
6291        let before = tree.get_log_block_probability();
6292        for bit_idx in 0..bits {
6293            let bit = fac_symbol_bit(symbol, msb_first, bits, bit_idx);
6294            tree.update(bit, bit_idx);
6295        }
6296        let after = tree.get_log_block_probability();
6297        for bit_idx in (0..bits).rev() {
6298            tree.revert(bit_idx);
6299        }
6300        after - before
6301    }
6302
6303    fn assert_symbol_scan_then_update_matches_plain(msb_first: bool) {
6304        let bits = 8usize;
6305        let mut with_scan = FacContextTree::new(7, bits);
6306        let mut plain = with_scan.clone();
6307        for &byte in b"pdf then update parity payload" {
6308            for bit_idx in 0..bits {
6309                let bit = fac_symbol_bit(byte, msb_first, bits, bit_idx);
6310                with_scan.update(bit, bit_idx);
6311                plain.update(bit, bit_idx);
6312            }
6313        }
6314
6315        scan_symbol_space(&mut with_scan, bits);
6316
6317        let observed = b'n';
6318        for bit_idx in 0..bits {
6319            let bit = fac_symbol_bit(observed, msb_first, bits, bit_idx);
6320            with_scan.update(bit, bit_idx);
6321            plain.update(bit, bit_idx);
6322        }
6323
6324        for sym in 0u8..=255u8 {
6325            let lp_scan = symbol_log_prob(&mut with_scan, sym, msb_first, bits);
6326            let lp_plain = symbol_log_prob(&mut plain, sym, msb_first, bits);
6327            let diff = (lp_scan - lp_plain).abs();
6328            assert!(
6329                diff < 1e-12,
6330                "symbol={sym} lp_scan={lp_scan} lp_plain={lp_plain} diff={diff}",
6331            );
6332        }
6333    }
6334
6335    fn assert_fill_fac_tree_log_probs_matches_direct_symbol_probs(msb_first: bool, bits: usize) {
6336        let bits = bits.clamp(1, 8);
6337        let min_logp = 1e-12f64.ln();
6338        let patterns = 1usize << bits;
6339        let alias_ln = if bits == 8 {
6340            0.0
6341        } else {
6342            ((1usize << (8 - bits)) as f64).ln()
6343        };
6344        let training = [0x0u8, 0x3, 0x5, 0x6, 0x9, 0xA, 0xC, 0xF, 0x7, 0x1];
6345        let mut tree = FacContextTree::new(7, bits);
6346        for &symbol in &training {
6347            for bit_idx in 0..bits {
6348                tree.update(fac_symbol_bit(symbol, msb_first, bits, bit_idx), bit_idx);
6349            }
6350        }
6351
6352        let log_before = tree.get_log_block_probability();
6353        let mut predict_zero_before = vec![0.0; bits];
6354        let mut predict_one_before = vec![0.0; bits];
6355        for bit_idx in 0..bits {
6356            predict_zero_before[bit_idx] = tree.predict(false, bit_idx);
6357            predict_one_before[bit_idx] = tree.predict(true, bit_idx);
6358        }
6359
6360        let mut out = [0.0; 256];
6361        fill_fac_tree_log_probs(&mut tree, bits, msb_first, min_logp, &mut out);
6362
6363        assert_close(tree.get_log_block_probability(), log_before);
6364        for bit_idx in 0..bits {
6365            assert_close(tree.predict(false, bit_idx), predict_zero_before[bit_idx]);
6366            assert_close(tree.predict(true, bit_idx), predict_one_before[bit_idx]);
6367        }
6368
6369        let mask = patterns - 1;
6370        for (byte, &actual) in out.iter().enumerate() {
6371            let symbol = if bits == 8 {
6372                byte as u8
6373            } else {
6374                (byte & mask) as u8
6375            };
6376            let expected =
6377                symbol_log_prob(&mut tree, symbol, msb_first, bits).max(min_logp) - alias_ln;
6378            assert_close(actual, expected);
6379        }
6380    }
6381
6382    #[test]
6383    fn fac_ctw_symbol_scan_then_update_matches_plain_msb() {
6384        assert_symbol_scan_then_update_matches_plain(true);
6385    }
6386
6387    #[test]
6388    fn fac_ctw_symbol_scan_then_update_matches_plain_lsb() {
6389        assert_symbol_scan_then_update_matches_plain(false);
6390    }
6391
6392    #[test]
6393    fn fill_fac_tree_log_probs_matches_direct_symbol_probs_for_subbyte_msb() {
6394        assert_fill_fac_tree_log_probs_matches_direct_symbol_probs(true, 4);
6395    }
6396
6397    #[test]
6398    fn fill_fac_tree_log_probs_matches_direct_symbol_probs_for_subbyte_lsb() {
6399        assert_fill_fac_tree_log_probs_matches_direct_symbol_probs(false, 4);
6400    }
6401}