use crate::config::Config;
use crate::embeddings::EmbeddingsConfig;
use crate::embeddings::EmbeddingsManager;
use thiserror::Error;
pub type EmbeddingVector = Vec<f32>;
#[derive(Debug, Error)]
pub enum EmbeddingError {
#[error("not implemented")]
NotImplemented,
#[error(transparent)]
Io(#[from] std::io::Error),
#[error("provider error: {0}")]
Provider(String),
#[error("embeddings error: {0}")]
EmbeddingsError(#[from] crate::embeddings::EmbeddingError),
}
#[allow(async_fn_in_trait)]
pub trait EmbeddingModel {
fn dimensions(&self) -> usize {
1536 }
async fn embed(&self, _text: &str) -> Result<EmbeddingVector, EmbeddingError> {
Err(EmbeddingError::NotImplemented)
}
async fn embed_batch(&self, _texts: &[String]) -> Result<Vec<EmbeddingVector>, EmbeddingError> {
Err(EmbeddingError::NotImplemented)
}
}
pub struct EmbeddingModelBridge {
manager: Option<EmbeddingsManager>,
}
impl EmbeddingModelBridge {
pub const fn new(manager: Option<EmbeddingsManager>) -> Self {
Self { manager }
}
pub fn from_embeddings_config(config: Option<EmbeddingsConfig>) -> Self {
let manager = EmbeddingsManager::new(config);
Self::new(Some(manager))
}
pub const fn from_config(_config: &Config) -> Self {
Self::new(None)
}
pub const fn disabled() -> Self {
Self::new(None)
}
pub fn is_enabled(&self) -> bool {
self.manager
.as_ref()
.map(|m| m.is_enabled())
.unwrap_or(false)
}
}
impl EmbeddingModel for EmbeddingModelBridge {
fn dimensions(&self) -> usize {
if let Some(manager) = &self.manager {
manager.current_dimensions().unwrap_or(1536)
} else {
1536 }
}
async fn embed(&self, text: &str) -> Result<EmbeddingVector, EmbeddingError> {
match &self.manager {
Some(manager) => {
if !manager.is_enabled() {
return Err(EmbeddingError::NotImplemented);
}
let result = manager.embed(text).await?;
match result {
Some(vector) => Ok(vector),
None => Err(EmbeddingError::NotImplemented),
}
}
None => Err(EmbeddingError::NotImplemented),
}
}
async fn embed_batch(&self, texts: &[String]) -> Result<Vec<EmbeddingVector>, EmbeddingError> {
match &self.manager {
Some(manager) => {
if !manager.is_enabled() {
return Err(EmbeddingError::NotImplemented);
}
let result = manager.embed_batch(texts).await?;
match result {
Some(vectors) => Ok(vectors),
None => Err(EmbeddingError::NotImplemented),
}
}
None => Err(EmbeddingError::NotImplemented),
}
}
}
pub struct NoOpEmbeddingModel;
impl EmbeddingModel for NoOpEmbeddingModel {
fn dimensions(&self) -> usize {
1536
}
async fn embed(&self, _text: &str) -> Result<EmbeddingVector, EmbeddingError> {
Err(EmbeddingError::NotImplemented)
}
async fn embed_batch(&self, _texts: &[String]) -> Result<Vec<EmbeddingVector>, EmbeddingError> {
Err(EmbeddingError::NotImplemented)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_disabled_bridge_has_zero_overhead() {
let bridge = EmbeddingModelBridge::disabled();
assert!(!bridge.is_enabled());
assert_eq!(bridge.dimensions(), 1536);
}
#[tokio::test]
async fn test_disabled_bridge_returns_not_implemented() {
let bridge = EmbeddingModelBridge::disabled();
let result = bridge.embed("test").await;
assert!(matches!(result, Err(EmbeddingError::NotImplemented)));
let batch_result = bridge.embed_batch(&["test".to_string()]).await;
assert!(matches!(batch_result, Err(EmbeddingError::NotImplemented)));
}
#[test]
fn test_noop_embedding_model() {
let model = NoOpEmbeddingModel;
assert_eq!(model.dimensions(), 1536);
}
#[tokio::test]
async fn test_noop_embedding_model_returns_not_implemented() {
let model = NoOpEmbeddingModel;
let result = model.embed("test").await;
assert!(matches!(result, Err(EmbeddingError::NotImplemented)));
let batch_result = model.embed_batch(&["test".to_string()]).await;
assert!(matches!(batch_result, Err(EmbeddingError::NotImplemented)));
}
}