use serde::Serialize;
use url::Url;
use crate::{
errors::OapiError,
rest::post::{Post, PostNoStream},
};
#[derive(Debug, Serialize, Default, Clone)]
pub struct EmbeddingRequest {
pub input: EmbeddingInput,
pub model: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub dimensions: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub encoding_format: Option<EncodingFormat>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
#[serde(flatten, skip_serializing_if = "Option::is_none")]
pub extra_body: Option<serde_json::Map<String, serde_json::Value>>,
}
#[derive(Debug, Serialize, Clone)]
#[serde(untagged)]
pub enum EmbeddingInput {
String(String),
StringArray(Vec<String>),
TokenArray(Vec<Vec<u32>>),
}
impl Default for EmbeddingInput {
fn default() -> Self {
Self::String(String::new())
}
}
#[derive(Debug, Serialize, Clone, Copy)]
#[serde(rename_all = "lowercase")]
pub enum EncodingFormat {
Float,
Base64,
}
impl EmbeddingRequest {
pub fn is_streaming(&self) -> bool {
false
}
}
impl Post for EmbeddingRequest {
fn is_streaming(&self) -> bool {
false
}
fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
url.path_segments_mut()
.map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
.push("embeddings");
Ok(url.to_string())
}
}
impl PostNoStream for EmbeddingRequest {
type Response = super::response::CreateEmbeddingResponse;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn string_input_serialization() {
let request = EmbeddingRequest {
input: EmbeddingInput::String("Hello".to_string()),
model: "text-embedding-v4".to_string(),
dimensions: Some(1024),
encoding_format: Some(EncodingFormat::Float),
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(json.contains(r#""input":"Hello""#), "json: {json}");
assert!(
json.contains(r#""model":"text-embedding-v4""#),
"json: {json}"
);
assert!(json.contains(r#""dimensions":1024"#), "json: {json}");
assert!(
json.contains(r#""encoding_format":"float""#),
"json: {json}"
);
}
#[test]
fn array_input_serialization() {
let request = EmbeddingRequest {
input: EmbeddingInput::StringArray(vec!["Hello".to_string(), "World".to_string()]),
model: "text-embedding-3-small".to_string(),
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(
json.contains(r#""input":["Hello","World"]"#),
"json: {json}"
);
let request = EmbeddingRequest {
input: EmbeddingInput::TokenArray(vec![vec![1234, 5678]]),
model: "text-embedding-3-small".to_string(),
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(json.contains(r#""input":[[1234,5678]]"#), "json: {json}");
}
#[test]
fn test_build_url() {
let request = EmbeddingRequest::default();
let url = request.build_url("https://api.openai.com/v1/").unwrap();
assert_eq!(url, "https://api.openai.com/v1/embeddings");
}
}