use std::path::PathBuf;
use std::sync::{Arc, Mutex};
pub const DEFAULT_MODEL: &str = "bge-small-en-v1.5";
#[derive(Debug, Clone, Copy, PartialEq, serde::Serialize, schemars::JsonSchema)]
#[serde(rename_all = "lowercase")]
pub enum EmbedderStatus {
Off,
Downloading,
Ready,
Failed,
}
struct Inner {
model_name: String,
status: Mutex<EmbedderStatus>,
engine: Mutex<Option<fastembed::TextEmbedding>>,
}
#[derive(Clone)]
pub struct Embedder {
inner: Arc<Inner>,
}
fn resolve_model(model_name: &str) -> Option<fastembed::EmbeddingModel> {
if model_name == DEFAULT_MODEL {
Some(fastembed::EmbeddingModel::BGESmallENV15)
} else {
None
}
}
impl Embedder {
pub fn disabled() -> Self {
Self {
inner: Arc::new(Inner {
model_name: DEFAULT_MODEL.into(),
status: Mutex::new(EmbedderStatus::Off),
engine: Mutex::new(None),
}),
}
}
pub fn start(model: Option<String>, cache_dir: PathBuf, ort_download: bool) -> Self {
let model_name = model.unwrap_or_else(|| DEFAULT_MODEL.into());
let e = Self {
inner: Arc::new(Inner {
model_name: model_name.clone(),
status: Mutex::new(EmbedderStatus::Downloading),
engine: Mutex::new(None),
}),
};
let Some(known_model) = resolve_model(&model_name) else {
*e.inner.status.lock().unwrap() = EmbedderStatus::Failed;
eprintln!(
"topodb-mcp: unknown embedding model {model_name:?} (only {DEFAULT_MODEL:?} is \
currently wired up); running text-only"
);
return e;
};
let init = e.clone();
std::thread::spawn(move || {
let resolve_result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
crate::ort_fetch::resolve(&cache_dir, ort_download, &crate::ort_fetch::http_fetch)
}));
let resolved = match resolve_result {
Ok(resolved) => resolved,
Err(_) => {
*init.inner.status.lock().unwrap() = EmbedderStatus::Failed;
eprintln!(
"topodb-mcp: embedding model {model_name} runtime resolution panicked; \
running text-only"
);
return;
}
};
match resolved {
crate::ort_fetch::OrtRuntime::EnvOverride
| crate::ort_fetch::OrtRuntime::System => {}
crate::ort_fetch::OrtRuntime::Local(dylib) => {
let init_result =
std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
ort::init_from(&dylib).map(|builder| {
let _ = builder.commit();
})
}));
match init_result {
Ok(Ok(())) => {}
Ok(Err(e)) => {
*init.inner.status.lock().unwrap() = EmbedderStatus::Failed;
eprintln!(
"topodb-mcp: embedding model {model_name} unavailable (ort \
init from {} failed: {e}); running text-only",
dylib.display()
);
return;
}
Err(_) => {
*init.inner.status.lock().unwrap() = EmbedderStatus::Failed;
eprintln!(
"topodb-mcp: embedding model {model_name} unavailable (ort \
init from {} panicked); running text-only",
dylib.display()
);
return;
}
}
}
crate::ort_fetch::OrtRuntime::Unavailable(reason) => {
*init.inner.status.lock().unwrap() = EmbedderStatus::Failed;
eprintln!(
"topodb-mcp: embedding model {model_name} unavailable ({reason}); \
running text-only"
);
return;
}
}
let result = std::panic::catch_unwind(|| {
fastembed::TextEmbedding::try_new(
fastembed::TextInitOptions::new(known_model).with_cache_dir(cache_dir),
)
});
match result {
Ok(Ok(engine)) => {
*init.inner.engine.lock().unwrap() = Some(engine);
*init.inner.status.lock().unwrap() = EmbedderStatus::Ready;
eprintln!("topodb-mcp: embedding model {model_name} ready");
}
Ok(Err(err)) => {
*init.inner.status.lock().unwrap() = EmbedderStatus::Failed;
eprintln!(
"topodb-mcp: embedding model {model_name} unavailable ({err}); running text-only"
);
}
Err(_) => {
*init.inner.status.lock().unwrap() = EmbedderStatus::Failed;
eprintln!(
"topodb-mcp: embedding model {model_name} init panicked; running text-only"
);
}
}
});
e
}
pub fn status(&self) -> EmbedderStatus {
*self.inner.status.lock().unwrap()
}
pub fn model_name(&self) -> String {
self.inner.model_name.clone()
}
pub fn embed(&self, text: &str) -> Option<Vec<f32>> {
let mut guard = self.inner.engine.lock().unwrap();
let engine = guard.as_mut()?;
match engine.embed(vec![text], None) {
Ok(mut vs) if !vs.is_empty() => Some(vs.remove(0)),
_ => None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn disabled_is_off_forever() {
let e = Embedder::disabled();
assert_eq!(e.status(), EmbedderStatus::Off);
assert_eq!(e.model_name(), DEFAULT_MODEL);
assert_eq!(e.embed("hello"), None);
}
#[test]
fn unknown_model_degrades_to_failed_without_crashing() {
let dir = tempfile::tempdir().unwrap();
let e = Embedder::start(
Some("not-a-real-model".into()),
dir.path().to_path_buf(),
true,
);
assert_eq!(e.status(), EmbedderStatus::Failed);
assert_eq!(e.embed("hello"), None);
}
#[test]
#[ignore]
fn real_ort_download_reaches_ready() {
let dir = tempfile::tempdir().unwrap();
let e = Embedder::start(None, dir.path().to_path_buf(), true);
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(600);
while e.status() == EmbedderStatus::Downloading {
assert!(
std::time::Instant::now() < deadline,
"init did not reach a terminal status within 10 minutes"
);
std::thread::sleep(std::time::Duration::from_millis(500));
}
assert_eq!(e.status(), EmbedderStatus::Ready);
let v = e.embed("hello embeddings").expect("Ready must embed");
assert_eq!(v.len(), 384, "bge-small-en-v1.5 is 384-dim");
}
}