use crate::traits::{Embedder, Extractor, Reranker, Summarizer};
use std::sync::Arc;
pub struct ModelRegistry {
pub embedder: Arc<dyn Embedder>,
pub extractor: Arc<dyn Extractor>,
pub reranker: Arc<dyn Reranker>,
pub summarizer: Arc<dyn Summarizer>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BackendChoice {
Api,
Deterministic,
Auto,
}
#[derive(Debug, Clone)]
pub struct BackendSelection {
pub extractor: BackendChoice,
pub reranker: BackendChoice,
pub summarizer: BackendChoice,
}
impl Default for BackendSelection {
fn default() -> Self {
Self {
extractor: BackendChoice::Auto,
reranker: BackendChoice::Auto,
summarizer: BackendChoice::Auto,
}
}
}
#[derive(Debug, Clone)]
pub struct BackendUsage {
pub embedder: String,
pub reranker: Option<String>,
}
use hippmem_core::config::EmbedderConfig;
pub fn build_embedder(
config: &EmbedderConfig,
) -> crate::error::ModelResult<std::sync::Arc<dyn crate::traits::Embedder>> {
match config {
EmbedderConfig::Hash { dimensions } => Ok(std::sync::Arc::new(
crate::deterministic::embed::DeterministicEmbedder::new(*dimensions),
)),
EmbedderConfig::Neural {
base_url,
model,
api_key,
dimensions,
} => {
let key = match api_key {
Some(k) if !k.is_empty() => k.clone(),
_ => std::env::var("OPENAI_API_KEY").unwrap_or_default(),
};
if key.is_empty() {
return Err(crate::error::ModelError::Auth(model.clone()));
}
let embedder = crate::api::openai::OpenAiEmbedder::new_with_base_url(
key,
base_url,
model,
*dimensions,
)?;
Ok(std::sync::Arc::new(embedder))
}
EmbedderConfig::Onnx { .. } => Err(crate::error::ModelError::Unavailable(
"onnx backend not yet implemented".to_string(),
)),
}
}
#[cfg(test)]
mod tests {
use super::*;
use hippmem_core::config::EmbedderConfig;
#[test]
fn build_embedder_hash_default() {
let cfg = EmbedderConfig::default(); let embedder = build_embedder(&cfg).unwrap();
assert_eq!(embedder.dim(), 256);
assert_eq!(embedder.backend_id(), "deterministic-hash");
}
#[test]
fn build_embedder_hash_custom_dim() {
let cfg = EmbedderConfig::Hash { dimensions: 512 };
let embedder = build_embedder(&cfg).unwrap();
assert_eq!(embedder.dim(), 512);
}
#[test]
fn onnx_returns_unavailable() {
let cfg = EmbedderConfig::Onnx {
model_name: "test-model".into(),
model_cache_dir: std::path::PathBuf::from("/tmp"),
dimensions: 512,
};
let result = build_embedder(&cfg);
match &result {
Err(e) => {
let err_msg = format!("{e}");
assert!(
err_msg.contains("onnx"),
"error message should mention onnx, got: {err_msg}"
);
}
Ok(_) => panic!("should return an error when ONNX is not implemented"),
}
}
}