use std::{error::Error, fmt, future::Future, pin::Pin};
pub type EmbeddingFuture<'a, T> =
Pin<Box<dyn Future<Output = Result<T, EmbeddingProviderError>> + Send + 'a>>;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct EmbeddingRequest {
pub inputs: Vec<String>,
pub model: String,
pub dimension: u32,
}
#[derive(Debug, Clone, PartialEq)]
pub struct EmbeddingVector {
pub values: Vec<f64>,
}
pub trait EmbeddingProvider: Send + Sync {
fn embed(&self, request: EmbeddingRequest) -> EmbeddingFuture<'_, Vec<EmbeddingVector>>;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProviderRetryClass {
Retryable,
Permanent,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct EmbeddingProviderError {
pub retry: ProviderRetryClass,
pub status_code: Option<u16>,
pub code: String,
pub message: String,
}
impl fmt::Display for EmbeddingProviderError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self.status_code {
Some(status) => write!(formatter, "{} ({status}): {}", self.code, self.message),
None => write!(formatter, "{}: {}", self.code, self.message),
}
}
}
impl Error for EmbeddingProviderError {}
#[cfg(test)]
#[path = "embedding_tests.rs"]
mod tests;