use std::future::Future;
use serde::Deserialize;
use serde::Serialize;
use tokio_util::sync::CancellationToken;
use crate::error::ProviderError;
use crate::language_model::ResponseMetadata;
use crate::shared::Headers;
use crate::shared::ModelId;
use crate::shared::ProviderId;
use crate::shared::ProviderMetadata;
use crate::shared::ProviderOptions;
use crate::shared::Warning;
pub type Embedding = Vec<f32>;
pub trait EmbeddingModel: Send + Sync + 'static {
fn provider(&self) -> &ProviderId;
fn model_id(&self) -> &ModelId;
fn max_embeddings_per_call(&self) -> Option<usize>;
fn max_input_bytes_per_call(&self) -> Option<usize> {
None
}
fn supports_parallel_calls(&self) -> bool;
fn do_embed(
&self,
options: EmbedOptions,
) -> impl Future<Output = Result<EmbedResult, ProviderError>> + Send;
}
#[derive(Debug, Clone, Default)]
pub struct EmbedOptions {
pub values: Vec<String>,
pub headers: Headers,
pub provider_options: ProviderOptions,
pub cancellation: CancellationToken,
}
impl EmbedOptions {
#[must_use]
pub fn new(values: Vec<String>) -> Self {
Self {
values,
..Self::default()
}
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct EmbeddingUsage {
pub tokens: u64,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct EmbedResult {
pub embeddings: Vec<Embedding>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub usage: Option<EmbeddingUsage>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_metadata: Option<ProviderMetadata>,
#[serde(default)]
pub response: ResponseMetadata,
#[serde(default)]
pub warnings: Vec<Warning>,
}