use serde::{Deserialize, Serialize};
use thiserror::Error;
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
pub struct Embedding {
pub document: String,
pub vec: Vec<f64>,
}
impl PartialEq for Embedding {
fn eq(&self, other: &Self) -> bool {
self.document == other.document
}
}
impl Eq for Embedding {}
#[derive(Debug, Error)]
pub enum EmbeddingError {
#[error("HTTP error: {0}")]
Http(String),
#[error("JSON error: {0}")]
Json(#[from] serde_json::Error),
#[error("Provider error ({status}): {message}")]
Provider { status: u16, message: String },
#[error("Response error: {0}")]
Response(String),
}
pub trait EmbeddingModel {
const MAX_DOCUMENTS: usize;
type Error: std::error::Error + 'static;
fn ndims(&self) -> usize;
fn embed_texts(
&self,
texts: Vec<String>,
) -> impl std::future::Future<Output = Result<Vec<Embedding>, Self::Error>>;
fn embed_text(
&self,
text: &str,
) -> impl std::future::Future<Output = Result<Embedding, Self::Error>> {
async move {
self.embed_texts(vec![text.to_owned()])
.await?
.into_iter()
.next()
.ok_or_else(|| unreachable!("embed_texts returned empty vec for one input"))
}
}
}