Skip to main content

qdrant_edge/edge/
bm25_embed.rs

1//! Edge-side BM25 wiring.
2//!
3//! Builds [`bm25::Bm25`] and runs it over `segment`'s tokenizer pipeline (so
4//! stopwords, stemming, and language defaults match server behavior). Emits
5//! [`sparse::common::sparse_vector::SparseVector`] which the rest of the edge
6//! API already understands.
7//!
8//! No `api` crate dependency: the public contract types belong to the REST/gRPC
9//! layer; edge consumers (e.g. Python bindings) construct [`EdgeBm25Config`]
10//! from their own input format.
11
12use 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/// Error returned by [`EdgeBm25::new`] for invalid configuration.
29#[derive(Debug, Clone, PartialEq)]
30pub enum EdgeBm25Error {
31    /// BM25 hyperparameters failed validation (see [`bm25::Bm25Error`]).
32    Bm25(crate::bm25::Bm25Error),
33    /// `language` did not match any supported [`Language`] variant.
34    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/// Configuration for an edge-side BM25 model.
55///
56/// JSON shape mirrors `lib/api`'s REST `Bm25Config` so configs are portable
57/// between cloud and edge: `k`, `b`, `avg_len`, `tokenizer`, plus the
58/// preprocessing fields (`language`, `lowercase`, `ascii_folding`, `stopwords`,
59/// `stemmer`, `min_token_len`, `max_token_len`).
60#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
61pub struct EdgeBm25Config {
62    /// Term-frequency saturation. Higher values make TF have more impact.
63    /// Default 1.2.
64    #[serde(default = "default_k")]
65    pub k: NotNan<f64>,
66    /// Document length normalization. 0 = none, 1 = full. Default 0.75.
67    #[serde(default = "default_b")]
68    pub b: NotNan<f64>,
69    /// Expected average document length in tokens. Default 256.
70    #[serde(default = "default_avg_len")]
71    pub avg_len: NotNan<f64>,
72    /// Tokenizer type used for text preprocessing.
73    #[serde(default)]
74    pub tokenizer: TokenizerType,
75    /// Language for default stopwords / stemmer (English if unset).
76    #[serde(default, skip_serializing_if = "Option::is_none")]
77    pub language: Option<String>,
78    /// Lowercase before tokenization. Default true.
79    #[serde(default, skip_serializing_if = "Option::is_none")]
80    pub lowercase: Option<bool>,
81    /// Fold accented characters to ASCII (e.g. `"ação" → "acao"`). Default false.
82    #[serde(default, skip_serializing_if = "Option::is_none")]
83    pub ascii_folding: Option<bool>,
84    /// Stopwords filter configuration; defaults from `language`.
85    #[serde(default, skip_serializing_if = "Option::is_none")]
86    pub stopwords: Option<StopwordsInterface>,
87    /// Stemmer configuration; defaults from `language`.
88    #[serde(default, skip_serializing_if = "Option::is_none")]
89    pub stemmer: Option<StemmingAlgorithm>,
90    /// Discard tokens shorter than this.
91    #[serde(default, skip_serializing_if = "Option::is_none")]
92    pub min_token_len: Option<usize>,
93    /// Discard tokens longer than this.
94    #[serde(default, skip_serializing_if = "Option::is_none")]
95    pub max_token_len: Option<usize>,
96}
97
98const fn default_k() -> NotNan<f64> {
99    // Safe: 1.2 is not NaN.
100    unsafe { NotNan::new_unchecked(1.2) }
101}
102
103const fn default_b() -> NotNan<f64> {
104    // Safe: 0.75 is not NaN.
105    unsafe { NotNan::new_unchecked(0.75) }
106}
107
108const fn default_avg_len() -> NotNan<f64> {
109    // Safe: 256.0 is not NaN.
110    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/// Edge-side BM25 model. Embeds raw text into [`SparseVector`].
132#[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    // Resolve language up-front so a typo / unsupported value fails the build
196    // instead of silently disabling stemming and stopwords.
197    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        // Stopwords ("the") get filtered, so we should still have content.
236        assert!(!q.indices.is_empty());
237        assert!(!d.indices.is_empty());
238        // Query is unit-weighted, document is TF-weighted (different by construction).
239        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        // "the", "a", "is" are stopwords — should not contribute distinct indices.
246        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}