use std::{future::Future, pin::Pin};
use runifold_core::Usage;
use crate::{RetrievalContext, RetrievalError};
#[cfg(not(target_arch = "wasm32"))]
pub type EmbeddingFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
#[cfg(target_arch = "wasm32")]
pub type EmbeddingFuture<'a, T> = Pin<Box<dyn Future<Output = T> + 'a>>;
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
#[non_exhaustive]
pub enum EmbeddingTask {
#[default]
Unspecified,
RetrievalQuery,
RetrievalDocument,
SemanticSimilarity,
Classification,
Clustering,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct EmbeddingRequest {
inputs: Vec<String>,
task: EmbeddingTask,
}
impl EmbeddingRequest {
pub fn new(inputs: Vec<String>, task: EmbeddingTask) -> Result<Self, RetrievalError> {
if let Some(index) = inputs.iter().position(|input| input.trim().is_empty()) {
return Err(RetrievalError::EmptyEmbeddingInput { index });
}
Ok(Self { inputs, task })
}
pub fn inputs(&self) -> &[String] {
&self.inputs
}
pub const fn task(&self) -> EmbeddingTask {
self.task
}
pub fn into_parts(self) -> (Vec<String>, EmbeddingTask) {
(self.inputs, self.task)
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct Embedding {
values: Vec<f64>,
squared_norm: f64,
}
impl Embedding {
pub fn new(values: Vec<f64>) -> Result<Self, RetrievalError> {
if values.is_empty() {
return Err(RetrievalError::EmptyEmbedding);
}
if let Some(index) = values.iter().position(|value| !value.is_finite()) {
return Err(RetrievalError::NonFiniteEmbedding { index });
}
let squared_norm = values.iter().map(|value| value.powi(2)).sum::<f64>();
if squared_norm == 0.0 {
return Err(RetrievalError::ZeroNormEmbedding);
}
Ok(Self {
values,
squared_norm,
})
}
pub fn values(&self) -> &[f64] {
&self.values
}
pub fn dimensions(&self) -> usize {
self.values.len()
}
pub(crate) fn squared_norm(&self) -> f64 {
self.squared_norm
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct EmbeddingBatch {
pub embeddings: Vec<Embedding>,
pub usage: Usage,
}
impl EmbeddingBatch {
pub fn validate_count(self, expected: usize) -> Result<Self, RetrievalError> {
let actual = self.embeddings.len();
if actual != expected {
return Err(RetrievalError::EmbeddingCountMismatch { expected, actual });
}
Ok(self)
}
}
pub trait EmbeddingModel: Send + Sync {
fn embed(
&self,
request: EmbeddingRequest,
context: RetrievalContext,
) -> EmbeddingFuture<'_, Result<EmbeddingBatch, RetrievalError>>;
}