Skip to main content

lb_tantivy/query/
set_query.rs

1use std::collections::HashMap;
2
3use tantivy_fst::raw::CompiledAddr;
4use tantivy_fst::{Automaton, Map};
5
6use crate::query::score_combiner::DoNothingCombiner;
7use crate::query::{AutomatonWeight, BooleanWeight, EnableScoring, Occur, Query, Weight};
8use crate::schema::{Field, Schema};
9use crate::Term;
10
11/// A Term Set Query matches all of the documents containing any of the Term provided
12#[derive(Debug, Clone)]
13pub struct TermSetQuery {
14    terms_map: HashMap<Field, Vec<Term>>,
15}
16
17impl TermSetQuery {
18    /// Create a Term Set Query
19    pub fn new<T: IntoIterator<Item = Term>>(terms: T) -> Self {
20        let mut terms_map: HashMap<_, Vec<_>> = HashMap::new();
21        for term in terms {
22            terms_map.entry(term.field()).or_default().push(term);
23        }
24
25        for terms in terms_map.values_mut() {
26            terms.sort_unstable();
27            terms.dedup();
28        }
29
30        TermSetQuery { terms_map }
31    }
32
33    fn specialized_weight(
34        &self,
35        schema: &Schema,
36    ) -> crate::Result<BooleanWeight<DoNothingCombiner>> {
37        let mut sub_queries: Vec<(_, Box<dyn Weight>)> = Vec::with_capacity(self.terms_map.len());
38
39        for (&field, sorted_terms) in self.terms_map.iter() {
40            let field_entry = schema.get_field_entry(field);
41            let field_type = field_entry.field_type();
42            if !field_type.is_indexed() {
43                let error_msg = format!("Field {:?} is not indexed.", field_entry.name());
44                return Err(crate::TantivyError::SchemaError(error_msg));
45            }
46
47            // In practice this won't fail because:
48            // - we are writing to memory, so no IoError
49            // - Terms are ordered
50            let map = Map::from_iter(
51                sorted_terms
52                    .iter()
53                    .map(|key| (key.serialized_value_bytes(), 0)),
54            )
55            .map_err(std::io::Error::other)?;
56
57            sub_queries.push((
58                Occur::Should,
59                Box::new(AutomatonWeight::new(field, SetDfaWrapper(map))),
60            ));
61        }
62
63        Ok(BooleanWeight::new(
64            sub_queries,
65            false,
66            Box::new(|| DoNothingCombiner),
67        ))
68    }
69}
70
71impl Query for TermSetQuery {
72    fn weight(&self, enable_scoring: EnableScoring<'_>) -> crate::Result<Box<dyn Weight>> {
73        Ok(Box::new(self.specialized_weight(enable_scoring.schema())?))
74    }
75
76    fn query_terms<'a>(&'a self, visitor: &mut dyn FnMut(&'a Term, bool)) {
77        for terms in self.terms_map.values() {
78            for term in terms {
79                visitor(term, false);
80            }
81        }
82    }
83}
84
85struct SetDfaWrapper(Map<Vec<u8>>);
86
87impl Automaton for SetDfaWrapper {
88    type State = Option<CompiledAddr>;
89
90    fn start(&self) -> Option<CompiledAddr> {
91        Some(self.0.as_ref().root().addr())
92    }
93
94    fn is_match(&self, state_opt: &Option<CompiledAddr>) -> bool {
95        if let Some(state) = state_opt {
96            self.0.as_ref().node(*state).is_final()
97        } else {
98            false
99        }
100    }
101
102    fn accept(&self, state_opt: &Option<CompiledAddr>, byte: u8) -> Option<CompiledAddr> {
103        let state = state_opt.as_ref()?;
104        let node = self.0.as_ref().node(*state);
105        let transition = node.find_input(byte)?;
106        Some(node.transition_addr(transition))
107    }
108
109    fn can_match(&self, state: &Self::State) -> bool {
110        state.is_some()
111    }
112}
113
114#[cfg(test)]
115mod tests {
116    use crate::collector::TopDocs;
117    use crate::query::{QueryParser, TermSetQuery};
118    use crate::schema::{Schema, TEXT};
119    use crate::{assert_nearly_equals, Index, IndexWriter, Term};
120
121    #[test]
122    pub fn test_term_set_query() -> crate::Result<()> {
123        let mut schema_builder = Schema::builder();
124        let field1 = schema_builder.add_text_field("field1", TEXT);
125        let field2 = schema_builder.add_text_field("field2", TEXT);
126        let schema = schema_builder.build();
127        let index = Index::create_in_ram(schema);
128        {
129            let mut index_writer: IndexWriter = index.writer_for_tests()?;
130            index_writer.add_document(doc!(
131                field1 => "doc1",
132                field2 => "val1",
133            ))?;
134            index_writer.add_document(doc!(
135                field1 => "doc2",
136                field2 => "val2",
137            ))?;
138            index_writer.add_document(doc!(
139                field1 => "doc3",
140                field2 => "val3",
141            ))?;
142            index_writer.add_document(doc!(
143                field1 => "val3",
144                field2 => "doc3",
145            ))?;
146            index_writer.commit()?;
147        }
148        let reader = index.reader()?;
149        let searcher = reader.searcher();
150
151        {
152            // single element
153            let terms = vec![Term::from_field_text(field1, "doc1")];
154
155            let term_set_query = TermSetQuery::new(terms);
156            let top_docs = searcher.search(&term_set_query, &TopDocs::with_limit(2))?;
157            assert_eq!(top_docs.len(), 1, "Expected 1 document");
158            let (score, _) = top_docs[0];
159            assert_nearly_equals!(1.0, score);
160        }
161
162        {
163            // single element, absent
164            let terms = vec![Term::from_field_text(field1, "doc4")];
165
166            let term_set_query = TermSetQuery::new(terms);
167            let top_docs = searcher.search(&term_set_query, &TopDocs::with_limit(1))?;
168            assert!(top_docs.is_empty(), "Expected 0 document");
169        }
170
171        {
172            // multiple elements
173            let terms = vec![
174                Term::from_field_text(field1, "doc1"),
175                Term::from_field_text(field1, "doc2"),
176            ];
177
178            let term_set_query = TermSetQuery::new(terms);
179            let top_docs = searcher.search(&term_set_query, &TopDocs::with_limit(2))?;
180            assert_eq!(top_docs.len(), 2, "Expected 2 documents");
181            for (score, _) in top_docs {
182                assert_nearly_equals!(1.0, score);
183            }
184        }
185
186        {
187            // multiple elements, mixed fields
188            let terms = vec![
189                Term::from_field_text(field1, "doc1"),
190                Term::from_field_text(field1, "doc1"),
191                Term::from_field_text(field2, "val2"),
192            ];
193
194            let term_set_query = TermSetQuery::new(terms);
195            let top_docs = searcher.search(&term_set_query, &TopDocs::with_limit(3))?;
196
197            assert_eq!(top_docs.len(), 2, "Expected 2 document");
198            for (score, _) in top_docs {
199                assert_nearly_equals!(1.0, score);
200            }
201        }
202
203        {
204            // no field crosstalk
205            let terms = vec![Term::from_field_text(field1, "doc3")];
206
207            let term_set_query = TermSetQuery::new(terms);
208            let top_docs = searcher.search(&term_set_query, &TopDocs::with_limit(3))?;
209            assert_eq!(top_docs.len(), 1, "Expected 1 document");
210
211            let terms = vec![Term::from_field_text(field2, "doc3")];
212
213            let term_set_query = TermSetQuery::new(terms);
214            let top_docs = searcher.search(&term_set_query, &TopDocs::with_limit(3))?;
215            assert_eq!(top_docs.len(), 1, "Expected 1 document");
216
217            let terms = vec![
218                Term::from_field_text(field1, "doc3"),
219                Term::from_field_text(field2, "doc3"),
220            ];
221
222            let term_set_query = TermSetQuery::new(terms);
223            let top_docs = searcher.search(&term_set_query, &TopDocs::with_limit(3))?;
224            assert_eq!(top_docs.len(), 2, "Expected 2 document");
225        }
226
227        Ok(())
228    }
229
230    #[test]
231    fn test_term_set_query_parser() -> crate::Result<()> {
232        let mut schema_builder = Schema::builder();
233        schema_builder.add_text_field("field", TEXT);
234        let schema = schema_builder.build();
235        let index = Index::create_in_ram(schema.clone());
236        let mut index_writer: IndexWriter = index.writer_for_tests()?;
237        let field = schema.get_field("field").unwrap();
238        index_writer.add_document(doc!(
239          field => "val1",
240        ))?;
241        index_writer.add_document(doc!(
242          field => "val2",
243        ))?;
244        index_writer.add_document(doc!(
245          field => "val3",
246        ))?;
247        index_writer.commit()?;
248        let reader = index.reader()?;
249        let searcher = reader.searcher();
250        let query_parser = QueryParser::for_index(&index, vec![]);
251        let query = query_parser.parse_query("field: IN [val1 val2]")?;
252        let top_docs = searcher.search(&query, &TopDocs::with_limit(3))?;
253        assert_eq!(top_docs.len(), 2);
254        Ok(())
255    }
256}