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#[derive(Debug, Clone)]
13pub struct TermSetQuery {
14 terms_map: HashMap<Field, Vec<Term>>,
15}
16
17impl TermSetQuery {
18 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 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 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 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 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 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 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}