Skip to main content

lb_tantivy/query/
all_query.rs

1use crate::docset::{DocSet, COLLECT_BLOCK_BUFFER_LEN, TERMINATED};
2use crate::index::SegmentReader;
3use crate::query::boost_query::BoostScorer;
4use crate::query::explanation::does_not_match;
5use crate::query::{EnableScoring, Explanation, Query, Scorer, Weight};
6use crate::{DocId, Score};
7
8/// Query that matches all of the documents.
9///
10/// All of the documents get the score 1.0.
11#[derive(Clone, Debug)]
12pub struct AllQuery;
13
14impl Query for AllQuery {
15    fn weight(&self, _: EnableScoring<'_>) -> crate::Result<Box<dyn Weight>> {
16        Ok(Box::new(AllWeight))
17    }
18}
19
20/// Weight associated with the `AllQuery` query.
21pub struct AllWeight;
22
23impl Weight for AllWeight {
24    fn scorer(&self, reader: &SegmentReader, boost: Score) -> crate::Result<Box<dyn Scorer>> {
25        let all_scorer = AllScorer::new(reader.max_doc());
26        Ok(Box::new(BoostScorer::new(all_scorer, boost)))
27    }
28
29    fn explain(&self, reader: &SegmentReader, doc: DocId) -> crate::Result<Explanation> {
30        if doc >= reader.max_doc() {
31            return Err(does_not_match(doc));
32        }
33        Ok(Explanation::new("AllQuery", 1.0))
34    }
35}
36
37/// Scorer associated with the `AllQuery` query.
38pub struct AllScorer {
39    doc: DocId,
40    max_doc: DocId,
41}
42
43impl AllScorer {
44    /// Creates a new AllScorer with `max_doc` docs.
45    pub fn new(max_doc: DocId) -> AllScorer {
46        AllScorer { doc: 0u32, max_doc }
47    }
48}
49
50impl DocSet for AllScorer {
51    #[inline(always)]
52    fn advance(&mut self) -> DocId {
53        if self.doc + 1 >= self.max_doc {
54            self.doc = TERMINATED;
55            return TERMINATED;
56        }
57        self.doc += 1;
58        self.doc
59    }
60
61    fn fill_buffer(&mut self, buffer: &mut [DocId; COLLECT_BLOCK_BUFFER_LEN]) -> usize {
62        if self.doc() == TERMINATED {
63            return 0;
64        }
65        let is_safe_distance = self.doc() + (buffer.len() as u32) < self.max_doc;
66        if is_safe_distance {
67            let num_items = buffer.len();
68            for buffer_val in buffer {
69                *buffer_val = self.doc();
70                self.doc += 1;
71            }
72            num_items
73        } else {
74            for (i, buffer_val) in buffer.iter_mut().enumerate() {
75                *buffer_val = self.doc();
76                if self.advance() == TERMINATED {
77                    return i + 1;
78                }
79            }
80            buffer.len()
81        }
82    }
83
84    #[inline(always)]
85    fn doc(&self) -> DocId {
86        self.doc
87    }
88
89    fn size_hint(&self) -> u32 {
90        self.max_doc
91    }
92}
93
94impl Scorer for AllScorer {
95    fn score(&mut self) -> Score {
96        1.0
97    }
98}
99
100#[cfg(test)]
101mod tests {
102    use super::AllQuery;
103    use crate::docset::{DocSet, COLLECT_BLOCK_BUFFER_LEN, TERMINATED};
104    use crate::query::{AllScorer, EnableScoring, Query};
105    use crate::schema::{Schema, TEXT};
106    use crate::{Index, IndexWriter};
107
108    fn create_test_index() -> crate::Result<Index> {
109        let mut schema_builder = Schema::builder();
110        let field = schema_builder.add_text_field("text", TEXT);
111        let schema = schema_builder.build();
112        let index = Index::create_in_ram(schema);
113        let mut index_writer: IndexWriter = index.writer_for_tests()?;
114        index_writer.add_document(doc!(field=>"aaa"))?;
115        index_writer.add_document(doc!(field=>"bbb"))?;
116        index_writer.commit()?;
117        index_writer.add_document(doc!(field=>"ccc"))?;
118        index_writer.commit()?;
119        Ok(index)
120    }
121
122    #[test]
123    fn test_all_query() -> crate::Result<()> {
124        let index = create_test_index()?;
125        let reader = index.reader()?;
126        let searcher = reader.searcher();
127        let weight = AllQuery.weight(EnableScoring::disabled_from_schema(&index.schema()))?;
128        {
129            let reader = searcher.segment_reader(0);
130            let mut scorer = weight.scorer(reader, 1.0)?;
131            assert_eq!(scorer.doc(), 0u32);
132            assert_eq!(scorer.advance(), 1u32);
133            assert_eq!(scorer.doc(), 1u32);
134            assert_eq!(scorer.advance(), TERMINATED);
135        }
136        {
137            let reader = searcher.segment_reader(1);
138            let mut scorer = weight.scorer(reader, 1.0)?;
139            assert_eq!(scorer.doc(), 0u32);
140            assert_eq!(scorer.advance(), TERMINATED);
141        }
142        Ok(())
143    }
144
145    #[test]
146    fn test_all_query_with_boost() -> crate::Result<()> {
147        let index = create_test_index()?;
148        let reader = index.reader()?;
149        let searcher = reader.searcher();
150        let weight = AllQuery.weight(EnableScoring::disabled_from_schema(searcher.schema()))?;
151        let reader = searcher.segment_reader(0);
152        {
153            let mut scorer = weight.scorer(reader, 2.0)?;
154            assert_eq!(scorer.doc(), 0u32);
155            assert_eq!(scorer.score(), 2.0);
156        }
157        {
158            let mut scorer = weight.scorer(reader, 1.5)?;
159            assert_eq!(scorer.doc(), 0u32);
160            assert_eq!(scorer.score(), 1.5);
161        }
162        Ok(())
163    }
164
165    #[test]
166    pub fn test_fill_buffer() {
167        let mut postings = AllScorer {
168            doc: 0u32,
169            max_doc: COLLECT_BLOCK_BUFFER_LEN as u32 * 2 + 9,
170        };
171        let mut buffer = [0u32; COLLECT_BLOCK_BUFFER_LEN];
172        assert_eq!(postings.fill_buffer(&mut buffer), COLLECT_BLOCK_BUFFER_LEN);
173        for i in 0u32..COLLECT_BLOCK_BUFFER_LEN as u32 {
174            assert_eq!(buffer[i as usize], i);
175        }
176        assert_eq!(postings.fill_buffer(&mut buffer), COLLECT_BLOCK_BUFFER_LEN);
177        for i in 0u32..COLLECT_BLOCK_BUFFER_LEN as u32 {
178            assert_eq!(buffer[i as usize], i + COLLECT_BLOCK_BUFFER_LEN as u32);
179        }
180        assert_eq!(postings.fill_buffer(&mut buffer), 9);
181    }
182}