Skip to main content

lb_tantivy/query/
regex_query.rs

1use std::clone::Clone;
2use std::sync::Arc;
3
4use tantivy_fst::Regex;
5
6use crate::error::TantivyError;
7use crate::query::{AutomatonWeight, EnableScoring, Query, Weight};
8use crate::schema::Field;
9
10/// A Regex Query matches all of the documents
11/// containing a specific term that matches
12/// a regex pattern.
13///
14/// Wildcard queries (e.g. ho*se) can be achieved
15/// by converting them to their regex counterparts.
16///
17/// ```rust
18/// use tantivy::collector::Count;
19/// use tantivy::query::RegexQuery;
20/// use tantivy::schema::{Schema, TEXT};
21/// use tantivy::{doc, Index, IndexWriter, Term};
22///
23/// # fn test() -> tantivy::Result<()> {
24/// let mut schema_builder = Schema::builder();
25/// let title = schema_builder.add_text_field("title", TEXT);
26/// let schema = schema_builder.build();
27/// let index = Index::create_in_ram(schema);
28/// {
29///     let mut index_writer: IndexWriter = index.writer(15_000_000)?;
30///     index_writer.add_document(doc!(
31///         title => "The Name of the Wind",
32///     ))?;
33///     index_writer.add_document(doc!(
34///         title => "The Diary of Muadib",
35///     ))?;
36///     index_writer.add_document(doc!(
37///         title => "A Dairy Cow",
38///     ))?;
39///     index_writer.add_document(doc!(
40///         title => "The Diary of a Young Girl",
41///     ))?;
42///     index_writer.commit()?;
43/// }
44///
45/// let reader = index.reader()?;
46/// let searcher = reader.searcher();
47///
48/// let term = Term::from_field_text(title, "Diary");
49/// let query = RegexQuery::from_pattern("d[ai]{2}ry", title)?;
50/// let count = searcher.search(&query, &Count)?;
51/// assert_eq!(count, 3);
52/// Ok(())
53/// # }
54/// # assert!(test().is_ok());
55/// ```
56#[derive(Debug, Clone)]
57pub struct RegexQuery {
58    regex: Arc<Regex>,
59    field: Field,
60}
61
62impl RegexQuery {
63    /// Creates a new RegexQuery from a given pattern
64    pub fn from_pattern(regex_pattern: &str, field: Field) -> crate::Result<Self> {
65        let regex = Regex::new(regex_pattern)
66            .map_err(|err| TantivyError::InvalidArgument(format!("RegexQueryError: {err}")))?;
67        Ok(RegexQuery::from_regex(regex, field))
68    }
69
70    /// Creates a new RegexQuery from a fully built Regex
71    pub fn from_regex<T: Into<Arc<Regex>>>(regex: T, field: Field) -> Self {
72        RegexQuery {
73            regex: regex.into(),
74            field,
75        }
76    }
77
78    fn specialized_weight(&self) -> AutomatonWeight<Regex> {
79        AutomatonWeight::new(self.field, self.regex.clone())
80    }
81}
82
83impl Query for RegexQuery {
84    fn weight(&self, _enabled_scoring: EnableScoring<'_>) -> crate::Result<Box<dyn Weight>> {
85        Ok(Box::new(self.specialized_weight()))
86    }
87}
88
89#[cfg(test)]
90mod test {
91    use std::sync::Arc;
92
93    use tantivy_fst::Regex;
94
95    use super::RegexQuery;
96    use crate::collector::TopDocs;
97    use crate::schema::{Field, Schema, TEXT};
98    use crate::{assert_nearly_equals, Index, IndexReader, IndexWriter};
99
100    fn build_test_index() -> crate::Result<(IndexReader, Field)> {
101        let mut schema_builder = Schema::builder();
102        let country_field = schema_builder.add_text_field("country", TEXT);
103        let schema = schema_builder.build();
104        let index = Index::create_in_ram(schema);
105        {
106            let mut index_writer: IndexWriter = index.writer_for_tests().unwrap();
107            index_writer.add_document(doc!(
108                country_field => "japan",
109            ))?;
110            index_writer.add_document(doc!(
111                country_field => "korea",
112            ))?;
113            index_writer.commit()?;
114        }
115        let reader = index.reader()?;
116
117        Ok((reader, country_field))
118    }
119
120    fn verify_regex_query(
121        query_matching_one: RegexQuery,
122        query_matching_zero: RegexQuery,
123        reader: IndexReader,
124    ) {
125        let searcher = reader.searcher();
126        {
127            let scored_docs = searcher
128                .search(&query_matching_one, &TopDocs::with_limit(2))
129                .unwrap();
130            assert_eq!(scored_docs.len(), 1, "Expected only 1 document");
131            let (score, _) = scored_docs[0];
132            assert_nearly_equals!(1.0, score);
133        }
134        let top_docs = searcher
135            .search(&query_matching_zero, &TopDocs::with_limit(2))
136            .unwrap();
137        assert!(top_docs.is_empty(), "Expected ZERO document");
138    }
139
140    #[test]
141    pub fn test_regex_query() -> crate::Result<()> {
142        let (reader, field) = build_test_index()?;
143
144        let matching_one = RegexQuery::from_pattern("jap[ao]n", field)?;
145        let matching_zero = RegexQuery::from_pattern("jap[A-Z]n", field)?;
146        verify_regex_query(matching_one, matching_zero, reader);
147        Ok(())
148    }
149
150    #[test]
151    pub fn test_construct_from_regex() -> crate::Result<()> {
152        let (reader, field) = build_test_index()?;
153
154        let matching_one = RegexQuery::from_regex(Regex::new("jap[ao]n").unwrap(), field);
155        let matching_zero = RegexQuery::from_regex(Regex::new("jap[A-Z]n").unwrap(), field);
156
157        verify_regex_query(matching_one, matching_zero, reader);
158        Ok(())
159    }
160
161    #[test]
162    pub fn test_construct_from_reused_regex() -> crate::Result<()> {
163        let r1 = Arc::new(Regex::new("jap[ao]n").unwrap());
164        let r2 = Arc::new(Regex::new("jap[A-Z]n").unwrap());
165
166        let (reader, field) = build_test_index()?;
167
168        let matching_one = RegexQuery::from_regex(r1.clone(), field);
169        let matching_zero = RegexQuery::from_regex(r2.clone(), field);
170
171        verify_regex_query(matching_one, matching_zero, reader.clone());
172
173        let matching_one = RegexQuery::from_regex(r1, field);
174        let matching_zero = RegexQuery::from_regex(r2, field);
175
176        verify_regex_query(matching_one, matching_zero, reader);
177        Ok(())
178    }
179
180    #[test]
181    pub fn test_pattern_error() {
182        let (_reader, field) = build_test_index().unwrap();
183
184        match RegexQuery::from_pattern(r"(foo", field) {
185            Err(crate::TantivyError::InvalidArgument(msg)) => {
186                assert!(msg.contains("error: unclosed group"))
187            }
188            res => panic!("unexpected result: {res:?}"),
189        }
190    }
191}