1use std::hash::{Hash, Hasher};
8
9use casefold::index_fold_char;
10
11use crate::MAX_SPARSE_GRAM_SIZE;
12use crate::ngram::NGram;
13use crate::table::{bigram_h, bigram_priority_rolling};
14
15const RING_MASK: u32 = MAX_SPARSE_GRAM_SIZE as u32 - 1;
17
18#[derive(Clone, Copy, Debug)]
19struct PosState {
20 index: u32,
21 value: u32,
22}
23
24#[derive(Clone, Debug, Default)]
25struct Queue {
26 idx_buf: [u32; MAX_SPARSE_GRAM_SIZE],
27 val_buf: [u32; MAX_SPARSE_GRAM_SIZE],
28 head: u32,
29 len: u32,
30}
31
32impl Queue {
33 fn new() -> Self {
34 Self {
35 idx_buf: [0; MAX_SPARSE_GRAM_SIZE],
36 val_buf: [0; MAX_SPARSE_GRAM_SIZE],
37 head: 0,
38 len: 0,
39 }
40 }
41
42 fn clear(&mut self) {
43 self.len = 0;
44 }
45
46 fn is_empty(&self) -> bool {
47 self.len == 0
48 }
49
50 fn front_idx(&self) -> u32 {
51 debug_assert!(!self.is_empty());
54 self.idx_buf[self.head as usize]
55 }
56
57 fn front_value(&self) -> u32 {
58 debug_assert!(!self.is_empty());
61 self.val_buf[self.head as usize]
62 }
63
64 fn pop_front_idx(&mut self) -> u32 {
65 debug_assert!(!self.is_empty());
66 let first = self.idx_buf[self.head as usize];
67 self.head = (self.head + 1) & RING_MASK;
68 self.len -= 1;
69 first
70 }
71
72 fn push(&mut self, state: PosState) {
73 while !self.is_empty() {
76 let slot = ((self.head + self.len - 1) & RING_MASK) as usize;
77 if self.val_buf[slot] <= state.value {
78 break;
79 }
80 self.len -= 1;
81 }
82 debug_assert!(self.len < MAX_SPARSE_GRAM_SIZE as u32);
83 let slot = ((self.head + self.len) & RING_MASK) as usize;
84 self.idx_buf[slot] = state.index;
85 self.val_buf[slot] = state.value;
86 self.len += 1;
87 }
88}
89
90#[derive(Clone)]
113pub struct QueryGrams {
114 queue: Queue,
116 content: u64,
118 content_end_idx: u32,
120 h: u32,
122}
123
124impl Eq for QueryGrams {}
125
126impl PartialEq for QueryGrams {
127 fn eq(&self, other: &Self) -> bool {
128 self.state() == other.state()
129 }
130}
131
132impl Hash for QueryGrams {
133 fn hash<H: Hasher>(&self, state: &mut H) {
134 self.state().hash(state);
135 }
136}
137
138impl Default for QueryGrams {
139 fn default() -> Self {
140 Self {
141 content: 0,
142 queue: Queue::new(),
143 content_end_idx: 0,
144 h: 0,
145 }
146 }
147}
148
149impl QueryGrams {
150 #[inline]
169 pub fn state(&self) -> (u32, u64) {
170 let begin = if self.queue.is_empty() {
171 0
172 } else {
173 self.queue.front_idx() - 2
177 };
178 let content_len = self.content_end_idx - begin;
179 debug_assert!(content_len < MAX_SPARSE_GRAM_SIZE as u32);
180 let mask = (1u64 << (content_len * 8)) - 1;
181 (content_len, self.content & mask)
182 }
183
184 pub fn min_priority(&self) -> u32 {
193 if self.queue.is_empty() {
194 u32::MAX
195 } else {
196 self.queue.front_value()
197 }
198 }
199
200 #[inline]
208 fn follow_byte(&self, end: u32) -> Option<u8> {
209 if end >= self.content_end_idx {
210 None
211 } else {
212 let shift = self.content_end_idx - 1 - end;
213 debug_assert!(shift < MAX_SPARSE_GRAM_SIZE as u32);
214 Some((self.content >> (shift * 8)) as u8)
215 }
216 }
217
218 fn extract_gram<F>(&mut self, begin_index: u32, end_index: u32, consumer: &mut F)
219 where
220 F: FnMut(NGram, u32, Option<u8>, &[u8]),
221 {
222 debug_assert!(end_index >= begin_index);
223 debug_assert!(end_index <= self.content_end_idx);
224
225 let len = (end_index - begin_index + 2) as usize;
226 let dist = self.content_end_idx - end_index;
227 let shifted = self.content >> (dist * 8);
228 let aligned = shifted << ((MAX_SPARSE_GRAM_SIZE - len) * 8);
232 let bytes = aligned.to_be_bytes();
233 let follow = self.follow_byte(end_index);
236 consumer(
241 NGram::from_window(aligned, len),
242 end_index,
243 follow,
244 &bytes[..len.min(MAX_SPARSE_GRAM_SIZE)],
245 );
246 }
247
248 pub fn append_char<F>(&mut self, c: char, consumer: F)
253 where
254 F: FnMut(NGram, u32, Option<u8>, &[u8]),
255 {
256 self.append_byte(index_fold_char(c), consumer);
257 }
258
259 pub fn append_byte<F>(&mut self, right: u8, mut consumer: F)
265 where
266 F: FnMut(NGram, u32, Option<u8>, &[u8]),
267 {
268 let left = (self.content & 0xFF) as u8;
269 self.content_end_idx += 1;
270 self.content = (self.content << 8) | right as u64;
271
272 let idx = self.content_end_idx;
273
274 if idx == 1 {
277 self.h = bigram_h(right);
278 } else {
279 let (value, h_b) = bigram_priority_rolling(left, right, self.h);
280 self.h = h_b;
281 if !self.queue.is_empty()
282 && let priority = self.queue.front_value()
283 && value < priority
284 {
285 let mut begin = self.queue.pop_front_idx();
286 while !self.queue.is_empty() && self.queue.front_value() == priority {
287 let end = self.queue.pop_front_idx();
288 self.extract_gram(begin, end, &mut consumer);
289 begin = end;
290 }
291 self.queue.clear();
292 self.queue.push(PosState { index: idx, value });
293 self.extract_gram(begin, idx, &mut consumer);
294 } else {
295 self.queue.push(PosState { index: idx, value });
296 }
297 if idx - self.queue.front_idx() + 2 >= MAX_SPARSE_GRAM_SIZE as u32 && self.queue.len > 1
298 {
299 let begin = self.queue.pop_front_idx();
300 let end = self.queue.front_idx();
301 self.extract_gram(begin, end, &mut consumer);
302 }
303 }
304 }
305
306 pub fn flush<F>(mut self, mut consumer: F)
308 where
309 F: FnMut(NGram, u32, Option<u8>, &[u8]),
310 {
311 if self.content_end_idx == 2 {
312 self.extract_gram(2, 2, &mut consumer);
313 } else {
314 while self.queue.len > 1 {
315 let begin = self.queue.pop_front_idx();
316 let end = self.queue.front_idx();
317 self.extract_gram(begin, end, &mut consumer);
318 }
319 }
320 }
321
322 pub fn consume_first<F>(&mut self, mut consumer: F)
324 where
325 F: FnMut(NGram, u32, Option<u8>, &[u8]),
326 {
327 if self.queue.len > 1 {
328 let begin = self.queue.pop_front_idx();
330 let end = self.queue.front_idx();
331 self.extract_gram(begin, end, &mut consumer);
332 } else if self.content_end_idx == 2 {
333 self.extract_gram(2, 2, &mut consumer);
335 }
336
337 if self.queue.len <= 1 {
341 self.queue.clear();
342 if self.content_end_idx > 1 {
343 self.content_end_idx = 1;
348 } else {
349 self.content_end_idx = 0;
355 }
356 }
357 }
358}
359
360#[cfg(test)]
361mod tests {
362 use super::*;
363 use crate::collect_sparse_grams_deque;
364 use std::collections::hash_map::DefaultHasher;
365 use std::ops::Range;
366
367 fn query_intervals(input: &str) -> Vec<Range<u32>> {
373 let mut q = QueryGrams::default();
374 let mut out = Vec::new();
375 for c in input.chars() {
376 q.append_char(c, |gram, end, _follow, _bytes| {
377 let begin = end + 1 - gram.len() as u32;
378 out.push(begin..end - 1);
379 });
380 }
381 q.flush(|gram, end, _follow, _bytes| {
382 let begin = end + 1 - gram.len() as u32;
383 out.push(begin..end - 1);
384 });
385 out
386 }
387
388 fn candidate_intervals(bytes: &[u8]) -> Vec<Range<u32>> {
393 let mut out = Vec::new();
394 collect_sparse_grams_deque(bytes, |gram, idx| {
395 let begin = idx + 1 - gram.len() as u32;
396 out.push(begin..idx - 1);
397 });
398 out
399 }
400
401 fn min_cover_size_query_semantics(m: u32, intervals: &[Range<u32>]) -> usize {
402 let inf = usize::MAX / 4;
403 let mut dp = vec![0usize; m as usize + 1];
404 for pos in (1..m).rev() {
405 let mut best = inf;
406 for iv in intervals {
407 if iv.start <= pos && iv.end > pos {
408 let next = iv.end.min(m) as usize;
409 best = best.min(1 + dp[next]);
410 }
411 }
412 dp[pos as usize] = best;
413 }
414 dp[1]
415 }
416
417 fn is_valid_query_cover_chain(m: u32, intervals: &[Range<u32>]) -> bool {
418 if m <= 1 {
419 return true;
420 }
421 let mut produced = 1u32;
422 for iv in intervals {
423 if iv.start > produced || iv.end <= produced {
424 return false;
425 }
426 produced = iv.end;
427 if produced >= m {
428 return true;
429 }
430 }
431 false
432 }
433
434 #[derive(Clone, Copy, Debug)]
435 struct Rng64(u64);
436
437 impl Rng64 {
438 fn new(seed: u64) -> Self {
439 Self(seed)
440 }
441
442 fn next_u64(&mut self) -> u64 {
443 let mut x = self.0;
445 x ^= x >> 12;
446 x ^= x << 25;
447 x ^= x >> 27;
448 self.0 = x;
449 x.wrapping_mul(0x2545_F491_4F6C_DD1D)
450 }
451
452 fn gen_range(&mut self, upper: usize) -> usize {
453 (self.next_u64() % upper as u64) as usize
454 }
455 }
456
457 #[test]
458 fn append_and_flush_emit_grams() {
459 let mut q = QueryGrams::default();
460 for c in "hello world".chars() {
461 q.append_char(c, |_gram, _idx, _follow, _bytes| {});
462 }
463 let mut out = Vec::new();
464 q.flush(|gram, idx, _follow, _bytes| out.push((gram, idx)));
465 assert!(!out.is_empty());
466 assert!(
467 out.iter()
468 .all(|(g, _)| (2..=MAX_SPARSE_GRAM_SIZE).contains(&g.len()))
469 );
470 }
471
472 #[test]
473 fn single_bigram_has_no_follow() {
474 let bytes: Vec<u8> = "ab".chars().map(crate::index_fold_char).collect();
477 let mut q = QueryGrams::default();
478 for &b in &bytes {
479 q.append_byte(b, |_gram, _end, _follow, _bytes| {});
480 }
481 let mut emitted = Vec::new();
482 q.flush(|gram, end, follow, _bytes| emitted.push((gram.len(), end, follow)));
483 assert_eq!(emitted, vec![(2usize, 2u32, None)]);
484 }
485
486 #[test]
487 fn emitted_follow_and_bytes_match_the_input() {
488 for input in [
492 "ab",
493 "abc",
494 "hello",
495 "querystream",
496 "aaaaaaaaaa",
497 "lardeee",
498 "sssdk",
499 ] {
500 let bytes: Vec<u8> = input.chars().map(crate::index_fold_char).collect();
501 let n = bytes.len() as u32;
502 let mut q = QueryGrams::default();
503 let mut check = |gram: NGram, end: u32, follow: Option<u8>, gram_bytes: &[u8]| {
504 assert_eq!(
505 gram_bytes,
506 &bytes[end as usize - gram.len()..end as usize],
507 "wrong gram bytes for {input:?} at end={end}"
508 );
509 assert_eq!(
510 gram,
511 NGram::from_bytes(gram_bytes),
512 "gram bytes disagree with the emitted key for {input:?} at end={end}"
513 );
514 if let Some(b) = follow {
515 assert!(
516 end < n,
517 "follow reported past input for {input:?}: end={end}"
518 );
519 assert_eq!(
520 b, bytes[end as usize],
521 "wrong follow byte for {input:?} at end={end}"
522 );
523 }
524 };
525 for &b in &bytes {
526 q.append_byte(b, &mut check);
527 }
528 q.flush(&mut check);
529 }
530 }
531
532 #[test]
533 fn state_tracks_active_window_each_step() {
534 let input = b"hello world";
543 let expected_len = [1u32, 2, 3, 2, 3, 2, 3, 2, 3, 4, 5];
545
546 let pack = |bytes: &[u8]| bytes.iter().fold(0u64, |acc, &b| (acc << 8) | b as u64);
547
548 let mut q = QueryGrams::default();
549 let mut prev = q.state();
550 assert_eq!(prev, (0, 0), "a fresh state must be empty");
551
552 for (i, &byte) in input.iter().enumerate() {
553 assert_eq!(
555 q.state(),
556 prev,
557 "state moved without an append before byte {i}"
558 );
559
560 q.append_byte(byte, |_, _, _, _| {});
561
562 let len = expected_len[i];
563 let window = &input[i + 1 - len as usize..=i];
564 let after = q.state();
565 assert_eq!(
566 after,
567 (len, pack(window)),
568 "unexpected state after byte {i} ({:?})",
569 char::from(byte),
570 );
571 prev = after;
572 }
573 }
574
575 #[test]
576 fn state_eq_and_hash_ignore_absolute_history() {
577 let mut a = QueryGrams::default();
578 let mut b = QueryGrams::default();
579
580 for c in "abc".chars() {
581 a.append_char(c, |_gram, _idx, _follow, _bytes| {});
582 }
583 for c in "zabc".chars() {
584 b.append_char(c, |_gram, _idx, _follow, _bytes| {});
585 }
586 b.consume_first(|_gram, _idx, _follow, _bytes| {});
587
588 if a == b {
590 let mut ha = DefaultHasher::new();
591 let mut hb = DefaultHasher::new();
592 a.hash(&mut ha);
593 b.hash(&mut hb);
594 assert_eq!(ha.finish(), hb.finish());
595 }
596 }
597
598 #[test]
599 fn query_flush_is_minimum_cover_on_small_inputs() {
600 for input in [
601 "abc",
602 "abcd",
603 "abcdef",
604 "hello",
605 "hello world",
606 "ababababab",
607 ] {
608 let bytes = input.as_bytes();
609 if bytes.len() < 3 {
610 continue;
611 }
612 let produced = query_intervals(input);
613 let candidates = candidate_intervals(bytes);
614 let m = bytes.len() as u32 - 1;
615
616 assert!(
617 is_valid_query_cover_chain(m, &produced),
618 "produced set is not a valid cover chain for {input:?}: {:?}",
619 produced
620 );
621 let optimum = min_cover_size_query_semantics(m, &candidates);
622 assert_eq!(
623 produced.len(),
624 optimum,
625 "produced set is not minimum-size query cover for {input:?}; produced={:?}; optimum={optimum}",
626 produced
627 );
628 }
629 }
630
631 #[test]
632 fn query_lardeee_diagnostic() {
633 let input = "lardeee";
634 let bytes = input.as_bytes();
635 let produced = query_intervals(input);
636 let candidates = candidate_intervals(bytes);
637 let m = bytes.len() as u32 - 1;
638
639 assert!(
640 is_valid_query_cover_chain(m, &produced),
641 "produced set is not a valid cover chain for {input:?}: {:?}",
642 produced
643 );
644
645 let optimum = min_cover_size_query_semantics(m, &candidates);
646
647 for iv in &produced {
650 assert!(
651 candidates
652 .iter()
653 .any(|c| c.start == iv.start && c.end == iv.end),
654 "produced interval not present in oracle candidates for {input:?}: {:?}; candidates={:?}",
655 iv,
656 candidates
657 );
658 }
659
660 assert_eq!(
662 produced.len(),
663 optimum,
664 "minimum-cover mismatch for {input:?}; produced={:?}; optimum={optimum}; candidates={:?}",
665 produced,
666 candidates
667 );
668 }
669
670 #[test]
671 fn query_sssdk_diagnostic() {
672 let input = "sssdk";
673 let bytes = input.as_bytes();
674 let produced = query_intervals(input);
675 let candidates = candidate_intervals(bytes);
676 let m = bytes.len() as u32 - 1;
677
678 assert!(
679 is_valid_query_cover_chain(m, &produced),
680 "produced set is not a valid cover chain for {input:?}: {:?}",
681 produced
682 );
683
684 let optimum = min_cover_size_query_semantics(m, &candidates);
685
686 for iv in &produced {
687 assert!(
688 candidates
689 .iter()
690 .any(|c| c.start == iv.start && c.end == iv.end),
691 "produced interval not present in oracle candidates for {input:?}: {:?}; candidates={:?}",
692 iv,
693 candidates
694 );
695 }
696
697 assert_eq!(
698 produced.len(),
699 optimum,
700 "minimum-cover mismatch for {input:?}; produced={:?}; optimum={optimum}; candidates={:?}",
701 produced,
702 candidates
703 );
704 }
705
706 #[test]
707 fn query_flush_is_minimum_cover_on_randomized_inputs() {
708 let mut rng = Rng64::new(0xA5A5_0123_89AB_CDEF);
709
710 for _ in 0..2000 {
711 let len = 3 + rng.gen_range(14); let mut bytes = vec![0u8; len];
713 for b in &mut bytes {
714 *b = b'a' + rng.gen_range(26) as u8; }
716 let input = std::str::from_utf8(&bytes).expect("ASCII should be valid UTF-8");
717
718 let produced = query_intervals(input);
719 let candidates = candidate_intervals(&bytes);
720 let m = bytes.len() as u32 - 1;
721
722 assert!(
723 is_valid_query_cover_chain(m, &produced),
724 "produced set is not a valid cover chain for randomized input {:?}: {:?}",
725 input,
726 produced
727 );
728 assert!(
729 produced.iter().all(|iv| {
730 let gram_len = (iv.end - iv.start + 2) as usize;
731 (2..=MAX_SPARSE_GRAM_SIZE).contains(&gram_len)
732 }),
733 "produced set contains out-of-range gram length for randomized input {:?}: {:?}",
734 input,
735 produced
736 );
737 let optimum = min_cover_size_query_semantics(m, &candidates);
738 assert_eq!(
739 produced.len(),
740 optimum,
741 "produced set is not minimum-size query cover for randomized input {:?}; produced={:?}; optimum={optimum}",
742 input,
743 produced
744 );
745 }
746 }
747
748 #[test]
749 fn query_consume_first_on_randomized_inputs() {
750 let mut rng = Rng64::new(0xC0DE_CAFE_1234_5678);
751
752 for _ in 0..2000 {
753 let len = 4 + rng.gen_range(13); let mut bytes = vec![0u8; len];
755 for b in &mut bytes {
756 *b = b'a' + rng.gen_range(26) as u8;
757 }
758 let input = std::str::from_utf8(&bytes).expect("ASCII should be valid UTF-8");
759
760 let split = 2 + rng.gen_range(len - 2);
761 let (prefix, suffix) = input.split_at(split);
762 let mut q = QueryGrams::default();
763
764 let mut first_half = Vec::new();
767 for c in prefix.chars() {
768 q.append_char(c, |gram, end, _follow, _bytes| {
769 let begin = end + 1 - gram.len() as u32;
770 first_half.push(begin..end - 1);
771 });
772 }
773
774 while !q.queue.is_empty() {
775 q.consume_first(|gram, end, _follow, _bytes| {
776 let begin = end + 1 - gram.len() as u32;
777 first_half.push(begin..end - 1);
778 });
779 }
780
781 assert!(first_half.iter().all(|iv| {
782 let gram_len = (iv.end - iv.start + 2) as usize;
783 (2..=MAX_SPARSE_GRAM_SIZE).contains(&gram_len)
784 }));
785
786 let prefix_bytes = prefix.as_bytes();
790 if prefix_bytes.len() >= 3 {
791 let prefix_m = prefix_bytes.len() as u32 - 1;
792 let prefix_candidates = candidate_intervals(prefix_bytes);
793 assert!(
794 is_valid_query_cover_chain(prefix_m, &first_half),
795 "first half is not a valid cover chain for prefix {:?}: {:?}",
796 prefix,
797 first_half
798 );
799 let prefix_optimum = min_cover_size_query_semantics(prefix_m, &prefix_candidates);
800 assert_eq!(
801 first_half.len(),
802 prefix_optimum,
803 "first half is not a minimum-size query cover for prefix {:?}; produced={:?}; optimum={prefix_optimum}",
804 prefix,
805 first_half
806 );
807 }
808
809 let retained = char::from((q.content & 0xFF) as u8);
810 let local_input = format!("{retained}{suffix}");
811
812 let mut remaining = Vec::new();
813 for c in suffix.chars() {
814 q.append_char(c, |gram, end, _follow, _bytes| {
815 let begin = end + 1 - gram.len() as u32;
816 remaining.push(begin..end - 1);
817 });
818 }
819 q.flush(|gram, end, _follow, _bytes| {
820 let begin = end + 1 - gram.len() as u32;
821 remaining.push(begin..end - 1);
822 });
823
824 assert_eq!(
825 remaining,
826 query_intervals(&local_input),
827 "post-drain continuation mismatch for randomized input {:?}; prefix={:?}; suffix={:?}; first_half={:?}; remaining={:?}; local_input={:?}",
828 input,
829 prefix,
830 suffix,
831 first_half,
832 remaining,
833 local_input
834 );
835
836 let local_bytes = local_input.as_bytes();
838 if local_bytes.len() >= 3 {
839 let local_m = local_bytes.len() as u32 - 1;
840 let local_candidates = candidate_intervals(local_bytes);
841 assert!(
842 is_valid_query_cover_chain(local_m, &remaining),
843 "second half is not a valid cover chain for {:?}: {:?}",
844 local_input,
845 remaining
846 );
847 let local_optimum = min_cover_size_query_semantics(local_m, &local_candidates);
848 assert_eq!(
849 remaining.len(),
850 local_optimum,
851 "second half is not a minimum-size query cover for {:?}; produced={:?}; optimum={local_optimum}",
852 local_input,
853 remaining
854 );
855 }
856 }
857 }
858
859 #[test]
865 fn consume_first_converges_to_default() {
866 for input in ["abc", "abcdef", "hello world", "ababababab", "mississippi"] {
867 let mut q = QueryGrams::default();
868 for c in input.chars() {
869 q.append_char(c, |_gram, _idx, _follow, _bytes| {});
870 }
871 for _ in 0..(input.len() + 8) {
874 q.consume_first(|_gram, _idx, _follow, _bytes| {});
875 }
876 assert_eq!(
877 q.state(),
878 QueryGrams::default().state(),
879 "consume_first did not converge to the default state for input {input:?}",
880 );
881 }
882 }
883}