use std::borrow::Cow;
use std::fmt;
use std::str::FromStr;
use std::sync::Arc;
use crate::bm25::SparseEmbedding;
use ordered_float::NotNan;
use crate::segment::data_types::index::{Language, StemmingAlgorithm, StopwordsInterface, TokenizerType};
use crate::segment::index::field_index::full_text_index::stop_words::StopwordsFilter;
use crate::segment::index::field_index::full_text_index::tokenizers::{
Stemmer, Tokenizer, TokensProcessor,
};
use crate::sparse::common::sparse_vector::SparseVector;
const DEFAULT_LANGUAGE: Language = Language::English;
#[derive(Debug, Clone, PartialEq)]
pub enum EdgeBm25Error {
Bm25(crate::bm25::Bm25Error),
UnsupportedLanguage(String),
}
impl fmt::Display for EdgeBm25Error {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Bm25(e) => write!(f, "{e}"),
Self::UnsupportedLanguage(lang) => write!(f, "unsupported language: {lang:?}"),
}
}
}
impl std::error::Error for EdgeBm25Error {}
impl From<crate::bm25::Bm25Error> for EdgeBm25Error {
fn from(e: crate::bm25::Bm25Error) -> Self {
Self::Bm25(e)
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct EdgeBm25Config {
#[serde(default = "default_k")]
pub k: NotNan<f64>,
#[serde(default = "default_b")]
pub b: NotNan<f64>,
#[serde(default = "default_avg_len")]
pub avg_len: NotNan<f64>,
#[serde(default)]
pub tokenizer: TokenizerType,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub language: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub lowercase: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub ascii_folding: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stopwords: Option<StopwordsInterface>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stemmer: Option<StemmingAlgorithm>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub min_token_len: Option<usize>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_token_len: Option<usize>,
}
const fn default_k() -> NotNan<f64> {
unsafe { NotNan::new_unchecked(1.2) }
}
const fn default_b() -> NotNan<f64> {
unsafe { NotNan::new_unchecked(0.75) }
}
const fn default_avg_len() -> NotNan<f64> {
unsafe { NotNan::new_unchecked(256.0) }
}
impl Default for EdgeBm25Config {
fn default() -> Self {
Self {
k: default_k(),
b: default_b(),
avg_len: default_avg_len(),
tokenizer: TokenizerType::default(),
language: None,
lowercase: None,
ascii_folding: None,
stopwords: None,
stemmer: None,
min_token_len: None,
max_token_len: None,
}
}
}
#[derive(Debug)]
pub struct EdgeBm25 {
bm25: crate::bm25::Bm25,
tokenizer: Tokenizer,
}
impl EdgeBm25 {
pub fn new(config: EdgeBm25Config) -> Result<Self, EdgeBm25Error> {
let params = crate::bm25::Bm25Params {
k1: config.k.into_inner(),
b: config.b.into_inner(),
avg_doc_len: config.avg_len.into_inner(),
};
let processor = build_tokens_processor(
config.language,
config.lowercase,
config.ascii_folding,
config.stopwords,
config.stemmer,
config.min_token_len,
config.max_token_len,
)?;
let tokenizer = Tokenizer::new(config.tokenizer, processor);
Ok(Self {
bm25: crate::bm25::Bm25::new(params)?,
tokenizer,
})
}
pub fn embed_query(&self, text: &str) -> SparseVector {
let mut tokens: Vec<Cow<'_, str>> = Vec::new();
self.tokenizer.tokenize_query(text, |t| tokens.push(t));
to_sparse_vector(self.bm25.embed_query(&tokens))
}
pub fn embed_document(&self, text: &str) -> SparseVector {
let mut tokens: Vec<Cow<'_, str>> = Vec::new();
self.tokenizer.tokenize_doc(text, |t| tokens.push(t));
to_sparse_vector(self.bm25.embed_document(&tokens))
}
}
fn to_sparse_vector(e: SparseEmbedding) -> SparseVector {
SparseVector {
indices: e.indices,
values: e.values,
}
}
fn build_tokens_processor(
language: Option<String>,
lowercase: Option<bool>,
ascii_folding: Option<bool>,
stopwords: Option<StopwordsInterface>,
stemmer: Option<StemmingAlgorithm>,
min_token_len: Option<usize>,
max_token_len: Option<usize>,
) -> Result<TokensProcessor, EdgeBm25Error> {
let lowercase = lowercase.unwrap_or(true);
let ascii_folding = ascii_folding.unwrap_or(false);
let resolved_language = match language {
Some(name) => {
Language::from_str(&name).map_err(|_| EdgeBm25Error::UnsupportedLanguage(name))?
}
None => DEFAULT_LANGUAGE,
};
let language_str = resolved_language.to_string();
let stemmer = match stemmer {
None => Stemmer::try_default_from_language(&language_str),
Some(algorithm) => Some(Stemmer::from_algorithm(&algorithm)),
};
let stopwords_config = match stopwords {
None => Some(StopwordsInterface::Language(resolved_language)),
Some(interface) => Some(interface),
};
Ok(TokensProcessor::new(
lowercase,
ascii_folding,
Arc::new(StopwordsFilter::new(&stopwords_config, lowercase)),
stemmer,
min_token_len,
max_token_len,
))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn defaults_construct_a_working_model() {
let model = EdgeBm25::new(EdgeBm25Config::default()).unwrap();
let text = "the quick brown fox jumps over the lazy dog";
let q = model.embed_query(text);
let d = model.embed_document(text);
assert!(!q.indices.is_empty());
assert!(!d.indices.is_empty());
assert!(q.values.iter().all(|&v| v == 1.0));
}
#[test]
fn english_stopwords_are_filtered_by_default() {
let model = EdgeBm25::new(EdgeBm25Config::default()).unwrap();
let with_stops = model.embed_query("the cat is a hunter");
let without_stops = model.embed_query("cat hunter");
assert_eq!(with_stops.indices.len(), without_stops.indices.len());
}
#[test]
fn custom_params_propagate() {
let cfg = EdgeBm25Config {
k: NotNan::new(2.0).unwrap(),
b: NotNan::new(0.5).unwrap(),
avg_len: NotNan::new(100.0).unwrap(),
..Default::default()
};
let model = EdgeBm25::new(cfg).unwrap();
let v = model.embed_document("alpha beta gamma");
assert_eq!(v.indices.len(), 3);
}
#[test]
fn unsupported_language_is_rejected() {
let cfg = EdgeBm25Config {
language: Some("klingon".to_string()),
..Default::default()
};
let err = EdgeBm25::new(cfg).expect_err("klingon should not be accepted");
assert!(matches!(err, EdgeBm25Error::UnsupportedLanguage(ref s) if s == "klingon"));
}
#[test]
fn invalid_avg_len_is_rejected() {
let cfg = EdgeBm25Config {
avg_len: NotNan::new(0.0).unwrap(),
..Default::default()
};
let err = EdgeBm25::new(cfg).expect_err("avg_len=0 should not be accepted");
assert!(matches!(err, EdgeBm25Error::Bm25(_)));
}
}