From b7c5ccf6848413ba1616daf7a8db091e2e7e1de2 Mon Sep 17 00:00:00 2001 From: Pascal Seitz Date: Thu, 1 Oct 2026 10:12:20 +0200 Subject: [PATCH] Share postings decoding across automatons and reuse traversal states Scan the term dictionary once and identify matching consumers from the existing automaton states instead of evaluating matching terms again. Decode each matching posting list once and return per-automaton bitsets. Expose streamer states and re-export BitSet. Add regression coverage for transition counts, overlapping patterns, duplicate patterns and empty results. --- src/index/inverted_index_reader.rs | 356 ++++++++++++++++++++++++++--- src/lib.rs | 2 +- sstable/src/streamer.rs | 36 +++ 3 files changed, 365 insertions(+), 29 deletions(-) diff --git a/src/index/inverted_index_reader.rs b/src/index/inverted_index_reader.rs index 16dd7c4e4..e125eb89f 100644 --- a/src/index/inverted_index_reader.rs +++ b/src/index/inverted_index_reader.rs @@ -1,6 +1,8 @@ use std::io; use common::json_path_writer::JSON_END_OF_PATH; +#[cfg(feature = "quickwit")] +use common::BitSet; use common::{BinarySerializable, ByteCount}; #[cfg(feature = "quickwit")] use futures_util::{FutureExt, StreamExt, TryStreamExt}; @@ -15,6 +17,72 @@ use crate::postings::{BlockSegmentPostings, SegmentPostings, TermInfo}; use crate::schema::{IndexRecordOption, Term, Type}; use crate::termdict::TermDictionary; +#[cfg(feature = "quickwit")] +// This is how many bytes we can hope to receive during a TTFB from S3 (~80MiB/s, 50ms). +const MERGE_HOLES_UNDER_BYTES: usize = (80 * 1024 * 1024 * 50) / 1000; + +#[cfg(feature = "quickwit")] +#[derive(Clone, Copy)] +struct AutomatonUnion<'a, A>(&'a [A]); + +#[cfg(feature = "quickwit")] +impl Automaton for AutomatonUnion<'_, A> { + type State = Vec; + + fn start(&self) -> Self::State { + self.0.iter().map(Automaton::start).collect() + } + + fn is_match(&self, state: &Self::State) -> bool { + self.0 + .iter() + .zip(state) + .any(|(automaton, state)| automaton.is_match(state)) + } + + fn can_match(&self, state: &Self::State) -> bool { + self.0 + .iter() + .zip(state) + .any(|(automaton, state)| automaton.can_match(state)) + } + + fn will_always_match(&self, state: &Self::State) -> bool { + self.0 + .iter() + .zip(state) + .any(|(automaton, state)| automaton.will_always_match(state)) + } + + fn accept(&self, state: &Self::State, byte: u8) -> Self::State { + self.0 + .iter() + .zip(state) + .map(|(automaton, state)| automaton.accept(state, byte)) + .collect() + } +} + +#[cfg(feature = "quickwit")] +fn send_coalesced_posting_ranges( + posting_ranges: impl Iterator>, + sender: futures_channel::mpsc::UnboundedSender>, +) -> io::Result<()> { + let merged_posting_ranges = posting_ranges.coalesce(|range1, range2| { + if range1.end + MERGE_HOLES_UNDER_BYTES >= range2.start { + Ok(range1.start..range2.end) + } else { + Err((range1, range2)) + } + }); + for posting_range in merged_posting_ranges { + sender + .unbounded_send(posting_range) + .map_err(|_| io::Error::other("failed to send posting range back"))?; + } + Ok(()) +} + /// The inverted index reader is in charge of accessing /// the inverted index associated with a specific field. /// @@ -392,6 +460,22 @@ impl InvertedIndexReader { Ok(true) } + async fn download_posting_ranges( + &self, + posting_ranges: futures_channel::mpsc::UnboundedReceiver>, + ) -> io::Result { + let slices_downloaded = posting_ranges + .map(|posting_range| { + self.postings_file_slice + .read_bytes_slice_async(posting_range) + .map(|result| result.map(|_slice| ())) + }) + .buffer_unordered(5) + .try_collect::>() + .await?; + Ok(!slices_downloaded.is_empty()) + } + /// Warmup a block postings given a range of `Term`s. /// This method is for an advanced usage only. /// @@ -410,10 +494,7 @@ impl InvertedIndexReader { where A::State: Clone, { - // merge holes under 4MiB, that's how many bytes we can hope to receive during a TTFB from - // S3 (~80MiB/s, and 50ms latency) - const MERGE_HOLES_UNDER_BYTES: usize = (80 * 1024 * 1024 * 50) / 1000; - // we build a first iterator to download everything. Simply calling the function already + // We build a first iterator to download everything. Simply calling the function already // download everything we need from the sstable, but doesn't start iterating over it. let _term_info_iter = self .get_term_range_async(.., automaton.clone(), None, MERGE_HOLES_UNDER_BYTES) @@ -432,42 +513,132 @@ impl InvertedIndexReader { // more leaky abstraction-wise, but a lot better than the alternative let mut stream = termdict.search(automaton).into_stream()?; - // we could do without an iterator, but this allows us access to coalesce which simplify - // things - let posting_ranges_iter = + let posting_ranges = std::iter::from_fn(move || stream.next().map(|(_k, v)| v.postings_range.clone())); + send_coalesced_posting_ranges(posting_ranges, sender) + }; + let task_handle = executor(Box::new(cpu_bound_task)); - let merged_posting_ranges_iter = posting_ranges_iter.coalesce(|range1, range2| { - if range1.end + MERGE_HOLES_UNDER_BYTES >= range2.start { - Ok(range1.start..range2.end) - } else { - Err((range1, range2)) + let (_, postings_found) = futures_util::future::try_join( + task_handle, + self.download_posting_ranges(posting_ranges_to_load_stream), + ) + .await?; + Ok(postings_found) + } + + /// Warms and decodes the postings matching several automatons in one term dictionary scan. + /// + /// The returned bitsets are in the same order as `automatons`. A posting list matching several + /// automatons is decoded only once. Matching automatons are identified from the dictionary + /// traversal states, without evaluating each matching term again. + pub async fn warm_postings_automatons< + A: Automaton + Send + 'static, + E: FnOnce(Box io::Result<()> + Send>) -> F, + F: std::future::Future>, + >( + &self, + automatons: Vec, + max_doc: crate::DocId, + executor: E, + ) -> io::Result> + where + A::State: Clone, + { + if automatons.is_empty() { + return Ok(Vec::new()); + } + + // Load only term dictionary blocks that can match at least one automaton before traversing + // them synchronously on the executor. + let automaton_union = AutomatonUnion(&automatons); + let term_info_stream = self + .termdict + .search(automaton_union) + .into_stream_async_merging_holes(MERGE_HOLES_UNDER_BYTES) + .await?; + drop(term_info_stream); + + let (posting_range_sender, posting_range_receiver) = futures_channel::mpsc::unbounded(); + let (downloads_done_sender, downloads_done_receiver) = std::sync::mpsc::channel(); + let (bitsets_sender, bitsets_receiver) = std::sync::mpsc::sync_channel(1); + let termdict = self.termdict.clone(); + let postings_file_slice = self.postings_file_slice.clone(); + let record_option = self.record_option; + let cpu_bound_task = move || { + let automaton_union = AutomatonUnion(&automatons); + let mut stream = termdict.search(automaton_union).into_stream()?; + let mut matching_terms: Vec<(TermInfo, Vec)> = Vec::new(); + let posting_ranges = std::iter::from_fn(|| { + if !stream.advance() { + return None; } + let matching_automaton_ids: Vec = automatons + .iter() + .zip(stream.automaton_state()) + .enumerate() + .filter_map(|(automaton_id, (automaton, state))| { + automaton.is_match(state).then_some(automaton_id) + }) + .collect(); + let term_info = stream.value(); + matching_terms.push((term_info.clone(), matching_automaton_ids)); + Some(term_info.postings_range.clone()) }); + send_coalesced_posting_ranges(posting_ranges, posting_range_sender)?; - for posting_range in merged_posting_ranges_iter { - if sender.unbounded_send(posting_range).is_err() { - // this should happen only when search is cancelled - return Err(io::Error::other("failed to send posting range back")); + downloads_done_receiver + .recv() + .map_err(|_| io::Error::other("posting downloader stopped unexpectedly"))??; + + let mut bitsets: Vec = (0..automatons.len()) + .map(|_| BitSet::with_max_value(max_doc)) + .collect(); + for (term_info, matching_automaton_ids) in matching_terms { + let postings_data = postings_file_slice.slice(term_info.postings_range.clone()); + let mut block_postings = BlockSegmentPostings::open( + term_info.doc_freq, + postings_data, + record_option, + IndexRecordOption::Basic, + )?; + loop { + let docs = block_postings.docs(); + if docs.is_empty() { + break; + } + for &doc in docs { + for &automaton_id in &matching_automaton_ids { + bitsets[automaton_id].insert(doc); + } + } + block_postings.advance(); } } + bitsets_sender + .send(bitsets) + .map_err(|_| io::Error::other("failed to send automaton bitsets"))?; Ok(()) }; let task_handle = executor(Box::new(cpu_bound_task)); - let posting_downloader = posting_ranges_to_load_stream - .map(|posting_slice| { - self.postings_file_slice - .read_bytes_slice_async(posting_slice) - .map(|result| result.map(|_slice| ())) - }) - .buffer_unordered(5) - .try_collect::>(); + let posting_downloader = async move { + let result = self + .download_posting_ranges(posting_range_receiver) + .await + .map(|_| ()); + let task_result = result + .as_ref() + .map(|_| ()) + .map_err(|error| io::Error::new(error.kind(), error.to_string())); + let _ = downloads_done_sender.send(task_result); + result + }; - let (_, slices_downloaded) = - futures_util::future::try_join(task_handle, posting_downloader).await?; - - Ok(!slices_downloaded.is_empty()) + futures_util::future::try_join(task_handle, posting_downloader).await?; + bitsets_receiver + .recv() + .map_err(|_| io::Error::other("automaton bitset task stopped unexpectedly")) } /// Warmup the block postings for all terms. @@ -492,3 +663,132 @@ impl InvertedIndexReader { .unwrap_or(0u32)) } } + +#[cfg(all(test, feature = "quickwit"))] +mod tests { + use std::io; + use std::sync::atomic::{AtomicUsize, Ordering}; + use std::sync::Arc; + + use futures::channel::oneshot; + use tantivy_fst::{Automaton, Regex}; + + use super::{AutomatonUnion, MERGE_HOLES_UNDER_BYTES}; + use crate::schema::{Schema, STRING}; + use crate::{Index, IndexWriter}; + + // Count byte transitions to detect a second evaluation of matching terms. + struct CountingAutomaton { + regex: Regex, + transitions: Arc, + } + + impl Automaton for CountingAutomaton { + type State = ::State; + + fn start(&self) -> Self::State { + self.regex.start() + } + + fn is_match(&self, state: &Self::State) -> bool { + self.regex.is_match(state) + } + + fn can_match(&self, state: &Self::State) -> bool { + self.regex.can_match(state) + } + + fn will_always_match(&self, state: &Self::State) -> bool { + self.regex.will_always_match(state) + } + + fn accept(&self, state: &Self::State, byte: u8) -> Self::State { + self.transitions.fetch_add(1, Ordering::Relaxed); + self.regex.accept(state, byte) + } + } + + fn execute_on_thread( + task: Box io::Result<()> + Send>, + ) -> impl std::future::Future> { + let (sender, receiver) = oneshot::channel(); + std::thread::spawn(move || { + let _ = sender.send(task()); + }); + async move { + receiver + .await + .map_err(|_| io::Error::other("executor task panicked"))? + } + } + + #[test] + fn test_warm_postings_automatons() -> crate::Result<()> { + let mut schema_builder = Schema::builder(); + let field = schema_builder.add_text_field("field", STRING); + let index = Index::create_in_ram(schema_builder.build()); + let mut writer: IndexWriter = index.writer_for_tests()?; + writer.add_document(doc!(field => "apple", field => "banana"))?; + writer.add_document(doc!(field => "apricot"))?; + writer.add_document(doc!(field => "banana", field => "berry"))?; + writer.add_document(doc!(field => "carrot"))?; + writer.commit()?; + + let reader = index.reader()?; + let searcher = reader.searcher(); + let segment_reader = searcher.segment_reader(0); + let inverted_index = segment_reader.inverted_index(field)?; + let transitions = Arc::new(AtomicUsize::new(0)); + let automatons: Vec<_> = ["b.*", "z.*", "a.*", ".*a.*", "b.*"] + .into_iter() + .map(|pattern| CountingAutomaton { + regex: Regex::new(pattern).unwrap(), + transitions: transitions.clone(), + }) + .collect(); + + // Count only block selection and dictionary traversal. Collecting hits in the warmup + // must not add byte transitions beyond these same operations. + let search = || inverted_index.termdict.search(AutomatonUnion(&automatons)); + drop(futures::executor::block_on( + search().into_stream_async_merging_holes(MERGE_HOLES_UNDER_BYTES), + )?); + let mut stream = search().into_stream()?; + while stream.advance() {} + drop(stream); + let traversal_transitions = transitions.swap(0, Ordering::Relaxed); + assert!(traversal_transitions > 0); + + let bitsets = futures::executor::block_on(inverted_index.warm_postings_automatons( + automatons, + segment_reader.max_doc(), + execute_on_thread, + ))?; + + let docs: Vec> = bitsets + .iter() + .map(|bitset| { + (0..bitset.max_value()) + .filter(|&doc| bitset.contains(doc)) + .collect() + }) + .collect(); + assert_eq!( + docs, + vec![vec![0, 2], vec![], vec![0, 1], vec![0, 1, 2, 3], vec![0, 2]] + ); + assert_eq!(transitions.load(Ordering::Relaxed), traversal_transitions); + + for patterns in [vec![], vec![Regex::new("z.*").unwrap()]] { + let num_patterns = patterns.len(); + let bitsets = futures::executor::block_on(inverted_index.warm_postings_automatons( + patterns, + segment_reader.max_doc(), + execute_on_thread, + ))?; + assert_eq!(bitsets.len(), num_patterns); + assert!(bitsets.iter().all(|bitset| bitset.len() == 0)); + } + Ok(()) + } +} diff --git a/src/lib.rs b/src/lib.rs index 53a4fb10f..abfd3fb0c 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -170,7 +170,7 @@ mod future_result; // Re-exports pub use columnar; -pub use common::{ByteCount, DateTime}; +pub use common::{BitSet, ByteCount, DateTime}; pub use query_grammar; pub use time; diff --git a/sstable/src/streamer.rs b/sstable/src/streamer.rs index 9203b3d0a..5c4d0931c 100644 --- a/sstable/src/streamer.rs +++ b/sstable/src/streamer.rs @@ -272,6 +272,16 @@ where &self.key } + /// Accesses the automaton state after consuming the current key. + /// + /// Only valid after `.advance()` returns `true` or `.next()` returns `Some`. + /// This reuses the traversal state without evaluating the key again. + pub fn automaton_state(&self) -> &A::State { + self.states + .last() + .expect("stream is not positioned on a key") + } + /// Accesses the current value. /// /// Calling `.value()` after the end of the stream will return the @@ -301,6 +311,7 @@ mod tests { use std::io; use common::OwnedBytes; + use tantivy_fst::{Automaton, Regex}; use crate::{Dictionary, MonotonicU64SSTable}; @@ -335,6 +346,31 @@ mod tests { Ok(()) } + #[test] + fn test_stream_automaton_state() -> io::Result<()> { + let mut builder = Dictionary::::builder(Vec::new())?; + for key in ["", "a", "ab", "abc", "b"] { + builder.insert(key.as_bytes(), &0)?; + } + let dict = + Dictionary::::from_bytes(OwnedBytes::new(builder.finish()?))?; + let automaton = Regex::new("a.*|").unwrap(); + let mut stream = dict.search(&automaton).into_stream()?; + let mut keys = Vec::new(); + while stream.advance() { + let state = stream.key().iter().fold(automaton.start(), |state, &byte| { + automaton.accept(&state, byte) + }); + assert_eq!(stream.automaton_state(), &state); + keys.push(stream.key().to_vec()); + } + assert_eq!( + keys, + [b"".to_vec(), b"a".to_vec(), b"ab".to_vec(), b"abc".to_vec()] + ); + Ok(()) + } + #[test] fn test_sstable_search() -> io::Result<()> { let term_dict = create_test_dictionary()?;