use serde::{Deserialize, Serialize};
use crate::client::OpenAiClient;
use crate::error::OpenAiError;
pub struct Embeddings<'a> {
pub(crate) client: &'a OpenAiClient,
}
impl Embeddings<'_> {
pub async fn create(
&self,
request: &EmbeddingsRequest,
) -> Result<EmbeddingsResponse, OpenAiError> {
self.client.post_json("/embeddings", request).await
}
}
#[derive(Debug, Clone, Serialize)]
pub struct EmbeddingsRequest {
model: String,
input: EmbeddingInput,
#[serde(skip_serializing_if = "Option::is_none")]
dimensions: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
encoding_format: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
user: Option<String>,
}
impl EmbeddingsRequest {
pub fn new(model: impl Into<String>, input: impl Into<EmbeddingInput>) -> Self {
Self {
model: model.into(),
input: input.into(),
dimensions: None,
encoding_format: None,
user: None,
}
}
pub fn dimensions(mut self, dimensions: u32) -> Self {
self.dimensions = Some(dimensions);
self
}
pub fn encoding_format(mut self, encoding_format: impl Into<String>) -> Self {
self.encoding_format = Some(encoding_format.into());
self
}
pub fn user(mut self, user: impl Into<String>) -> Self {
self.user = Some(user.into());
self
}
}
#[derive(Debug, Clone, Serialize)]
#[serde(untagged)]
pub enum EmbeddingInput {
Text(String),
Texts(Vec<String>),
}
impl From<&str> for EmbeddingInput {
fn from(text: &str) -> Self {
EmbeddingInput::Text(text.to_string())
}
}
impl From<String> for EmbeddingInput {
fn from(text: String) -> Self {
EmbeddingInput::Text(text)
}
}
impl From<Vec<String>> for EmbeddingInput {
fn from(texts: Vec<String>) -> Self {
EmbeddingInput::Texts(texts)
}
}
impl From<Vec<&str>> for EmbeddingInput {
fn from(texts: Vec<&str>) -> Self {
EmbeddingInput::Texts(texts.into_iter().map(str::to_string).collect())
}
}
#[derive(Debug, Clone, Deserialize)]
pub struct EmbeddingsResponse {
pub object: Option<String>,
#[serde(default)]
pub data: Vec<Embedding>,
pub model: Option<String>,
pub usage: Option<EmbeddingsUsage>,
}
#[derive(Debug, Clone, Deserialize)]
pub struct Embedding {
pub object: Option<String>,
#[serde(default)]
pub index: u32,
pub embedding: Vec<f32>,
}
#[derive(Debug, Clone, Deserialize)]
pub struct EmbeddingsUsage {
#[serde(default)]
pub prompt_tokens: u64,
#[serde(default)]
pub total_tokens: u64,
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn serializes_single_and_batch_input() {
let single = EmbeddingsRequest::new("text-embedding-3-small", "hello").dimensions(256);
let value = serde_json::to_value(&single).unwrap();
assert_eq!(value["input"], "hello");
assert_eq!(value["dimensions"], 256);
let batch = EmbeddingsRequest::new("text-embedding-3-large", vec!["a", "b"]);
let value = serde_json::to_value(&batch).unwrap();
assert_eq!(value["input"], json!(["a", "b"]));
}
#[test]
fn deserializes_response() {
let body = json!({
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, -0.2]}],
"model": "text-embedding-3-small",
"usage": {"prompt_tokens": 2, "total_tokens": 2}
});
let response: EmbeddingsResponse = serde_json::from_value(body).unwrap();
assert_eq!(response.data[0].embedding, vec![0.1, -0.2]);
assert_eq!(response.usage.unwrap().total_tokens, 2);
}
}