hermes_core/query/
prefix.rs1use std::sync::Arc;
7
8use crate::dsl::Field;
9use crate::segment::SegmentReader;
10use crate::structures::{BlockPostingList, TERMINATED};
11use crate::{DocId, Score};
12
13use super::docset::{DocSet, SortedVecDocSet};
14use super::traits::{CountFuture, EmptyScorer, Query, Scorer, ScorerFuture};
15
16#[derive(Debug, Clone)]
18pub struct PrefixQuery {
19 pub field: Field,
20 pub prefix: Vec<u8>,
21}
22
23impl std::fmt::Display for PrefixQuery {
24 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
25 write!(
26 f,
27 "Prefix({}:\"{}*\")",
28 self.field.0,
29 String::from_utf8_lossy(&self.prefix)
30 )
31 }
32}
33
34impl PrefixQuery {
35 pub fn new(field: Field, prefix: impl Into<Vec<u8>>) -> Self {
37 Self {
38 field,
39 prefix: prefix.into(),
40 }
41 }
42
43 pub fn text(field: Field, text: &str) -> Self {
45 Self {
46 field,
47 prefix: text.to_lowercase().into_bytes(),
48 }
49 }
50}
51
52fn reject_chunked(reader: &SegmentReader, field: Field) -> crate::Result<()> {
56 if reader.is_chunked_field(field) {
57 return Err(crate::Error::Query(format!(
58 "PrefixQuery is not supported on chunked text field '{}'; use a MatchQuery or PhraseQuery",
59 reader.schema().get_field_name(field).unwrap_or("?")
60 )));
61 }
62 Ok(())
63}
64
65impl Query for PrefixQuery {
66 fn scorer<'a>(&self, reader: &'a SegmentReader, _limit: usize) -> ScorerFuture<'a> {
67 let field = self.field;
68 let prefix = self.prefix.clone();
69 Box::pin(async move {
70 reject_chunked(reader, field)?;
71 let postings = reader.get_prefix_postings(field, &prefix).await?;
72 if postings.is_empty() {
73 return Ok(Box::new(EmptyScorer) as Box<dyn Scorer>);
74 }
75 let docs = materialize_union(&postings, reader.num_docs());
76 if docs.is_empty() {
77 return Ok(Box::new(EmptyScorer) as Box<dyn Scorer>);
78 }
79 Ok(Box::new(PrefixScorer::new(docs)) as Box<dyn Scorer>)
80 })
81 }
82
83 #[cfg(feature = "sync")]
84 fn scorer_sync<'a>(
85 &self,
86 reader: &'a SegmentReader,
87 _limit: usize,
88 ) -> crate::Result<Box<dyn Scorer + 'a>> {
89 reject_chunked(reader, self.field)?;
90 let postings = reader.get_prefix_postings_sync(self.field, &self.prefix)?;
91 if postings.is_empty() {
92 return Ok(Box::new(EmptyScorer) as Box<dyn Scorer>);
93 }
94 let docs = materialize_union(&postings, reader.num_docs());
95 if docs.is_empty() {
96 return Ok(Box::new(EmptyScorer) as Box<dyn Scorer>);
97 }
98 Ok(Box::new(PrefixScorer::new(docs)) as Box<dyn Scorer>)
99 }
100
101 fn count_estimate<'a>(&self, reader: &'a SegmentReader) -> CountFuture<'a> {
102 let field = self.field;
103 let prefix = self.prefix.clone();
104 Box::pin(async move {
105 let postings = reader.get_prefix_postings(field, &prefix).await?;
106 Ok(postings
107 .iter()
108 .fold(0u32, |sum, posting| sum.saturating_add(posting.doc_count()))
109 .min(reader.num_docs()))
110 })
111 }
112
113 fn is_filter(&self) -> bool {
114 true
115 }
116
117 #[cfg(feature = "sync")]
118 fn as_doc_predicate<'a>(&self, reader: &'a SegmentReader) -> Option<super::DocPredicate<'a>> {
119 let bitset = self.as_doc_bitset(reader)?;
120 Some(Box::new(move |doc_id: DocId| bitset.contains(doc_id)))
121 }
122
123 #[cfg(feature = "sync")]
124 fn as_doc_bitset(&self, reader: &SegmentReader) -> Option<super::DocBitset> {
125 if reader.is_chunked_field(self.field) {
126 return None;
127 }
128 let postings = reader
129 .get_prefix_postings_sync(self.field, &self.prefix)
130 .ok()?;
131 let mut bitset = super::DocBitset::new(reader.num_docs());
132 for posting in &postings {
133 let mut iter = posting.iterator();
134 loop {
135 let d = iter.doc();
136 if d == TERMINATED {
137 break;
138 }
139 bitset.set(d);
140 iter.advance();
141 }
142 }
143 Some(bitset)
144 }
145}
146
147struct PrefixScorer {
151 inner: SortedVecDocSet,
152}
153
154impl PrefixScorer {
155 fn new(docs: Vec<u32>) -> Self {
156 Self {
157 inner: SortedVecDocSet::new(Arc::new(docs)),
158 }
159 }
160}
161
162impl DocSet for PrefixScorer {
163 #[inline]
164 fn doc(&self) -> DocId {
165 self.inner.doc()
166 }
167
168 #[inline]
169 fn advance(&mut self) -> DocId {
170 self.inner.advance()
171 }
172
173 fn seek(&mut self, target: DocId) -> DocId {
174 self.inner.seek(target)
175 }
176
177 fn size_hint(&self) -> u32 {
178 self.inner.size_hint()
179 }
180}
181
182impl Scorer for PrefixScorer {
183 fn score(&self) -> Score {
184 1.0
185 }
186}
187
188fn materialize_union(postings: &[BlockPostingList], num_docs: u32) -> Vec<u32> {
194 let posting_count = postings.iter().fold(0usize, |sum, posting| {
195 sum.saturating_add(posting.doc_count() as usize)
196 });
197 let posting_bytes = posting_count.saturating_mul(std::mem::size_of::<u32>());
198 let bitset_bytes = (num_docs as usize)
199 .div_ceil(64)
200 .saturating_mul(std::mem::size_of::<u64>());
201
202 if posting_bytes <= bitset_bytes {
203 let mut docs = Vec::with_capacity(posting_count);
204 for posting in postings {
205 let mut iter = posting.iterator();
206 loop {
207 let d = iter.doc();
208 if d == TERMINATED {
209 break;
210 }
211 docs.push(d);
212 iter.advance();
213 }
214 }
215 docs.sort_unstable();
216 docs.dedup();
217 return docs;
218 }
219
220 let mut bitset = super::DocBitset::new(num_docs);
221 for posting in postings {
222 let mut iter = posting.iterator();
223 loop {
224 let d = iter.doc();
225 if d == TERMINATED {
226 break;
227 }
228 bitset.set(d);
229 iter.advance();
230 }
231 }
232
233 let mut docs = Vec::with_capacity(bitset.count() as usize);
234 for (word_idx, &word) in bitset.bits.iter().enumerate() {
235 let mut remaining = word;
236 while remaining != 0 {
237 let bit = remaining.trailing_zeros() as usize;
238 docs.push((word_idx * 64 + bit) as u32);
239 remaining &= remaining - 1;
240 }
241 }
242 docs
243}
244
245#[cfg(test)]
246mod tests {
247 use super::*;
248
249 #[test]
250 fn test_materialize_union_empty() {
251 let docs = materialize_union(&[], 0);
252 assert!(docs.is_empty());
253 }
254
255 #[test]
256 fn test_materialize_union_deduplicates() {
257 let mut left = crate::structures::PostingList::new();
258 left.push(1, 1);
259 left.push(5, 1);
260 left.push(9, 1);
261 let mut right = crate::structures::PostingList::new();
262 right.push(2, 1);
263 right.push(5, 1);
264 right.push(10, 1);
265 let postings = vec![
266 BlockPostingList::from_posting_list(&left).unwrap(),
267 BlockPostingList::from_posting_list(&right).unwrap(),
268 ];
269
270 assert_eq!(materialize_union(&postings, 11), vec![1, 2, 5, 9, 10]);
271 assert_eq!(
274 materialize_union(&postings[..1], 1_000_000_000),
275 vec![1, 5, 9]
276 );
277 }
278
279 #[test]
280 fn test_prefix_scorer_basic() {
281 let mut scorer = PrefixScorer::new(vec![1, 5, 10, 20]);
282 assert_eq!(scorer.doc(), 1);
283 assert_eq!(scorer.score(), 1.0);
284 assert_eq!(scorer.advance(), 5);
285 assert_eq!(scorer.seek(10), 10);
286 assert_eq!(scorer.advance(), 20);
287 assert_eq!(scorer.advance(), TERMINATED);
288 }
289
290 #[test]
291 fn test_prefix_scorer_seek_past() {
292 let mut scorer = PrefixScorer::new(vec![1, 5, 10, 20]);
293 assert_eq!(scorer.seek(7), 10);
294 assert_eq!(scorer.seek(100), TERMINATED);
295 }
296
297 #[test]
298 fn test_prefix_query_display() {
299 let q = PrefixQuery::text(Field(0), "abc");
300 assert_eq!(format!("{}", q), "Prefix(0:\"abc*\")");
301 }
302
303 #[test]
304 fn test_prefix_query_is_filter() {
305 let q = PrefixQuery::text(Field(0), "test");
306 assert!(q.is_filter());
307 }
308}