use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use crate::client::LlmClient;
use crate::error::Result;
use crate::types::EmbeddingRequest;
#[cfg_attr(alef, alef(skip))]
pub trait EmbeddingProvider: Send + Sync + 'static {
fn embed<'a>(&'a self, text: &'a str) -> Pin<Box<dyn Future<Output = Result<Vec<f32>>> + Send + 'a>>;
fn dim(&self) -> usize;
}
#[cfg_attr(alef, alef(skip))]
pub struct SelfHostedEmbeddingProvider {
client: Arc<dyn LlmClient>,
model: String,
dim: usize,
}
impl SelfHostedEmbeddingProvider {
#[must_use]
pub fn new(client: Arc<dyn LlmClient>, model: impl Into<String>, dim: usize) -> Self {
Self {
client,
model: model.into(),
dim,
}
}
}
impl EmbeddingProvider for SelfHostedEmbeddingProvider {
fn embed<'a>(&'a self, text: &'a str) -> Pin<Box<dyn Future<Output = Result<Vec<f32>>> + Send + 'a>> {
let req = EmbeddingRequest {
model: self.model.clone(),
input: crate::types::EmbeddingInput::Single(text.to_owned()),
encoding_format: None,
dimensions: Some(self.dim as u32),
user: None,
};
let client = Arc::clone(&self.client);
Box::pin(async move {
let resp = client.embed(req).await?;
let vec: Vec<f32> = resp
.data
.into_iter()
.next()
.map(|obj| obj.embedding.into_iter().map(|x| x as f32).collect())
.unwrap_or_default();
Ok(vec)
})
}
fn dim(&self) -> usize {
self.dim
}
}
#[cfg_attr(alef, alef(skip))]
#[derive(Debug, Clone)]
pub struct NoOpEmbeddingProvider {
pub dim: usize,
}
impl EmbeddingProvider for NoOpEmbeddingProvider {
fn embed<'a>(&'a self, _text: &'a str) -> Pin<Box<dyn Future<Output = Result<Vec<f32>>> + Send + 'a>> {
Box::pin(std::future::ready(Ok(vec![0.0_f32; self.dim])))
}
fn dim(&self) -> usize {
self.dim
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::*;
use crate::client::LlmClient;
use crate::client::{BoxFuture, BoxStream};
use crate::error::LiterLlmError;
use crate::types::{
ChatCompletionChunk, ChatCompletionRequest, ChatCompletionResponse, EmbeddingObject, EmbeddingRequest,
EmbeddingResponse, ModelsListResponse, Usage,
audio::{CreateSpeechRequest, CreateTranscriptionRequest, TranscriptionResponse},
image::{CreateImageRequest, ImagesResponse},
moderation::{ModerationRequest, ModerationResponse},
ocr::{OcrRequest, OcrResponse},
rerank::{RerankRequest, RerankResponse},
search::{SearchRequest, SearchResponse},
};
#[tokio::test]
async fn no_op_embedding_provider_returns_zero_vector() {
let provider = NoOpEmbeddingProvider { dim: 4 };
let vec = provider.embed("hello world").await.unwrap();
assert_eq!(vec.len(), 4);
assert!(vec.iter().all(|&x| x == 0.0));
}
#[tokio::test]
async fn no_op_embedding_provider_dim_is_consistent() {
let provider = NoOpEmbeddingProvider { dim: 128 };
assert_eq!(provider.dim(), 128);
let vec = provider.embed("test").await.unwrap();
assert_eq!(vec.len(), provider.dim());
}
#[derive(Clone)]
struct MockEmbedClient {
embedding: Vec<f32>,
}
impl MockEmbedClient {
fn new(embedding: Vec<f32>) -> Self {
Self { embedding }
}
}
impl LlmClient for MockEmbedClient {
fn chat(&self, _req: ChatCompletionRequest) -> BoxFuture<'_, crate::error::Result<ChatCompletionResponse>> {
Box::pin(async {
Err(LiterLlmError::EndpointNotSupported {
endpoint: "chat".into(),
provider: "mock".into(),
})
})
}
fn chat_stream(
&self,
_req: ChatCompletionRequest,
) -> BoxFuture<'_, crate::error::Result<BoxStream<'static, crate::error::Result<ChatCompletionChunk>>>>
{
Box::pin(async {
Err(LiterLlmError::EndpointNotSupported {
endpoint: "chat_stream".into(),
provider: "mock".into(),
})
})
}
fn embed(&self, req: EmbeddingRequest) -> BoxFuture<'_, crate::error::Result<EmbeddingResponse>> {
let embedding: Vec<f64> = self.embedding.iter().map(|&v| f64::from(v)).collect();
Box::pin(async move {
Ok(EmbeddingResponse {
object: "list".into(),
data: vec![EmbeddingObject {
object: "embedding".into(),
embedding,
index: 0,
}],
model: req.model,
usage: Some(Usage {
prompt_tokens: 4,
completion_tokens: 0,
total_tokens: 4,
prompt_tokens_details: None,
}),
})
})
}
fn list_models(&self) -> BoxFuture<'_, crate::error::Result<ModelsListResponse>> {
Box::pin(async {
Ok(ModelsListResponse {
object: "list".into(),
data: vec![],
})
})
}
fn image_generate(&self, _req: CreateImageRequest) -> BoxFuture<'_, crate::error::Result<ImagesResponse>> {
Box::pin(async {
Ok(ImagesResponse {
created: 0,
data: vec![],
})
})
}
fn speech(&self, _req: CreateSpeechRequest) -> BoxFuture<'_, crate::error::Result<bytes::Bytes>> {
Box::pin(async { Ok(bytes::Bytes::new()) })
}
fn transcribe(
&self,
_req: CreateTranscriptionRequest,
) -> BoxFuture<'_, crate::error::Result<TranscriptionResponse>> {
Box::pin(async {
Ok(TranscriptionResponse {
text: String::new(),
language: None,
duration: None,
segments: None,
})
})
}
fn moderate(&self, _req: ModerationRequest) -> BoxFuture<'_, crate::error::Result<ModerationResponse>> {
Box::pin(async {
Ok(ModerationResponse {
id: String::new(),
model: String::new(),
results: vec![],
})
})
}
fn rerank(&self, _req: RerankRequest) -> BoxFuture<'_, crate::error::Result<RerankResponse>> {
Box::pin(async {
Ok(RerankResponse {
id: None,
results: vec![],
meta: None,
})
})
}
fn search(&self, _req: SearchRequest) -> BoxFuture<'_, crate::error::Result<SearchResponse>> {
Box::pin(async {
Err(LiterLlmError::EndpointNotSupported {
endpoint: "search".into(),
provider: "mock".into(),
})
})
}
fn ocr(&self, _req: OcrRequest) -> BoxFuture<'_, crate::error::Result<OcrResponse>> {
Box::pin(async {
Err(LiterLlmError::EndpointNotSupported {
endpoint: "ocr".into(),
provider: "mock".into(),
})
})
}
}
#[tokio::test]
async fn self_hosted_embedding_provider_round_trips_through_mock_client() {
let expected_vec = vec![0.1_f32, 0.2, 0.3, 0.4];
let client = Arc::new(MockEmbedClient::new(expected_vec.clone()));
let provider = SelfHostedEmbeddingProvider::new(client, "openai/text-embedding-3-small", 4);
let result = provider.embed("hello world").await.unwrap();
assert_eq!(result, expected_vec, "should return the mock client's embedding vector");
}
#[tokio::test]
async fn self_hosted_embedding_provider_dim_matches_constructor() {
let client = Arc::new(MockEmbedClient::new(vec![0.0; 1536]));
let provider = SelfHostedEmbeddingProvider::new(client, "openai/text-embedding-3-small", 1536);
assert_eq!(provider.dim(), 1536);
}
}