diff --git a/common/src/key_tracking.rs b/common/src/key_tracking.rs new file mode 100644 index 000000000..d1afa5bbc --- /dev/null +++ b/common/src/key_tracking.rs @@ -0,0 +1,86 @@ +/// Initial capacity of the `Vec` key buffer, large enough for most keys. +const DEFAULT_KEY_CAPACITY: usize = 100; + +/// Buffer in which a term dictionary streamer keeps track of the current key. +/// +/// - `Vec` keeps track of the current key. +/// - [`WithoutKeys`] skips this work, for callers that only need term ordinals or values. +pub trait KeyTracking { + /// Returns the buffer a streamer starts with. + fn make_default() -> Self; + + /// Sets the current key to `key`. + fn set_key(&mut self, key: &[u8]); + + /// Sets the current key to its first `common_prefix_len` bytes, followed by `suffix`. + /// + /// The current key is not always a key of the dictionary. For instance, when positioning + /// itself on its lower bound, the sstable streamer sets the current key to a prefix of + /// the lower bound, and then applies its first entry with this method. + fn update_with_prefix(&mut self, common_prefix_len: usize, suffix: &[u8]); +} + +impl KeyTracking for Vec { + #[inline] + fn make_default() -> Self { + Vec::with_capacity(DEFAULT_KEY_CAPACITY) + } + + #[inline(always)] + fn set_key(&mut self, key: &[u8]) { + self.clear(); + self.extend_from_slice(key); + } + + #[inline(always)] + fn update_with_prefix(&mut self, common_prefix_len: usize, suffix: &[u8]) { + self.truncate(common_prefix_len); + self.extend_from_slice(suffix); + } +} + +/// [`KeyTracking`] that does not keep track of keys. +/// +/// A streamer using it only gives access to term ordinals and values. +#[derive(Default)] +pub struct WithoutKeys; + +impl KeyTracking for WithoutKeys { + #[inline] + fn make_default() -> Self { + WithoutKeys + } + + #[inline(always)] + fn set_key(&mut self, _key: &[u8]) {} + + #[inline(always)] + fn update_with_prefix(&mut self, _common_prefix_len: usize, _suffix: &[u8]) {} +} + +#[cfg(test)] +mod tests { + use super::KeyTracking; + + #[test] + fn test_vec_key_tracking() { + let mut key: Vec = Vec::make_default(); + key.update_with_prefix(0, b"abc"); + assert_eq!(key, b"abc"); + key.update_with_prefix(2, b"xy"); + assert_eq!(key, b"abxy"); + key.update_with_prefix(0, b"z"); + assert_eq!(key, b"z"); + } + + #[test] + fn test_vec_set_key() { + let mut key: Vec = Vec::make_default(); + key.set_key(b"abc"); + assert_eq!(key, b"abc"); + key.set_key(b"de"); + assert_eq!(key, b"de"); + key.set_key(b""); + assert!(key.is_empty()); + } +} diff --git a/common/src/lib.rs b/common/src/lib.rs index 4e64af11c..3f47aa8d6 100644 --- a/common/src/lib.rs +++ b/common/src/lib.rs @@ -11,6 +11,7 @@ mod datetime; pub mod file_slice; mod group_by; pub mod json_path_writer; +mod key_tracking; mod serialize; mod vint; mod writer; @@ -19,6 +20,7 @@ pub use byte_count::ByteCount; pub use datetime::{DateTime, DateTimePrecision}; pub use group_by::GroupByIteratorExtended; pub use json_path_writer::JsonPathWriter; +pub use key_tracking::{KeyTracking, WithoutKeys}; pub use ownedbytes::{OwnedBytes, StableDeref}; pub use serialize::{BinarySerializable, DeserializeFrom, FixedSize}; pub use vint::{ diff --git a/src/aggregation/agg_data.rs b/src/aggregation/agg_data.rs index 6dae08c26..20d1d42c4 100644 --- a/src/aggregation/agg_data.rs +++ b/src/aggregation/agg_data.rs @@ -1065,7 +1065,11 @@ fn for_each_matching_term_ord( crate::TantivyError::InvalidArgument(format!("Invalid regex `{}`: {}", pattern, e)) })?; // TODO: we can handle patterns like `^prefix.*` more efficiently - let mut stream = str_col.dictionary().search(re).into_stream()?; + let mut stream = str_col + .dictionary() + .search(re) + .without_keys() + .into_stream()?; while stream.advance() { cb(stream.term_ord() as u32); } diff --git a/src/index/inverted_index_reader.rs b/src/index/inverted_index_reader.rs index 16dd7c4e4..5c0f9fde6 100644 --- a/src/index/inverted_index_reader.rs +++ b/src/index/inverted_index_reader.rs @@ -298,7 +298,8 @@ impl InvertedIndexReader { A::State: Clone, { use std::ops::Bound; - let range_builder = self.termdict.search(automaton); + // Matching terms are only used for their `TermInfo`. + let range_builder = self.termdict.search(automaton).without_keys(); let range_builder = match terms.start_bound() { Bound::Included(bound) => range_builder.ge(bound.serialized_value_bytes()), Bound::Excluded(bound) => range_builder.gt(bound.serialized_value_bytes()), @@ -319,7 +320,7 @@ impl InvertedIndexReader { .into_stream_async_merging_holes(merge_holes_under_bytes) .await?; - let iter = std::iter::from_fn(move || stream.next().map(|(_k, v)| v.clone())); + let iter = std::iter::from_fn(move || stream.next_without_key().cloned()); // limit on stream is only an optimization to load less data, the stream may still return // more than limit elements. @@ -430,12 +431,15 @@ impl InvertedIndexReader { // We build things from this closure otherwise we get into lifetime issues that can only // be solved with self referential strucs. Returning an io::Result from here is a bit // more leaky abstraction-wise, but a lot better than the alternative - let mut stream = termdict.search(automaton).into_stream()?; + let mut stream = termdict.search(automaton).without_keys().into_stream()?; // we could do without an iterator, but this allows us access to coalesce which simplify // things - let posting_ranges_iter = - std::iter::from_fn(move || stream.next().map(|(_k, v)| v.postings_range.clone())); + let posting_ranges_iter = std::iter::from_fn(move || { + stream + .next_without_key() + .map(|value| value.postings_range.clone()) + }); let merged_posting_ranges_iter = posting_ranges_iter.coalesce(|range1, range2| { if range1.end + MERGE_HOLES_UNDER_BYTES >= range2.start { diff --git a/src/query/automaton_weight.rs b/src/query/automaton_weight.rs index 5f1053fb6..1efd5b7b9 100644 --- a/src/query/automaton_weight.rs +++ b/src/query/automaton_weight.rs @@ -9,7 +9,7 @@ use crate::index::SegmentReader; use crate::postings::TermInfo; use crate::query::{BitSetDocSet, ConstScorer, Explanation, Scorer, Weight}; use crate::schema::{Field, IndexRecordOption}; -use crate::termdict::{TermDictionary, TermStreamer}; +use crate::termdict::{TermDictionary, TermStreamer, WithoutKeys}; use crate::{DocId, Score, TantivyError}; /// A weight struct for Fuzzy Term and Regex Queries @@ -52,9 +52,10 @@ where fn automaton_stream<'a>( &'a self, term_dict: &'a TermDictionary, - ) -> io::Result> { + ) -> io::Result> { let automaton: &A = &self.automaton; - let mut term_stream_builder = term_dict.search(automaton); + // Matching terms are only used for their `TermInfo`. + let mut term_stream_builder = term_dict.search(automaton).without_keys(); if let Some(json_path_bytes) = &self.json_path_bytes { term_stream_builder = term_stream_builder.ge(json_path_bytes); diff --git a/src/query/range_query/range_query.rs b/src/query/range_query/range_query.rs index a597c8dca..633a40988 100644 --- a/src/query/range_query/range_query.rs +++ b/src/query/range_query/range_query.rs @@ -3,6 +3,7 @@ use std::ops::Bound; use common::bounds::{map_bound, BoundsRange}; use common::BitSet; +use tantivy_fst::automaton::AlwaysMatch; use super::range_query_fastfield::FastFieldRangeWeight; use crate::index::SegmentReader; @@ -10,7 +11,7 @@ use crate::query::explanation::does_not_match; use crate::query::range_query::is_type_valid_for_fastfield_range_query; use crate::query::{BitSetDocSet, ConstScorer, EnableScoring, Explanation, Query, Scorer, Weight}; use crate::schema::{Field, IndexRecordOption, Term, Type}; -use crate::termdict::{TermDictionary, TermStreamer}; +use crate::termdict::{TermDictionary, TermStreamer, WithoutKeys}; use crate::{DocId, Score}; /// `RangeQuery` matches all documents that have at least one term within a defined range. @@ -190,9 +191,13 @@ impl InvertedIndexRangeWeight { } } - fn term_range<'a>(&self, term_dict: &'a TermDictionary) -> io::Result> { + fn term_range<'a>( + &self, + term_dict: &'a TermDictionary, + ) -> io::Result> { use std::ops::Bound::*; - let mut term_stream_builder = term_dict.range(); + // Terms in the range are only used for their `TermInfo`. + let mut term_stream_builder = term_dict.range().without_keys(); term_stream_builder = match self.lower_bound { Included(ref term_val) => term_stream_builder.ge(term_val), Excluded(ref term_val) => term_stream_builder.gt(term_val), diff --git a/src/termdict/fst_termdict/streamer.rs b/src/termdict/fst_termdict/streamer.rs index d2e31421f..60943fde9 100644 --- a/src/termdict/fst_termdict/streamer.rs +++ b/src/termdict/fst_termdict/streamer.rs @@ -1,5 +1,7 @@ use std::io; +use std::marker::PhantomData; +use common::{KeyTracking, WithoutKeys}; use tantivy_fst::automaton::AlwaysMatch; use tantivy_fst::map::{Stream, StreamBuilder}; use tantivy_fst::{Automaton, IntoStreamer, Streamer}; @@ -10,23 +12,45 @@ use crate::termdict::TermOrdinal; /// `TermStreamerBuilder` is a helper object used to define /// a range of terms that should be streamed. -pub struct TermStreamerBuilder<'a, A = AlwaysMatch> -where A: Automaton +pub struct TermStreamerBuilder<'a, A = AlwaysMatch, K = Vec> +where + A: Automaton, + K: KeyTracking, { fst_map: &'a TermDictionary, stream_builder: StreamBuilder<'a, A>, + _key_tracking: PhantomData, } -impl<'a, A> TermStreamerBuilder<'a, A> +impl<'a, A> TermStreamerBuilder<'a, A, Vec> where A: Automaton { pub(crate) fn new(fst_map: &'a TermDictionary, stream_builder: StreamBuilder<'a, A>) -> Self { TermStreamerBuilder { fst_map, stream_builder, + _key_tracking: PhantomData, } } + /// Makes the resulting [`TermStreamer`] skip copying keys. + /// + /// Use this when only term ordinals or values are needed. + /// `.key()` and `.next()` are not available on the resulting [`TermStreamer`]. + pub fn without_keys(self) -> TermStreamerBuilder<'a, A, WithoutKeys> { + TermStreamerBuilder { + fst_map: self.fst_map, + stream_builder: self.stream_builder, + _key_tracking: PhantomData, + } + } +} + +impl<'a, A, K> TermStreamerBuilder<'a, A, K> +where + A: Automaton, + K: KeyTracking, +{ /// Limit the range to terms greater or equal to the bound pub fn ge>(mut self, bound: T) -> Self { self.stream_builder = self.stream_builder.ge(bound); @@ -59,12 +83,12 @@ where A: Automaton /// Creates the stream corresponding to the range /// of terms defined using the `TermStreamerBuilder`. - pub fn into_stream(self) -> io::Result> { + pub fn into_stream(self) -> io::Result> { Ok(TermStreamer { fst_map: self.fst_map, stream: self.stream_builder.into_stream(), term_ord: 0u64, - current_key: Vec::with_capacity(100), + current_key: K::make_default(), current_value: TermInfo::default(), }) } @@ -72,26 +96,29 @@ where A: Automaton /// `TermStreamer` acts as a cursor over a range of terms of a segment. /// Terms are guaranteed to be sorted. -pub struct TermStreamer<'a, A = AlwaysMatch> -where A: Automaton +pub struct TermStreamer<'a, A = AlwaysMatch, K = Vec> +where + A: Automaton, + K: KeyTracking, { pub(crate) fst_map: &'a TermDictionary, pub(crate) stream: Stream<'a, A>, term_ord: TermOrdinal, - current_key: Vec, + current_key: K, current_value: TermInfo, } -impl TermStreamer<'_, A> -where A: Automaton +impl TermStreamer<'_, A, K> +where + A: Automaton, + K: KeyTracking, { /// Advance position the stream on the next item. /// Before the first call to `.advance()`, the stream /// is an uninitialized state. pub fn advance(&mut self) -> bool { if let Some((term, term_ord)) = self.stream.next() { - self.current_key.clear(); - self.current_key.extend_from_slice(term); + self.current_key.set_key(term); self.term_ord = term_ord; self.current_value = self.fst_map.term_info_from_ord(term_ord); true @@ -108,6 +135,23 @@ where A: Automaton self.term_ord } + /// Accesses the current value. + /// + /// Calling `.value()` after the end of the stream will return the + /// last `.value()` encountered. + /// + /// # Panics + /// + /// Calling `.value()` before the first call to `.advance()` returns + /// `V::default()`. + pub fn value(&self) -> &TermInfo { + &self.current_value + } +} + +impl TermStreamer<'_, A, Vec> +where A: Automaton +{ /// Accesses the current key. /// /// `.key()` should return the key that was returned @@ -122,19 +166,6 @@ where A: Automaton &self.current_key } - /// Accesses the current value. - /// - /// Calling `.value()` after the end of the stream will return the - /// last `.value()` encountered. - /// - /// # Panics - /// - /// Calling `.value()` before the first call to `.advance()` returns - /// `V::default()`. - pub fn value(&self) -> &TermInfo { - &self.current_value - } - /// Return the next `(key, value)` pair. #[expect(clippy::should_implement_trait)] pub fn next(&mut self) -> Option<(&[u8], &TermInfo)> { diff --git a/src/termdict/mod.rs b/src/termdict/mod.rs index 4000b08d4..0ade14e71 100644 --- a/src/termdict/mod.rs +++ b/src/termdict/mod.rs @@ -38,6 +38,7 @@ use std::io; use common::file_slice::FileSlice; use common::BinarySerializable; +pub use common::{KeyTracking, WithoutKeys}; use tantivy_fst::Automaton; use self::termdict::{ diff --git a/src/termdict/sstable_termdict/mod.rs b/src/termdict/sstable_termdict/mod.rs index cc82eba3c..1164455b4 100644 --- a/src/termdict/sstable_termdict/mod.rs +++ b/src/termdict/sstable_termdict/mod.rs @@ -25,13 +25,14 @@ pub type TermDictionaryBuilder = sstable::Writer; /// `TermStreamer` acts as a cursor over a range of terms of a segment. /// Terms are guaranteed to be sorted. -pub type TermStreamer<'a, A = AlwaysMatch> = sstable::Streamer<'a, TermSSTable, A>; +pub type TermStreamer<'a, A = AlwaysMatch, K = Vec> = sstable::Streamer<'a, TermSSTable, A, K>; /// SSTable used to store TermInfo objects. #[derive(Clone)] pub struct TermSSTable; -pub type TermStreamerBuilder<'a, A = AlwaysMatch> = sstable::StreamerBuilder<'a, TermSSTable, A>; +pub type TermStreamerBuilder<'a, A = AlwaysMatch, K = Vec> = + sstable::StreamerBuilder<'a, TermSSTable, A, K>; impl SSTable for TermSSTable { type Value = TermInfo; diff --git a/src/termdict/tests.rs b/src/termdict/tests.rs index 71b3f1c3e..736e483de 100644 --- a/src/termdict/tests.rs +++ b/src/termdict/tests.rs @@ -429,3 +429,56 @@ fn test_automaton_search() -> crate::Result<()> { assert!(!range.advance()); Ok(()) } + +#[test] +fn test_range_without_keys() { + let term_dict = stream_range_test_dict().unwrap(); + let mut stream = term_dict + .range() + .ge([2u8]) + .lt([5u8]) + .without_keys() + .into_stream() + .unwrap(); + for term_ord in 2u64..5u64 { + assert!(stream.advance()); + assert_eq!(stream.term_ord(), term_ord); + assert_eq!(stream.value(), &make_term_info(term_ord)); + } + assert!(!stream.advance()); +} + +#[test] +fn test_search_without_keys() { + const COUNTRIES: [&str; 7] = [ + "San Marino", + "Serbia", + "Slovakia", + "Slovenia", + "Spain", + "Sweden", + "Switzerland", + ]; + let buffer: Vec = { + let mut term_dictionary_builder = TermDictionaryBuilder::create(Vec::new()).unwrap(); + for (term_ord, term) in COUNTRIES.iter().enumerate() { + term_dictionary_builder + .insert(term.as_bytes(), &make_term_info(term_ord as u64)) + .unwrap(); + } + term_dictionary_builder.finish().unwrap() + }; + let term_dict = TermDictionary::open(FileSlice::from(buffer)).unwrap(); + let regex = tantivy_fst::Regex::new("S[lw].*").unwrap(); + let mut stream = term_dict + .search(regex) + .without_keys() + .into_stream() + .unwrap(); + for term_ord in [2u64, 3, 5, 6] { + assert!(stream.advance()); + assert_eq!(stream.term_ord(), term_ord); + assert_eq!(stream.value(), &make_term_info(term_ord)); + } + assert!(!stream.advance()); +} diff --git a/sstable/benches/stream_bench.rs b/sstable/benches/stream_bench.rs index f8235ceea..4b05ad68f 100644 --- a/sstable/benches/stream_bench.rs +++ b/sstable/benches/stream_bench.rs @@ -6,7 +6,7 @@ use common::file_slice::FileSlice; use criterion::{Criterion, criterion_group, criterion_main}; use rand::rngs::StdRng; use rand::{Rng, SeedableRng}; -use tantivy_fst::Automaton; +use tantivy_fst::{Automaton, Regex}; use tantivy_sstable::{Dictionary, MonotonicU64SSTable}; const CHARSET: &[u8] = b"abcdefghij"; @@ -131,6 +131,68 @@ fn automaton_bench( count } +fn stream_bench_without_keys( + dictionary: &Dictionary, + lower: &[u8], + upper: &[u8], +) -> usize { + let mut stream = dictionary + .range() + .ge(lower) + .lt(upper) + .without_keys() + .into_stream() + .unwrap(); + let mut count = 0; + while stream.advance() { + count += 1; + } + count +} + +fn automaton_bench_without_keys( + dictionary: &Dictionary, + can_match_hint: bool, + always_match_hint: bool, +) -> usize { + let mut stream = dictionary + .search(HintedPrefixAutomaton::new( + AUTOMATON_PREFIX, + black_box(can_match_hint), + black_box(always_match_hint), + )) + .without_keys() + .into_stream() + .unwrap(); + let mut count = 0; + while stream.advance() { + count += 1; + } + count +} + +fn regex_bench(dictionary: &Dictionary, regex: &Regex) -> usize { + let mut stream = dictionary.search(regex).into_stream().unwrap(); + let mut count = 0; + while stream.advance() { + count += 1; + } + count +} + +fn regex_bench_without_keys(dictionary: &Dictionary, regex: &Regex) -> usize { + let mut stream = dictionary + .search(regex) + .without_keys() + .into_stream() + .unwrap(); + let mut count = 0; + while stream.advance() { + count += 1; + } + count +} + pub fn criterion_benchmark(c: &mut Criterion) { let dict = prepare_sstable().unwrap(); c.bench_function("short_scan_init", |b| { @@ -168,6 +230,51 @@ pub fn criterion_benchmark(c: &mut Criterion) { c.bench_function("full_scan_prefix_automaton_both_hints", |b| { b.iter(|| assert_eq!(automaton_bench(&dict, true, true), NUM_AUTOMATON_MATCHES)) }); + + c.bench_function( + "full_scan_init_and_scan_full_with_bound_without_keys", + |b| { + b.iter(|| { + assert_eq!(stream_bench_without_keys(&dict, b"", b"z"), 100_000); + }) + }, + ); + c.bench_function("full_scan_init_and_scan_full_no_bounds_without_keys", |b| { + b.iter(|| { + let mut stream = dict.range().without_keys().into_stream().unwrap(); + let mut count = 0; + while stream.advance() { + count += 1; + } + count + }) + }); + c.bench_function("full_scan_prefix_automaton_no_hints_without_keys", |b| { + b.iter(|| { + assert_eq!( + automaton_bench_without_keys(&dict, false, false), + NUM_AUTOMATON_MATCHES + ) + }) + }); + c.bench_function("full_scan_prefix_automaton_both_hints_without_keys", |b| { + b.iter(|| { + assert_eq!( + automaton_bench_without_keys(&dict, true, true), + NUM_AUTOMATON_MATCHES + ) + }) + }); + + // A regex that cannot prune any block: every term goes through the automaton. + let regex = Regex::new(".*ab.*").unwrap(); + let num_regex_matches = regex_bench(&dict, ®ex); + c.bench_function("full_scan_regex", |b| { + b.iter(|| assert_eq!(regex_bench(&dict, ®ex), num_regex_matches)) + }); + c.bench_function("full_scan_regex_without_keys", |b| { + b.iter(|| assert_eq!(regex_bench_without_keys(&dict, ®ex), num_regex_matches)) + }); } criterion_group!(benches, criterion_benchmark); diff --git a/sstable/src/delta.rs b/sstable/src/delta.rs index 986700fef..4dbb0945a 100644 --- a/sstable/src/delta.rs +++ b/sstable/src/delta.rs @@ -67,6 +67,23 @@ impl DeltaKeyComparator { } self.compare(target, common_prefix_len, suffix) } + + /// Compares the key `prefix + suffix` with `target`. + /// + /// Like a call to `compare_across_blocks` with a zero `common_prefix_len`, this resets the + /// comparator: subsequent keys can then be compared incrementally. + pub(crate) fn compare_prefix_and_suffix( + &mut self, + target: &[u8], + prefix: &[u8], + suffix: &[u8], + ) -> Ordering { + if self.compare_across_blocks(target, 0, prefix) == Ordering::Greater { + // `prefix` is already past `target`, and so is any key starting with it. + return Ordering::Greater; + } + self.compare(target, prefix.len(), suffix) + } } pub struct DeltaWriter @@ -328,4 +345,51 @@ mod tests { Ordering::Greater ); } + + #[test] + fn test_delta_key_comparator_prefix_and_suffix() { + let mut keys: Vec> = vec![Vec::new()]; + for len in 1..=3 { + let shorter_keys: Vec> = keys + .iter() + .filter(|key| key.len() == len - 1) + .cloned() + .collect(); + for shorter_key in shorter_keys { + for &b in b"ab" { + let mut key = shorter_key.clone(); + key.push(b); + keys.push(key); + } + } + } + for target in &keys { + for key in &keys { + for split in 0..=key.len() { + let mut comparator = DeltaKeyComparator::new(); + let ordering = + comparator.compare_prefix_and_suffix(target, &key[..split], &key[split..]); + assert_eq!(ordering, key.cmp(target)); + if ordering == Ordering::Greater { + // A streamer stops at the first key past its target. + continue; + } + // The comparator keeps comparing the following keys incrementally. + for next_key in keys.iter().filter(|next_key| *next_key > key) { + let mut comparator = DeltaKeyComparator::new(); + comparator.compare_prefix_and_suffix(target, &key[..split], &key[split..]); + let common_prefix_len = crate::common_prefix_len(key, next_key); + assert_eq!( + comparator.compare( + target, + common_prefix_len, + &next_key[common_prefix_len..] + ), + next_key.cmp(target) + ); + } + } + } + } + } } diff --git a/sstable/src/lib.rs b/sstable/src/lib.rs index d925ad540..f0a2ddcb2 100644 --- a/sstable/src/lib.rs +++ b/sstable/src/lib.rs @@ -50,6 +50,7 @@ pub mod value; mod index; pub use index::{BlockAddr, SSTableIndex, SSTableIndexBuilder}; pub(crate) mod vint; +pub use common::{KeyTracking, WithoutKeys}; pub use dictionary::{Dictionary, TermOrdHit}; pub use streamer::{Streamer, StreamerBuilder}; diff --git a/sstable/src/streamer.rs b/sstable/src/streamer.rs index 9e5d11541..447fc9f3c 100644 --- a/sstable/src/streamer.rs +++ b/sstable/src/streamer.rs @@ -1,7 +1,9 @@ use std::cmp::Ordering; use std::io; +use std::marker::PhantomData; use std::ops::Bound; +use common::{KeyTracking, WithoutKeys}; use tantivy_fst::Automaton; use tantivy_fst::automaton::AlwaysMatch; @@ -11,17 +13,19 @@ use crate::{DeltaReader, SSTable, TermOrdinal}; /// `StreamerBuilder` is a helper object used to define /// a range of terms that should be streamed. -pub struct StreamerBuilder<'a, TSSTable, A = AlwaysMatch> +pub struct StreamerBuilder<'a, TSSTable, A = AlwaysMatch, K = Vec> where A: Automaton, A::State: Clone, TSSTable: SSTable, + K: KeyTracking, { term_dict: &'a Dictionary, automaton: A, lower: Bound>, upper: Bound>, limit: Option, + _key_tracking: PhantomData, } fn bound_as_byte_slice(bound: &Bound>) -> Bound<&[u8]> { @@ -32,6 +36,22 @@ fn bound_as_byte_slice(bound: &Bound>) -> Bound<&[u8]> { } } +/// Same as `matches_upper_bound`, for a key given as `prefix + suffix` rather than as a delta. +fn prefix_and_suffix_match_upper_bound( + comparator: &mut DeltaKeyComparator, + upper_bound: &Bound>, + prefix: &[u8], + suffix: &[u8], +) -> bool { + let (upper_bound_key, inclusive) = match upper_bound { + Bound::Unbounded => return true, + Bound::Included(upper_bound_key) => (upper_bound_key, true), + Bound::Excluded(upper_bound_key) => (upper_bound_key, false), + }; + let ordering = comparator.compare_prefix_and_suffix(upper_bound_key, prefix, suffix); + ordering == Ordering::Less || inclusive && ordering == Ordering::Equal +} + #[inline(always)] fn matches_upper_bound( comparator: &mut DeltaKeyComparator, @@ -48,7 +68,7 @@ fn matches_upper_bound( ordering == Ordering::Less || inclusive && ordering == Ordering::Equal } -impl<'a, TSSTable, A> StreamerBuilder<'a, TSSTable, A> +impl<'a, TSSTable, A> StreamerBuilder<'a, TSSTable, A, Vec> where A: Automaton, A::State: Clone, @@ -61,9 +81,30 @@ where lower: Bound::Unbounded, upper: Bound::Unbounded, limit: None, + _key_tracking: PhantomData, } } + /// Makes the resulting [`Streamer`] skip rebuilding keys (as an optimisation). + pub fn without_keys(self) -> StreamerBuilder<'a, TSSTable, A, WithoutKeys> { + StreamerBuilder { + term_dict: self.term_dict, + automaton: self.automaton, + lower: self.lower, + upper: self.upper, + limit: self.limit, + _key_tracking: PhantomData, + } + } +} + +impl<'a, TSSTable, A, K> StreamerBuilder<'a, TSSTable, A, K> +where + A: Automaton, + A::State: Clone, + TSSTable: SSTable, + K: KeyTracking, +{ /// Limit the range to terms greater or equal to the bound pub fn ge>(mut self, bound: T) -> Self { self.lower = Bound::Included(bound.as_ref().to_owned()); @@ -128,7 +169,7 @@ where fn into_stream_given_delta_reader( self, delta_reader: DeltaReader<::ValueReader>, - ) -> io::Result> { + ) -> io::Result> { let start_state = self.automaton.start(); let start_key = bound_as_byte_slice(&self.lower); @@ -153,7 +194,7 @@ where states: vec![start_state], always_match_at, delta_reader, - key: Vec::new(), + key: K::make_default(), term_ord: first_term.checked_sub(1), lower_bound_reached: self.lower == Bound::Unbounded, lower_bound: self.lower, @@ -164,7 +205,7 @@ where } /// See `into_stream(..)` - pub async fn into_stream_async(self) -> io::Result> { + pub async fn into_stream_async(self) -> io::Result> { self.into_stream_async_merging_holes(0).await } @@ -173,14 +214,14 @@ where pub async fn into_stream_async_merging_holes( self, merge_holes_under_bytes: usize, - ) -> io::Result> { + ) -> io::Result> { let delta_reader = self.delta_reader_async(merge_holes_under_bytes).await?; self.into_stream_given_delta_reader(delta_reader) } /// Creates the stream corresponding to the range /// of terms defined using the `StreamerBuilder`. - pub fn into_stream(self) -> io::Result> { + pub fn into_stream(self) -> io::Result> { let delta_reader = self.delta_reader()?; self.into_stream_given_delta_reader(delta_reader) } @@ -188,16 +229,17 @@ where /// `Streamer` acts as a cursor over a range of terms of a segment. /// Terms are guaranteed to be sorted. -pub struct Streamer<'a, TSSTable, A = AlwaysMatch> +pub struct Streamer<'a, TSSTable, A = AlwaysMatch, K = Vec> where A: Automaton, A::State: Clone, TSSTable: SSTable, + K: KeyTracking, { automaton: A, states: Vec, delta_reader: crate::DeltaReader, - key: Vec, + key: K, term_ord: Option, lower_bound: Bound>, upper_bound: Bound>, @@ -228,11 +270,12 @@ where TSSTable: SSTable } } -impl Streamer<'_, TSSTable, A> +impl Streamer<'_, TSSTable, A, K> where A: Automaton, A::State: Clone, TSSTable: SSTable, + K: KeyTracking, { #[inline(always)] fn advance_delta_reader(&mut self) -> bool { @@ -254,8 +297,9 @@ where /// Make progress up to the lower bound /// - /// Returns whether the reader was positioned on a key matching the lower bound. - /// If false, the delta_reader has been exhausted without finding such a key. + /// Returns whether the reader was positioned on a key within both the lower and the upper + /// bound. If false, there is no such key: either the delta_reader has been exhausted without + /// reaching the lower bound, or the first key past the lower bound exceeds the upper bound. fn initialize(&mut self) -> bool { debug_assert!(!self.lower_bound_reached); let mut lower_bound_comparator = DeltaKeyComparator::new(); @@ -275,17 +319,24 @@ where let match_lower_bound = ordering == Ordering::Greater || inclusive && ordering == Ordering::Equal; if match_lower_bound { - self.key.clear(); - self.key - .extend_from_slice(&lower_bound_key[..common_prefix_len]); - self.key.extend_from_slice(suffix); + // The previous key is unknown, but the comparator guarantees it starts with + // `implied_prefix`. Replaying `implied_prefix` as the previous key turns the + // current entry into a regular delta. + let implied_prefix = &lower_bound_key[..common_prefix_len]; + self.key.set_key(implied_prefix); + self.key.update_with_prefix(common_prefix_len, suffix); let mut state: A::State = self.states.last().unwrap().clone(); - for b in &self.key { - state = self.automaton.accept(&state, *b); + for b in implied_prefix.iter().copied().chain(suffix.iter().copied()) { + state = self.automaton.accept(&state, b); self.states.push(state.clone()); } self.lower_bound_reached = true; - return true; + return prefix_and_suffix_match_upper_bound( + &mut self.upper_bound_comparator, + &self.upper_bound, + implied_prefix, + suffix, + ); } } self.lower_bound_reached = true; @@ -298,15 +349,7 @@ where pub fn advance(&mut self) -> bool { if !self.lower_bound_reached { if !self.initialize() { - // no key higher than lower-bound at all - return false; - } - if !matches_upper_bound( - &mut self.upper_bound_comparator, - &self.upper_bound, - 0, - &self.key, - ) { + // no key within the bounds at all return false; } if self.automaton.is_match(self.states.last().unwrap()) { @@ -388,8 +431,8 @@ where #[inline(always)] fn reconstruct_key_and_check_upper_bound(&mut self) -> bool { let common_prefix_len = self.delta_reader.common_prefix_len(); - self.key.truncate(common_prefix_len); - self.key.extend_from_slice(self.delta_reader.suffix()); + self.key + .update_with_prefix(common_prefix_len, self.delta_reader.suffix()); // TODO there is an idea where we only look at the upper bound when our delta_reader // reached the last block (if we pruned blocks beforehand (do we always?) we cannot @@ -411,6 +454,26 @@ where self.term_ord.unwrap_or(0u64) } + /// Accesses the current value. + /// + /// Calling `.value()` after the end of the stream will return the + /// last `.value()` encountered. + /// + /// # Panics + /// + /// Calling `.value()` before the first call to `.advance()` returns + /// `V::default()`. + pub fn value(&self) -> &TSSTable::Value { + self.delta_reader.value() + } +} + +impl Streamer<'_, TSSTable, A, Vec> +where + A: Automaton, + A::State: Clone, + TSSTable: SSTable, +{ /// Accesses the current key. /// /// `.key()` should return the key that was returned @@ -425,21 +488,9 @@ where &self.key } - /// Accesses the current value. - /// - /// Calling `.value()` after the end of the stream will return the - /// last `.value()` encountered. - /// - /// # Panics - /// - /// Calling `.value()` before the first call to `.advance()` returns - /// `V::default()`. - pub fn value(&self) -> &TSSTable::Value { - self.delta_reader.value() - } - /// Return the next `(key, value)` pair. #[expect(clippy::should_implement_trait)] + #[inline(always)] pub fn next(&mut self) -> Option<(&[u8], &TSSTable::Value)> { if self.advance() { Some((self.key(), self.value())) @@ -449,6 +500,23 @@ where } } +impl Streamer<'_, TSSTable, A, WithoutKeys> +where + A: Automaton, + A::State: Clone, + TSSTable: SSTable, +{ + /// Return the next `(key, value)` pair. + #[inline(always)] + pub fn next_without_key(&mut self) -> Option<&TSSTable::Value> { + if self.advance() { + Some(self.value()) + } else { + None + } + } +} + #[cfg(test)] mod tests { use std::io; @@ -503,6 +571,61 @@ mod tests { Ok(()) } + #[test] + fn test_sstable_search_without_keys() { + let term_dict = create_test_dictionary().unwrap(); + let ptn = tantivy_fst::Regex::new("ab.*t.*").unwrap(); + let mut term_streamer = term_dict.search(ptn).without_keys().into_stream().unwrap(); + assert!(term_streamer.advance()); + assert_eq!(term_streamer.term_ord(), 1); + assert_eq!(term_streamer.value(), &1u64); + assert!(term_streamer.advance()); + assert_eq!(term_streamer.term_ord(), 2); + assert_eq!(term_streamer.value(), &2u64); + assert!(!term_streamer.advance()); + } + + #[test] + fn test_sstable_range_first_key_past_upper_bound() { + let term_dict = create_test_dictionary().unwrap(); + // "abalation" is the first key past the lower bound, and it already exceeds the upper + // bound. + let mut keyed_stream = term_dict + .range() + .gt("abaisance") + .lt("abal") + .into_stream() + .unwrap(); + assert!(!keyed_stream.advance()); + let mut keyless_stream = term_dict + .range() + .gt("abaisance") + .lt("abal") + .without_keys() + .into_stream() + .unwrap(); + assert!(!keyless_stream.advance()); + } + + #[test] + fn test_sstable_range_without_keys() { + let term_dict = create_test_dictionary().unwrap(); + let mut term_streamer = term_dict + .range() + .ge("abal") + .le("abalienate") + .without_keys() + .into_stream() + .unwrap(); + assert!(term_streamer.advance()); + assert_eq!(term_streamer.term_ord(), 1); + assert_eq!(term_streamer.value(), &1u64); + assert!(term_streamer.advance()); + assert_eq!(term_streamer.term_ord(), 2); + assert_eq!(term_streamer.value(), &2u64); + assert!(!term_streamer.advance()); + } + // TODO add test for sparse search with a block of poison (starts with 0xffffffff) => such a // block instantly causes an unexpected EOF error }