1use std::borrow::Cow;
13use std::fmt;
14use std::str::FromStr;
15use std::sync::Arc;
16
17use crate::bm25::SparseEmbedding;
18use ordered_float::NotNan;
19use crate::segment::data_types::index::{Language, StemmingAlgorithm, StopwordsInterface, TokenizerType};
20use crate::segment::index::field_index::full_text_index::stop_words::StopwordsFilter;
21use crate::segment::index::field_index::full_text_index::tokenizers::{
22 Stemmer, Tokenizer, TokensProcessor,
23};
24use crate::sparse::common::sparse_vector::SparseVector;
25
26const DEFAULT_LANGUAGE: Language = Language::English;
27
28#[derive(Debug, Clone, PartialEq)]
30pub enum EdgeBm25Error {
31 Bm25(crate::bm25::Bm25Error),
33 UnsupportedLanguage(String),
35}
36
37impl fmt::Display for EdgeBm25Error {
38 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
39 match self {
40 Self::Bm25(e) => write!(f, "{e}"),
41 Self::UnsupportedLanguage(lang) => write!(f, "unsupported language: {lang:?}"),
42 }
43 }
44}
45
46impl std::error::Error for EdgeBm25Error {}
47
48impl From<crate::bm25::Bm25Error> for EdgeBm25Error {
49 fn from(e: crate::bm25::Bm25Error) -> Self {
50 Self::Bm25(e)
51 }
52}
53
54#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
61pub struct EdgeBm25Config {
62 #[serde(default = "default_k")]
65 pub k: NotNan<f64>,
66 #[serde(default = "default_b")]
68 pub b: NotNan<f64>,
69 #[serde(default = "default_avg_len")]
71 pub avg_len: NotNan<f64>,
72 #[serde(default)]
74 pub tokenizer: TokenizerType,
75 #[serde(default, skip_serializing_if = "Option::is_none")]
77 pub language: Option<String>,
78 #[serde(default, skip_serializing_if = "Option::is_none")]
80 pub lowercase: Option<bool>,
81 #[serde(default, skip_serializing_if = "Option::is_none")]
83 pub ascii_folding: Option<bool>,
84 #[serde(default, skip_serializing_if = "Option::is_none")]
86 pub stopwords: Option<StopwordsInterface>,
87 #[serde(default, skip_serializing_if = "Option::is_none")]
89 pub stemmer: Option<StemmingAlgorithm>,
90 #[serde(default, skip_serializing_if = "Option::is_none")]
92 pub min_token_len: Option<usize>,
93 #[serde(default, skip_serializing_if = "Option::is_none")]
95 pub max_token_len: Option<usize>,
96}
97
98const fn default_k() -> NotNan<f64> {
99 unsafe { NotNan::new_unchecked(1.2) }
101}
102
103const fn default_b() -> NotNan<f64> {
104 unsafe { NotNan::new_unchecked(0.75) }
106}
107
108const fn default_avg_len() -> NotNan<f64> {
109 unsafe { NotNan::new_unchecked(256.0) }
111}
112
113impl Default for EdgeBm25Config {
114 fn default() -> Self {
115 Self {
116 k: default_k(),
117 b: default_b(),
118 avg_len: default_avg_len(),
119 tokenizer: TokenizerType::default(),
120 language: None,
121 lowercase: None,
122 ascii_folding: None,
123 stopwords: None,
124 stemmer: None,
125 min_token_len: None,
126 max_token_len: None,
127 }
128 }
129}
130
131#[derive(Debug)]
133pub struct EdgeBm25 {
134 bm25: crate::bm25::Bm25,
135 tokenizer: Tokenizer,
136}
137
138impl EdgeBm25 {
139 pub fn new(config: EdgeBm25Config) -> Result<Self, EdgeBm25Error> {
140 let params = crate::bm25::Bm25Params {
141 k1: config.k.into_inner(),
142 b: config.b.into_inner(),
143 avg_doc_len: config.avg_len.into_inner(),
144 };
145
146 let processor = build_tokens_processor(
147 config.language,
148 config.lowercase,
149 config.ascii_folding,
150 config.stopwords,
151 config.stemmer,
152 config.min_token_len,
153 config.max_token_len,
154 )?;
155 let tokenizer = Tokenizer::new(config.tokenizer, processor);
156
157 Ok(Self {
158 bm25: crate::bm25::Bm25::new(params)?,
159 tokenizer,
160 })
161 }
162
163 pub fn embed_query(&self, text: &str) -> SparseVector {
164 let mut tokens: Vec<Cow<'_, str>> = Vec::new();
165 self.tokenizer.tokenize_query(text, |t| tokens.push(t));
166 to_sparse_vector(self.bm25.embed_query(&tokens))
167 }
168
169 pub fn embed_document(&self, text: &str) -> SparseVector {
170 let mut tokens: Vec<Cow<'_, str>> = Vec::new();
171 self.tokenizer.tokenize_doc(text, |t| tokens.push(t));
172 to_sparse_vector(self.bm25.embed_document(&tokens))
173 }
174}
175
176fn to_sparse_vector(e: SparseEmbedding) -> SparseVector {
177 SparseVector {
178 indices: e.indices,
179 values: e.values,
180 }
181}
182
183fn build_tokens_processor(
184 language: Option<String>,
185 lowercase: Option<bool>,
186 ascii_folding: Option<bool>,
187 stopwords: Option<StopwordsInterface>,
188 stemmer: Option<StemmingAlgorithm>,
189 min_token_len: Option<usize>,
190 max_token_len: Option<usize>,
191) -> Result<TokensProcessor, EdgeBm25Error> {
192 let lowercase = lowercase.unwrap_or(true);
193 let ascii_folding = ascii_folding.unwrap_or(false);
194
195 let resolved_language = match language {
198 Some(name) => {
199 Language::from_str(&name).map_err(|_| EdgeBm25Error::UnsupportedLanguage(name))?
200 }
201 None => DEFAULT_LANGUAGE,
202 };
203 let language_str = resolved_language.to_string();
204
205 let stemmer = match stemmer {
206 None => Stemmer::try_default_from_language(&language_str),
207 Some(algorithm) => Some(Stemmer::from_algorithm(&algorithm)),
208 };
209
210 let stopwords_config = match stopwords {
211 None => Some(StopwordsInterface::Language(resolved_language)),
212 Some(interface) => Some(interface),
213 };
214
215 Ok(TokensProcessor::new(
216 lowercase,
217 ascii_folding,
218 Arc::new(StopwordsFilter::new(&stopwords_config, lowercase)),
219 stemmer,
220 min_token_len,
221 max_token_len,
222 ))
223}
224
225#[cfg(test)]
226mod tests {
227 use super::*;
228
229 #[test]
230 fn defaults_construct_a_working_model() {
231 let model = EdgeBm25::new(EdgeBm25Config::default()).unwrap();
232 let text = "the quick brown fox jumps over the lazy dog";
233 let q = model.embed_query(text);
234 let d = model.embed_document(text);
235 assert!(!q.indices.is_empty());
237 assert!(!d.indices.is_empty());
238 assert!(q.values.iter().all(|&v| v == 1.0));
240 }
241
242 #[test]
243 fn english_stopwords_are_filtered_by_default() {
244 let model = EdgeBm25::new(EdgeBm25Config::default()).unwrap();
245 let with_stops = model.embed_query("the cat is a hunter");
247 let without_stops = model.embed_query("cat hunter");
248 assert_eq!(with_stops.indices.len(), without_stops.indices.len());
249 }
250
251 #[test]
252 fn custom_params_propagate() {
253 let cfg = EdgeBm25Config {
254 k: NotNan::new(2.0).unwrap(),
255 b: NotNan::new(0.5).unwrap(),
256 avg_len: NotNan::new(100.0).unwrap(),
257 ..Default::default()
258 };
259 let model = EdgeBm25::new(cfg).unwrap();
260 let v = model.embed_document("alpha beta gamma");
261 assert_eq!(v.indices.len(), 3);
262 }
263
264 #[test]
265 fn unsupported_language_is_rejected() {
266 let cfg = EdgeBm25Config {
267 language: Some("klingon".to_string()),
268 ..Default::default()
269 };
270 let err = EdgeBm25::new(cfg).expect_err("klingon should not be accepted");
271 assert!(matches!(err, EdgeBm25Error::UnsupportedLanguage(ref s) if s == "klingon"));
272 }
273
274 #[test]
275 fn invalid_avg_len_is_rejected() {
276 let cfg = EdgeBm25Config {
277 avg_len: NotNan::new(0.0).unwrap(),
278 ..Default::default()
279 };
280 let err = EdgeBm25::new(cfg).expect_err("avg_len=0 should not be accepted");
281 assert!(matches!(err, EdgeBm25Error::Bm25(_)));
282 }
283}