use base64::Engine as _;
use serde::{Deserialize, Serialize};
use super::common::{ImageUrl, Usage};
use crate::cost;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum EmbeddingFormat {
Float,
Base64,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct EmbeddingRequest {
pub model: String,
pub input: EmbeddingInput,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub encoding_format: Option<EmbeddingFormat>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub dimensions: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum EmbeddingInput {
Single(String),
Multiple(Vec<String>),
Multimodal(Vec<EmbeddingContentPart>),
}
#[cfg_attr(alef, alef(skip))]
impl Default for EmbeddingInput {
fn default() -> Self {
Self::Single(String::new())
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum EmbeddingContentPart {
#[serde(rename = "text")]
Text { text: String },
#[serde(rename = "image_url")]
ImageUrl { image_url: ImageUrl },
#[serde(rename = "image_base64")]
ImageBase64 { image_base64: String },
}
impl EmbeddingContentPart {
pub fn text(text: impl Into<String>) -> Self {
Self::Text { text: text.into() }
}
pub fn image_url(url: impl Into<String>) -> Self {
Self::ImageUrl {
image_url: ImageUrl {
url: url.into(),
detail: None,
},
}
}
pub fn image_base64(data_url: impl Into<String>) -> Self {
Self::ImageBase64 {
image_base64: data_url.into(),
}
}
pub fn image_bytes(bytes: &[u8], mime_type: Option<&str>) -> Self {
Self::image_base64(crate::image::encode_data_url(bytes, mime_type))
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct EmbeddingResponse {
pub object: String,
pub data: Vec<EmbeddingObject>,
pub model: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub usage: Option<Usage>,
}
impl EmbeddingResponse {
#[cfg_attr(alef, alef(skip))]
#[must_use]
pub fn estimated_cost(&self) -> Option<f64> {
let usage = self.usage.as_ref()?;
cost::completion_cost(&self.model, usage.prompt_tokens, usage.completion_tokens)
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct EmbeddingObject {
pub object: String,
#[serde(deserialize_with = "deserialize_embedding")]
pub embedding: Vec<f32>,
pub index: u32,
}
fn deserialize_embedding<'de, D>(deserializer: D) -> Result<Vec<f32>, D::Error>
where
D: serde::Deserializer<'de>,
{
struct EmbeddingVisitor;
impl<'de> serde::de::Visitor<'de> for EmbeddingVisitor {
type Value = Vec<f32>;
fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
formatter.write_str("a float array or a base64-encoded string of little-endian f32 bytes")
}
fn visit_str<E>(self, value: &str) -> Result<Vec<f32>, E>
where
E: serde::de::Error,
{
let bytes = base64::engine::general_purpose::STANDARD
.decode(value)
.map_err(|e| E::custom(format!("invalid base64 embedding: {e}")))?;
if bytes.len() % 4 != 0 {
return Err(E::custom(format!(
"base64 embedding length {} is not a multiple of 4",
bytes.len()
)));
}
Ok(bytes
.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect())
}
fn visit_seq<A>(self, mut seq: A) -> Result<Vec<f32>, A::Error>
where
A: serde::de::SeqAccess<'de>,
{
let mut out = Vec::with_capacity(seq.size_hint().unwrap_or(0));
while let Some(value) = seq.next_element::<f32>()? {
out.push(value);
}
Ok(out)
}
}
deserializer.deserialize_any(EmbeddingVisitor)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn multimodal_input_round_trips_text_and_image_url() {
let input = EmbeddingInput::Multimodal(vec![
EmbeddingContentPart::text("a red bicycle"),
EmbeddingContentPart::image_url("https://example.com/bicycle.png"),
]);
let json = serde_json::to_string(&input).expect("serialization should not fail");
assert_eq!(
json,
r#"[{"type":"text","text":"a red bicycle"},{"type":"image_url","image_url":{"url":"https://example.com/bicycle.png"}}]"#
);
let parsed: EmbeddingInput = serde_json::from_str(&json).expect("deserialization should not fail");
assert_eq!(parsed, input);
}
#[test]
fn image_bytes_encode_as_data_url() {
let part = EmbeddingContentPart::image_bytes(b"image", Some("image/png"));
assert_eq!(
part,
EmbeddingContentPart::ImageBase64 {
image_base64: "data:image/png;base64,aW1hZ2U=".into(),
}
);
}
fn embedding_body(embedding_json: &str) -> String {
format!(r#"{{"object":"embedding","index":0,"embedding":{embedding_json}}}"#)
}
#[test]
fn base64_embedding_round_trips_bit_exact() {
let src: [f32; 5] = [1.0, -2.5, 12.375, 0.0, f32::MIN_POSITIVE];
let mut bytes = Vec::with_capacity(src.len() * 4);
for v in src {
bytes.extend_from_slice(&v.to_le_bytes());
}
let encoded = base64::engine::general_purpose::STANDARD.encode(&bytes);
let body = embedding_body(&format!("{encoded:?}"));
let obj: EmbeddingObject = serde_json::from_str(&body).expect("base64 embedding should deserialize");
assert_eq!(obj.embedding, src, "decoded floats must match source bit-exactly");
}
#[test]
fn openai_base64_anchor_decodes_to_one() {
let body = embedding_body(r#""AACAPw==""#);
let obj: EmbeddingObject = serde_json::from_str(&body).expect("anchor base64 embedding should deserialize");
assert_eq!(obj.embedding, vec![1.0_f32]);
}
#[test]
fn float_array_body_still_parses() {
let body = embedding_body("[0.1,0.2,0.3]");
let obj: EmbeddingObject = serde_json::from_str(&body).expect("float array embedding should deserialize");
assert_eq!(obj.embedding, vec![0.1_f32, 0.2, 0.3]);
}
#[test]
fn odd_length_base64_errors_with_multiple_of_four_message() {
let encoded = base64::engine::general_purpose::STANDARD.encode(b"abcdef");
let body = embedding_body(&format!("{encoded:?}"));
let err = serde_json::from_str::<EmbeddingObject>(&body).expect_err("6-byte base64 payload must error");
assert!(
err.to_string().contains("not a multiple of 4"),
"expected multiple-of-4 message, got: {err}"
);
}
#[test]
fn invalid_base64_errors() {
let body = embedding_body(r#""not valid base64!!!""#);
let err = serde_json::from_str::<EmbeddingObject>(&body).expect_err("non-base64 string must error");
assert!(
err.to_string().contains("invalid base64 embedding"),
"expected invalid base64 message, got: {err}"
);
}
#[test]
fn serialization_stays_a_json_array() {
let body = embedding_body(r#""AACAPw==""#);
let obj: EmbeddingObject = serde_json::from_str(&body).expect("anchor base64 embedding should deserialize");
let serialized = serde_json::to_value(&obj).expect("serialize back to JSON");
assert!(
serialized["embedding"].is_array(),
"expected embedding field to serialize as an array, got: {serialized}"
);
}
}