use crate::completion::Usage;
use crate::error::ProviderError;
use serde::{Deserialize, Serialize};
impl<W, T> crate::driver::Model<W, T>
where
W: crate::wire::Wire<Op = crate::operation::Embedding>,
T: crate::driver::Transport<W>,
{
pub async fn embed_text(&self, text: &str) -> Result<Embedding, ProviderError> {
last_embedding(self.call(vec![text.to_owned()]).await?)
}
}
impl crate::driver::DynModel<crate::operation::Embedding> {
pub async fn embed_text(&self, text: &str) -> Result<Embedding, ProviderError> {
last_embedding(self.call(vec![text.to_owned()]).await?)
}
}
fn last_embedding(response: EmbeddingResponse) -> Result<Embedding, ProviderError> {
let mut embeddings = response.embeddings;
embeddings.pop().ok_or_else(|| {
ProviderError::Response(
"embedding provider returned an empty response for embed_text".to_string(),
)
})
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EmbeddingResponse {
pub embeddings: Vec<Embedding>,
#[serde(default)]
pub usage: Usage,
pub provider: String,
#[serde(default)]
pub model: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub response_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_request_id: Option<String>,
#[serde(default, skip_serializing_if = "serde_json::Value::is_null")]
pub raw: serde_json::Value,
}
impl EmbeddingResponse {
pub fn new(embeddings: Vec<Embedding>) -> Self {
Self {
embeddings,
usage: Usage::default(),
provider: String::new(),
model: None,
response_id: None,
provider_request_id: None,
raw: serde_json::Value::Null,
}
}
pub(crate) fn from_vectors(vectors: impl IntoIterator<Item = Vec<f64>>) -> Self {
Self::new(
vectors
.into_iter()
.map(|vec| Embedding {
document: String::new(),
vec,
})
.collect(),
)
}
}
#[derive(Clone, Default, Deserialize, Serialize, Debug)]
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 {}
#[cfg(test)]
mod provider_response_tests;
pub fn image_media_type(bytes: &[u8]) -> Option<&'static str> {
if bytes.starts_with(b"\x89PNG\r\n\x1a\n") {
Some("image/png")
} else if bytes.starts_with(b"\xff\xd8\xff") {
Some("image/jpeg")
} else if bytes.starts_with(b"GIF87a") || bytes.starts_with(b"GIF89a") {
Some("image/gif")
} else if bytes.starts_with(b"RIFF") && bytes.get(8..12) == Some(b"WEBP".as_slice()) {
Some("image/webp")
} else {
None
}
}
pub fn image_document(bytes: &[u8]) -> String {
use base64::Engine as _;
use sha2::Digest as _;
let media_type = image_media_type(bytes).unwrap_or("application/octet-stream");
let digest = sha2::Sha256::digest(bytes);
format!(
"{media_type};sha256={}",
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest)
)
}