Skip to main content

lance_index/scalar/inverted/tokenizer/
lance_tokenizer.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4use arrow_schema::{DataType, Field};
5use lance_arrow::ARROW_EXT_NAME_KEY;
6use lance_arrow::json::JSON_EXT_NAME;
7use serde_json::Value;
8use tantivy::tokenizer::{BoxTokenStream, Token, TokenStream};
9
10/// Document type for full text search.
11#[derive(Debug, Clone)]
12pub enum DocType {
13    Text,
14    Json,
15}
16
17impl AsRef<str> for DocType {
18    fn as_ref(&self) -> &str {
19        match self {
20            Self::Text => "text",
21            Self::Json => "json",
22        }
23    }
24}
25
26impl TryFrom<&Field> for DocType {
27    type Error = lance_core::Error;
28
29    fn try_from(field: &Field) -> Result<Self, Self::Error> {
30        match field.data_type() {
31            DataType::Utf8 | DataType::LargeUtf8 => Ok(Self::Text),
32            DataType::List(field) | DataType::LargeList(field)
33                if matches!(field.data_type(), DataType::Utf8 | DataType::LargeUtf8) =>
34            {
35                Ok(Self::Text)
36            }
37            DataType::LargeBinary => match field.metadata().get(ARROW_EXT_NAME_KEY) {
38                Some(name) if name.as_str() == JSON_EXT_NAME => Ok(Self::Json),
39                _ => Err(lance_core::Error::invalid_input_source(
40                    format!("field {} is not json", field.name()).into(),
41                )),
42            },
43            _ => Err(lance_core::Error::invalid_input_source(
44                format!("field {} is not json", field.name()).into(),
45            )),
46        }
47    }
48}
49
50impl DocType {
51    /// Get the length of the prefix before value.
52    ///  - JSON Token: path,type,value
53    ///  - Text Token: value
54    pub fn prefix_len(&self, token: &str) -> usize {
55        match self {
56            Self::Json => {
57                if let Some(pos) = token.find(',')
58                    && let Some(second_pos) = token[pos + 1..].find(',')
59                {
60                    return pos + second_pos + 2;
61                }
62                panic!("json token must be in format of <path>,<type>,<value>")
63            }
64            Self::Text => 0,
65        }
66    }
67}
68
69/// Lance full text search tokenizer.
70///
71/// `LanceTokenizer` defines 2 methods for tokenization, normally they are the same, but sometimes
72/// tokenizer needs different behavior for search and index. Take json document as an example:
73/// 1. Query text is a triplet <path,type,value>, something like `a.b,str,123`. We shouldn't use
74///    json in search, because it would be too complicated.
75/// 2. Document text is a json string.
76pub trait LanceTokenizer: Send + Sync + std::fmt::Debug {
77    /// Tokenize query text for search.
78    fn token_stream_for_search<'a>(&'a mut self, query_text: &'a str) -> BoxTokenStream<'a>;
79    /// Tokenize document text for index.
80    fn token_stream_for_doc<'a>(&'a mut self, text: &'a str) -> BoxTokenStream<'a>;
81    /// Clone the tokenizer.
82    fn box_clone(&self) -> Box<dyn LanceTokenizer>;
83    /// Get document type.
84    fn doc_type(&self) -> DocType;
85}
86
87impl Clone for Box<dyn LanceTokenizer> {
88    fn clone(&self) -> Self {
89        self.box_clone()
90    }
91}
92
93#[derive(Clone)]
94pub struct TextTokenizer {
95    tokenizer: tantivy::tokenizer::TextAnalyzer,
96}
97
98impl std::fmt::Debug for TextTokenizer {
99    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
100        write!(f, "TextTokenizer")
101    }
102}
103
104impl TextTokenizer {
105    pub fn new(tokenizer: tantivy::tokenizer::TextAnalyzer) -> Self {
106        Self { tokenizer }
107    }
108}
109
110impl LanceTokenizer for TextTokenizer {
111    fn token_stream_for_search<'a>(&'a mut self, query_text: &'a str) -> BoxTokenStream<'a> {
112        self.tokenizer.token_stream(query_text)
113    }
114
115    fn token_stream_for_doc<'a>(&'a mut self, text: &'a str) -> BoxTokenStream<'a> {
116        self.tokenizer.token_stream(text)
117    }
118
119    fn box_clone(&self) -> Box<dyn LanceTokenizer> {
120        Box::new(self.clone())
121    }
122
123    fn doc_type(&self) -> DocType {
124        DocType::Text
125    }
126}
127
128#[derive(Clone)]
129pub struct JsonTokenizer {
130    tokenizer: tantivy::tokenizer::TextAnalyzer,
131}
132
133impl JsonTokenizer {
134    pub fn new(tokenizer: tantivy::tokenizer::TextAnalyzer) -> Self {
135        Self { tokenizer }
136    }
137}
138
139impl std::fmt::Debug for JsonTokenizer {
140    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
141        write!(f, "JsonTokenizer")
142    }
143}
144
145impl LanceTokenizer for JsonTokenizer {
146    fn token_stream_for_search<'a>(&'a mut self, query_text: &'a str) -> BoxTokenStream<'a> {
147        let tokens = flatten_triplet(query_text, &mut self.tokenizer).unwrap();
148        BoxTokenStream::new(TTStream { tokens, index: 0 })
149    }
150
151    fn token_stream_for_doc<'a>(&'a mut self, text: &'a str) -> BoxTokenStream<'a> {
152        let value: Value = match serde_json::from_slice(text.as_bytes()) {
153            Ok(v) => v,
154            Err(e) => {
155                panic!("JSON parse error: {:?}", e);
156            }
157        };
158        let mut tokens = vec![];
159        let mut position = 0;
160        flatten_json(&value, "", &mut tokens, &mut position, &mut self.tokenizer);
161        BoxTokenStream::new(TTStream { tokens, index: 0 })
162    }
163
164    fn box_clone(&self) -> Box<dyn LanceTokenizer> {
165        Box::new(self.clone())
166    }
167
168    fn doc_type(&self) -> DocType {
169        DocType::Json
170    }
171}
172
173fn flatten_triplet(
174    text: &str,
175    tokenizer: &mut tantivy::tokenizer::TextAnalyzer,
176) -> lance_core::Result<Vec<Token>> {
177    let mut token_vec = Vec::new();
178    let mut idx = 0;
179
180    for triple in text.split(';') {
181        let parts: Vec<&str> = triple.splitn(3, ',').collect();
182        if parts.len() != 3 {
183            return Err(lance_core::Error::invalid_input_source(
184                format!("Invalid triple format: {}", triple).into(),
185            ));
186        }
187        let field = parts[0];
188        let v_type = parts[1];
189        let value = parts[2];
190
191        match v_type {
192            "number" | "bool" | "null" => {
193                let token = Token {
194                    offset_from: 0,
195                    offset_to: 0,
196                    position: idx,
197                    text: format!("{},{},{}", field, v_type, value),
198                    position_length: 1,
199                };
200                token_vec.push(token);
201                idx += 1;
202            }
203            "str" => {
204                let mut tokens = tokenizer.token_stream(value);
205                while let Some(token) = tokens.next() {
206                    token_vec.push(Token {
207                        offset_from: 0,
208                        offset_to: 0,
209                        position: idx,
210                        text: format!("{},{},{}", field, v_type, token.text),
211                        position_length: 1,
212                    });
213                    idx += 1;
214                }
215            }
216            _ => {
217                return Err(lance_core::Error::invalid_input_source(
218                    format!("Invalid triple type: {}", v_type).into(),
219                ));
220            }
221        }
222    }
223    Ok(token_vec)
224}
225
226fn flatten_json(
227    value: &Value,
228    prefix: &str,
229    out: &mut Vec<Token>,
230    position: &mut usize,
231    tokenizer: &mut tantivy::tokenizer::TextAnalyzer,
232) {
233    match value {
234        Value::Object(map) => {
235            for (k, v) in map {
236                let next_prefix = if prefix.is_empty() {
237                    k.clone()
238                } else {
239                    format!("{}.{}", prefix, k)
240                };
241                flatten_json(v, &next_prefix, out, position, tokenizer);
242            }
243        }
244        Value::Array(arr) => {
245            for v in arr.iter() {
246                flatten_json(v, prefix, out, position, tokenizer);
247            }
248        }
249        Value::String(text) => {
250            let mut tokens = tokenizer.token_stream(text);
251            while let Some(token) = tokens.next() {
252                let token = Token {
253                    offset_from: 0,
254                    offset_to: 0,
255                    position: *position,
256                    text: format!("{},{},{}", prefix, "str", token.text),
257                    position_length: 1,
258                };
259                *position += 1;
260                out.push(token);
261            }
262        }
263        _ => {
264            let value_type = match value {
265                Value::Null => "null",
266                Value::Bool(_) => "bool",
267                Value::Number(_) => "number",
268                _ => unreachable!(),
269            };
270            let token = Token {
271                offset_from: 0,
272                offset_to: 0,
273                position: *position,
274                text: format!("{},{},{}", prefix, value_type, value),
275                position_length: 1,
276            };
277            *position += 1;
278            out.push(token);
279        }
280    }
281}
282
283struct TTStream {
284    tokens: Vec<Token>,
285    index: usize,
286}
287
288impl TokenStream for TTStream {
289    fn advance(&mut self) -> bool {
290        if self.index < self.tokens.len() {
291            self.index += 1;
292            true
293        } else {
294            false
295        }
296    }
297
298    fn token(&self) -> &Token {
299        &self.tokens[self.index - 1]
300    }
301
302    fn token_mut(&mut self) -> &mut Token {
303        &mut self.tokens[self.index - 1]
304    }
305}
306
307#[cfg(test)]
308mod tests {
309    use crate::scalar::inverted::tokenizer::lance_tokenizer::{
310        JsonTokenizer, LanceTokenizer, flatten_json, flatten_triplet,
311    };
312    use serde_json::Value;
313    use tantivy::tokenizer::{SimpleTokenizer, Token};
314
315    #[test]
316    fn test_json_tokenizer() {
317        let text = r#"{
318          "a": 1,
319          "b": [
320            {"c": "d"},
321            {"c": "e"}
322          ]
323        }"#;
324        let mut tokenizer = JsonTokenizer::new(
325            tantivy::tokenizer::TextAnalyzer::builder(SimpleTokenizer::default()).build(),
326        );
327        let mut stream = tokenizer.token_stream_for_doc(text);
328
329        let mut tokens: Vec<Token> = vec![];
330        while let Some(token) = stream.next() {
331            tokens.push(token.clone());
332        }
333
334        assert_eq!(tokens.len(), 3);
335        assert_token(&tokens[0], 0, "a,number,1");
336        assert_token(&tokens[1], 1, "b.c,str,d");
337        assert_token(&tokens[2], 2, "b.c,str,e");
338    }
339
340    #[test]
341    fn test_flatten_json_text() {
342        let json = r#"{
343              "a": 1,
344              "b": [
345                {"c": "hello world"},
346                {"c": "e"}
347              ],
348              "c": true,
349              "d": null,
350              "e": {
351                "f": 1.0
352              }
353          }"#;
354        let value: Value = serde_json::from_str(json).unwrap();
355
356        let mut tokens = vec![];
357        let mut tokenizer =
358            tantivy::tokenizer::TextAnalyzer::builder(SimpleTokenizer::default()).build();
359        let mut position = 0;
360        flatten_json(&value, "", &mut tokens, &mut position, &mut tokenizer);
361
362        assert_eq!(7, tokens.len());
363        assert_token(&tokens[0], 0, "a,number,1");
364        assert_token(&tokens[1], 1, "b.c,str,hello");
365        assert_token(&tokens[2], 2, "b.c,str,world");
366        assert_token(&tokens[3], 3, "b.c,str,e");
367        assert_token(&tokens[4], 4, "c,bool,true");
368        assert_token(&tokens[5], 5, "d,null,null");
369        assert_token(&tokens[6], 6, "e.f,number,1.0");
370    }
371
372    #[test]
373    fn test_flatten_triplet() {
374        let text = r#"a,number,1;b.c,str,d;b.c,str,e;d,str,hello world;e,number,1.0"#;
375        let mut tokenizer =
376            tantivy::tokenizer::TextAnalyzer::builder(SimpleTokenizer::default()).build();
377        let tokens = flatten_triplet(text, &mut tokenizer).unwrap();
378
379        assert_eq!(tokens.len(), 6);
380        assert_token(&tokens[0], 0, "a,number,1");
381        assert_token(&tokens[1], 1, "b.c,str,d");
382        assert_token(&tokens[2], 2, "b.c,str,e");
383        assert_token(&tokens[3], 3, "d,str,hello");
384        assert_token(&tokens[4], 4, "d,str,world");
385        assert_token(&tokens[5], 5, "e,number,1.0");
386    }
387
388    fn assert_token(token: &Token, position: usize, text: &str) {
389        assert_eq!(
390            token.position, position,
391            "expected position {position} but {token:?}"
392        );
393        assert_eq!(
394            token.text.as_str(),
395            text,
396            "expected text {text} but {token:?}"
397        );
398    }
399}