enhance: collect doc_id from posting list directly for text match (#43899)

issue: https://github.com/milvus-io/milvus/issues/43898

---------

Signed-off-by: SpadeA <tangchenjie1210@gmail.com>
This commit is contained in:
Spade A
2025-08-27 10:39:52 +08:00
committed by GitHub
parent e205c30f7d
commit 90a7e63665
7 changed files with 136 additions and 39 deletions
@@ -0,0 +1,80 @@
use tantivy::{
collector::{Collector, SegmentCollector},
schema::IndexRecordOption,
DocId, DocSet, Score, SegmentOrdinal, SegmentReader, Term, COLLECT_BLOCK_BUFFER_LEN,
TERMINATED,
};
use crate::bitset_wrapper::BitsetWrapper;
// only support for text match query.
pub(crate) struct DirectBitsetCollector {
pub(crate) bitset_wrapper: BitsetWrapper,
pub(crate) terms: Vec<Term>,
}
pub(crate) struct DirectBitsetChildCollector {}
impl Collector for DirectBitsetCollector {
type Fruit = ();
type Child = DirectBitsetChildCollector;
fn collect_segment(
&self,
weight: &dyn tantivy::query::Weight,
_segment_ord: u32,
reader: &SegmentReader,
) -> tantivy::Result<<Self::Child as SegmentCollector>::Fruit> {
let mut buffer = [0u32; 4096];
for term in self.terms.iter() {
let inv_index = reader.inverted_index(term.field())?;
if let Some(mut posting) = inv_index.read_postings(term, IndexRecordOption::Basic)? {
while posting.doc() != TERMINATED {
let mut len = 0;
while posting.doc() != TERMINATED && len < 4096 {
buffer[len] = posting.doc();
len += 1;
posting.advance();
}
self.bitset_wrapper.batch_set(&buffer[..len]);
}
}
}
Ok(())
}
fn for_segment(
&self,
_segment_local_id: SegmentOrdinal,
_segment: &SegmentReader,
) -> tantivy::Result<Self::Child> {
Ok(DirectBitsetChildCollector {})
}
fn merge_fruits(
&self,
_segment_fruits: Vec<<Self::Child as SegmentCollector>::Fruit>,
) -> tantivy::Result<Self::Fruit> {
Ok(())
}
fn requires_scoring(&self) -> bool {
false
}
}
impl SegmentCollector for DirectBitsetChildCollector {
type Fruit = ();
fn collect(&mut self, _doc: DocId, _score: Score) {
unreachable!();
}
fn collect_block(&mut self, _docs: &[DocId]) {
unreachable!();
}
fn harvest(self) -> Self::Fruit {}
}
@@ -63,7 +63,7 @@ impl IndexWriterWrapper {
#[cfg(test)]
mod tests {
use std::ffi::c_void;
use std::{collections::HashSet, ffi::c_void};
use tempfile::TempDir;
@@ -107,11 +107,11 @@ mod tests {
writer.commit().unwrap();
let reader = writer.create_reader(set_bitset).unwrap();
let mut res: Vec<u32> = vec![];
let mut res: HashSet<u32> = HashSet::new();
reader
.ngram_match_query("ic", 2, 3, &mut res as *mut _ as *mut c_void)
.unwrap();
assert_eq!(res, vec![2, 4, 5]);
assert_eq!(res, vec![2, 4, 5].into_iter().collect::<HashSet<u32>>());
}
#[test]
@@ -136,22 +136,22 @@ mod tests {
writer.commit().unwrap();
let reader = writer.create_reader(set_bitset).unwrap();
let mut res: Vec<u32> = vec![];
let mut res: HashSet<u32> = HashSet::new();
reader
.ngram_match_query("测试", 2, 3, &mut res as *mut _ as *mut c_void)
.unwrap();
assert_eq!(res, vec![0, 1, 2, 4]);
assert_eq!(res, vec![0, 1, 2, 4].into_iter().collect::<HashSet<u32>>());
let mut res: Vec<u32> = vec![];
let mut res: HashSet<u32> = HashSet::new();
reader
.ngram_match_query("m测试", 2, 3, &mut res as *mut _ as *mut c_void)
.unwrap();
assert_eq!(res, vec![0, 2]);
assert_eq!(res, vec![0, 2].into_iter().collect::<HashSet<u32>>());
let mut res: Vec<u32> = vec![];
let mut res: HashSet<u32> = HashSet::new();
reader
.ngram_match_query("需要被测试", 2, 3, &mut res as *mut _ as *mut c_void)
.unwrap();
assert_eq!(res, vec![4]);
assert_eq!(res, vec![4].into_iter().collect::<HashSet<u32>>());
}
}
@@ -577,6 +577,7 @@ impl IndexReaderWrapper {
#[cfg(test)]
mod test {
use std::{
collections::HashSet,
ffi::{c_void, CString},
sync::Arc,
};
@@ -607,7 +608,7 @@ mod test {
index_writer.commit().unwrap();
let index_shared = Arc::new(index);
let mut res: Vec<u32> = vec![];
let mut res: HashSet<u32> = HashSet::new();
let index_reader_wrapper =
IndexReaderWrapper::from_index(index_shared, set_bitset).unwrap();
index_reader_wrapper
@@ -683,13 +684,13 @@ mod test {
let arrays: Vec<*const libc::c_char> =
arrays.iter().map(|s| s.as_ptr()).collect::<Vec<_>>();
let mut res: Vec<u32> = vec![];
let mut res: HashSet<u32> = HashSet::new();
reader_wrapper
.terms_query_keyword(&arrays, &mut res as *mut _ as *mut c_void)
.unwrap();
assert_eq!(res.len(), 1000);
for i in 0..1000 {
assert_eq!(res[i], i as u32);
assert!(res.contains(&(i as u32)));
}
let arrays = (0..20000)
@@ -697,7 +698,7 @@ mod test {
.collect::<Vec<_>>();
let arrays: Vec<*const libc::c_char> =
arrays.iter().map(|s| s.as_ptr()).collect::<Vec<_>>();
let mut res: Vec<u32> = vec![];
let mut res: HashSet<u32> = HashSet::new();
reader_wrapper
.terms_query_keyword(&arrays, &mut res as *mut _ as *mut c_void)
.unwrap();
@@ -6,8 +6,12 @@ use tantivy::{
Term,
};
use crate::error::Result;
use crate::{analyzer::standard_analyzer, index_reader::IndexReaderWrapper};
use crate::{
analyzer::standard_analyzer, error::TantivyBindingError, index_reader::IndexReaderWrapper,
};
use crate::{
bitset_wrapper::BitsetWrapper, direct_bitset_collector::DirectBitsetCollector, error::Result,
};
impl IndexReaderWrapper {
// split the query string into multiple tokens using index's default tokenizer,
@@ -25,8 +29,15 @@ impl IndexReaderWrapper {
let token = token_stream.token();
terms.push(Term::from_field_text(self.field, &token.text));
}
let query = BooleanQuery::new_multiterms_query(terms);
self.search(&query, bitset)
let collector = DirectBitsetCollector {
bitset_wrapper: BitsetWrapper::new(bitset, self.set_bitset),
terms,
};
let query = BooleanQuery::new_multiterms_query(vec![]);
let searcher = self.reader.searcher();
searcher
.search(&query, &collector)
.map_err(TantivyBindingError::TantivyError)
}
// split the query string into multiple tokens using index's default tokenizer,
@@ -65,9 +76,8 @@ impl IndexReaderWrapper {
#[cfg(test)]
mod tests {
use std::ffi::c_void;
use std::{collections::HashSet, ffi::c_void};
use tantivy::query::TermQuery;
use tempfile::TempDir;
use crate::{index_writer::IndexWriterWrapper, util::set_bitset, TantivyIndexVersion};
@@ -95,19 +105,19 @@ mod tests {
let slop = 1;
let reader = writer.create_reader(set_bitset).unwrap();
let mut res: Vec<u32> = vec![];
let mut res: HashSet<u32> = HashSet::new();
reader
.phrase_match_query("网球滑雪", slop, &mut res as *mut _ as *mut c_void)
.unwrap();
assert_eq!(res, vec![0]);
assert_eq!(res, vec![0].into_iter().collect::<HashSet<u32>>());
let slop = 2;
let mut res: Vec<u32> = vec![];
let mut res: HashSet<u32> = HashSet::new();
let reader = writer.create_reader(set_bitset).unwrap();
reader
.phrase_match_query("网球滑雪", slop, &mut res as *mut _ as *mut c_void)
.unwrap();
assert_eq!(res, vec![0, 1]);
assert_eq!(res, vec![0, 1].into_iter().collect::<HashSet<u32>>());
}
#[test]
@@ -125,22 +135,17 @@ mod tests {
)
.unwrap();
for i in 0..10000 {
for i in 0..100000 {
writer.add("hello world", Some(i)).unwrap();
}
writer.commit().unwrap();
let reader = writer.create_reader(set_bitset).unwrap();
let query = TermQuery::new(
tantivy::Term::from_field_text(reader.field.clone(), "hello"),
tantivy::schema::IndexRecordOption::Basic,
);
let mut res: Vec<u32> = vec![];
let mut res: HashSet<u32> = HashSet::new();
reader
.search(&query, &mut res as *mut _ as *mut c_void)
.match_query("hello world", &mut res as *mut _ as *mut c_void)
.unwrap();
assert_eq!(res, (0..10000).collect::<Vec<u32>>());
assert_eq!(res, (0..100000).collect::<HashSet<u32>>());
}
}
@@ -2,7 +2,6 @@ use core::slice;
use std::sync::Arc;
use either::Either;
use futures::executor::block_on;
use libc::c_char;
use log::info;
use tantivy_5::schema::{
@@ -272,8 +271,15 @@ impl IndexWriterWrapperImpl {
match self.index_writer {
Either::Left(mut index_writer) => {
index_writer.commit()?;
// self.manual_merge();
block_on(index_writer.garbage_collect_files())?;
// merge all segments
let segment_ids = index_writer.index().searchable_segment_ids()?;
if segment_ids.len() > 1 {
let _ = index_writer.merge(&segment_ids).wait();
}
index_writer.garbage_collect_files().wait()?;
index_writer.wait_merging_threads()?;
}
Either::Right(single_segment_index_writer) => {
@@ -4,6 +4,7 @@ mod array;
mod bitset_wrapper;
mod data_type;
mod demo_c;
mod direct_bitset_collector;
mod docid_collector;
mod error;
mod hashmap_c;
@@ -1,10 +1,12 @@
use crate::convert_to_rust_slice;
use crate::error::Result;
use core::slice;
use std::collections::HashSet;
use std::ffi::CStr;
use std::ffi::{c_char, c_void};
use std::ops::Bound;
use tantivy::{directory::MmapDirectory, Index};
use crate::error::Result;
#[inline]
pub fn c_ptr_to_str(ptr: *const c_char) -> Result<&'static str> {
Ok(unsafe { CStr::from_ptr(ptr) }.to_str()?)
@@ -38,7 +40,9 @@ pub fn free_binding<T>(ptr: *mut c_void) {
#[cfg(test)]
pub extern "C" fn set_bitset(bitset: *mut c_void, doc_id: *const u32, len: usize) {
let bitset = unsafe { &mut *(bitset as *mut Vec<u32>) };
let bitset = unsafe { &mut *(bitset as *mut HashSet<u32>) };
let docs = unsafe { convert_to_rust_slice!(doc_id, len) };
bitset.extend_from_slice(docs);
for doc in docs {
bitset.insert(*doc);
}
}