1use rust_stemmers::{Algorithm, Stemmer};
11use serde::{Deserialize, Serialize};
12use std::borrow::Cow;
13use std::collections::BTreeMap;
14use unicode_normalization::{UnicodeNormalization, char::is_combining_mark};
15
16#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize, PartialEq, Eq)]
19#[serde(rename_all = "snake_case")]
20pub enum Stopwords {
21 #[default]
22 None,
23 English,
24}
25
26const ENGLISH_STOPWORDS: &[&str] = &[
29 "a", "an", "and", "are", "as", "at", "be", "but", "by", "for", "if", "in", "into", "is", "it",
30 "no", "not", "of", "on", "or", "such", "that", "the", "their", "then", "there", "these",
31 "they", "this", "to", "was", "will", "with",
32];
33
34#[derive(Clone, Debug, Default, Deserialize, Serialize, PartialEq)]
35#[serde(deny_unknown_fields, default)]
36pub struct TextAnalyzer {
37 pub stem: bool,
39 pub stopwords: Stopwords,
41 pub ascii_folding: bool,
43 pub max_token_length: Option<usize>,
45}
46
47impl TextAnalyzer {
48 pub fn plain() -> Self {
50 Self::default()
51 }
52 pub fn english() -> Self {
54 Self {
55 stem: true,
56 stopwords: Stopwords::English,
57 ascii_folding: true,
58 max_token_length: Some(40),
59 }
60 }
61 pub fn validate(&self) -> Result<(), String> {
62 if self.max_token_length.is_some_and(|n| n == 0 || n > 256) {
63 return Err("analyzer max_token_length must be in 1..=256".into());
64 }
65 Ok(())
66 }
67 pub fn description(&self) -> String {
69 if *self == Self::plain() {
70 "unicode-alphanumeric-lowercase-v1".into()
71 } else if *self == Self::english() {
72 "unicode-alphanumeric-english-fold-stop-stem-40-v1".into()
73 } else {
74 serde_json::to_string(self).unwrap_or_else(|_| "custom".into())
75 }
76 }
77}
78
79pub struct Analyzer {
82 config: TextAnalyzer,
83 stemmer: Option<Stemmer>,
84}
85
86impl Clone for Analyzer {
87 fn clone(&self) -> Self {
88 Self::new(self.config.clone())
89 }
90}
91
92impl Default for Analyzer {
93 fn default() -> Self {
94 Self::new(TextAnalyzer::plain())
95 }
96}
97
98impl Analyzer {
99 pub fn new(config: TextAnalyzer) -> Self {
100 let stemmer = config.stem.then(|| Stemmer::create(Algorithm::English));
101 Self { config, stemmer }
102 }
103 pub fn analyze(&self, text: &str) -> BTreeMap<String, f32> {
106 let mut counts = BTreeMap::new();
107 for word in text
108 .split(|c: char| !c.is_alphanumeric())
109 .filter(|s| !s.is_empty())
110 {
111 let mut token: Cow<str> = if self.config.ascii_folding {
112 word.nfkd()
113 .filter(|c| !is_combining_mark(*c))
114 .flat_map(char::to_lowercase)
115 .collect::<String>()
116 .into()
117 } else {
118 word.to_lowercase().into()
119 };
120 if let Some(limit) = self.config.max_token_length
121 && token.chars().nth(limit).is_some()
122 {
123 token = token.chars().take(limit).collect::<String>().into();
124 }
125 if ENGLISH_STOPWORDS.binary_search(&token.as_ref()).is_ok()
126 && self.config.stopwords == Stopwords::English
127 {
128 continue;
129 }
130 if let Some(stemmer) = &self.stemmer {
131 token = stemmer.stem(token.as_ref()).into_owned().into();
132 }
133 *counts.entry(token.into_owned()).or_default() += 1.0;
134 }
135 counts
136 }
137}
138
139#[cfg(test)]
140mod tests {
141 use super::*;
142
143 fn keys(analyzer: &Analyzer, text: &str) -> Vec<String> {
144 analyzer.analyze(text).into_keys().collect()
145 }
146
147 #[test]
148 fn plain_reproduces_the_historical_tokenizer() {
149 let analyzer = Analyzer::default();
150 assert_eq!(
151 keys(&analyzer, "Don't STOP. café's E123 — repair!"),
152 ["café", "don", "e123", "repair", "s", "stop", "t"]
153 );
154 }
155
156 #[test]
157 fn plain_counts_repeated_terms() {
158 let analyzer = Analyzer::default();
159 let counts = analyzer.analyze("cat Cat dog cat");
160 assert_eq!(counts["cat"], 3.0);
161 assert_eq!(counts["dog"], 1.0);
162 }
163
164 #[test]
165 fn english_folds_stems_and_removes_stop_words() {
166 let analyzer = Analyzer::new(TextAnalyzer::english());
167 let counts = analyzer.analyze("The runners are running toward the CAFÉs");
168 assert!(counts.contains_key("run"));
169 assert!(!counts.contains_key("runners"));
170 assert!(counts.contains_key("cafe"));
171 for stop in ["the", "are"] {
172 assert!(!counts.contains_key(stop), "{stop} should be filtered");
173 }
174 assert!(counts.contains_key("toward"), "non-list words survive");
175 }
176
177 #[test]
178 fn english_truncates_long_tokens_at_the_limit() {
179 let analyzer = Analyzer::new(TextAnalyzer::english());
180 let long = "x".repeat(41);
181 assert_eq!(keys(&analyzer, &long), ["x".repeat(40)]);
182 }
183
184 #[test]
185 fn folding_preserves_non_ascii_that_has_no_decomposition() {
186 let analyzer = Analyzer::new(TextAnalyzer {
187 ascii_folding: true,
188 ..TextAnalyzer::plain()
189 });
190 assert_eq!(keys(&analyzer, "Ünïcödé fi 東京"), ["fi", "unicode", "東京"]);
191 }
192
193 #[test]
194 fn individual_options_compose() {
195 let stem_only = Analyzer::new(TextAnalyzer {
196 stem: true,
197 ..TextAnalyzer::plain()
198 });
199 assert_eq!(keys(&stem_only, "the running"), ["run", "the"]);
200 let stop_only = Analyzer::new(TextAnalyzer {
201 stopwords: Stopwords::English,
202 ..TextAnalyzer::plain()
203 });
204 assert_eq!(keys(&stop_only, "the running"), ["running"]);
205 }
206
207 #[test]
208 fn stopword_list_is_sorted_for_binary_search() {
209 assert!(ENGLISH_STOPWORDS.windows(2).all(|w| w[0] < w[1]));
210 }
211
212 #[test]
213 fn config_validation_and_description() {
214 assert!(TextAnalyzer::plain().validate().is_ok());
215 assert!(TextAnalyzer::english().validate().is_ok());
216 for bad in [
217 TextAnalyzer {
218 max_token_length: Some(0),
219 ..TextAnalyzer::plain()
220 },
221 TextAnalyzer {
222 max_token_length: Some(257),
223 ..TextAnalyzer::plain()
224 },
225 ] {
226 assert!(bad.validate().is_err());
227 }
228 assert_eq!(
229 TextAnalyzer::plain().description(),
230 "unicode-alphanumeric-lowercase-v1"
231 );
232 assert!(
233 TextAnalyzer::english()
234 .description()
235 .starts_with("unicode-alphanumeric-english")
236 );
237 }
238
239 #[test]
240 fn serde_defaults_to_plain_and_round_trips() {
241 let parsed: TextAnalyzer = serde_json::from_str("{}").unwrap();
242 assert_eq!(parsed, TextAnalyzer::plain());
243 let english = TextAnalyzer::english();
244 let json = serde_json::to_string(&english).unwrap();
245 assert_eq!(
246 serde_json::from_str::<TextAnalyzer>(&json).unwrap(),
247 english
248 );
249 assert!(serde_json::from_str::<TextAnalyzer>(r#"{"stem":false,"unknown":1}"#).is_err());
250 }
251}