openai_interface/embeddings/
request.rs1use serde::Serialize;
2use url::Url;
3
4use crate::{
5 errors::OapiError,
6 rest::post::{Post, PostNoStream},
7};
8
9#[derive(Debug, Serialize, Default, Clone)]
11pub struct EmbeddingRequest {
12 pub input: EmbeddingInput,
18 pub model: String,
20 #[serde(skip_serializing_if = "Option::is_none")]
24 pub dimensions: Option<u32>,
25 #[serde(skip_serializing_if = "Option::is_none")]
31 pub encoding_format: Option<EncodingFormat>,
32 #[serde(skip_serializing_if = "Option::is_none")]
35 pub user: Option<String>,
36 #[serde(flatten, skip_serializing_if = "Option::is_none")]
38 pub extra_body: Option<serde_json::Map<String, serde_json::Value>>,
39}
40
41#[derive(Debug, Serialize, Clone)]
43#[serde(untagged)]
44pub enum EmbeddingInput {
45 String(String),
47 StringArray(Vec<String>),
49 Tokens(Vec<u32>),
51 TokenArray(Vec<Vec<u32>>),
53}
54
55impl Default for EmbeddingInput {
56 fn default() -> Self {
57 Self::String(String::new())
58 }
59}
60
61#[derive(Debug, Serialize, Clone, Copy)]
63#[serde(rename_all = "lowercase")]
64pub enum EncodingFormat {
65 Float,
66 Base64,
67}
68
69impl EmbeddingRequest {
70 pub fn is_streaming(&self) -> bool {
71 false
72 }
73}
74
75impl Post for EmbeddingRequest {
76 fn is_streaming(&self) -> bool {
77 false
78 }
79
80 fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
84 let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
85 url.path_segments_mut()
86 .map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
87 .push("embeddings");
88
89 Ok(url.to_string())
90 }
91}
92
93impl PostNoStream for EmbeddingRequest {
94 type Response = super::response::CreateEmbeddingResponse;
95}
96
97#[cfg(test)]
98mod tests {
99 use super::*;
100
101 #[test]
103 fn string_input_serialization() {
104 let request = EmbeddingRequest {
105 input: EmbeddingInput::String("Hello".to_string()),
106 model: "text-embedding-v4".to_string(),
107 dimensions: Some(1024),
108 encoding_format: Some(EncodingFormat::Float),
109 ..Default::default()
110 };
111
112 let json = serde_json::to_string(&request).unwrap();
113 assert!(json.contains(r#""input":"Hello""#), "json: {json}");
114 assert!(
115 json.contains(r#""model":"text-embedding-v4""#),
116 "json: {json}"
117 );
118 assert!(json.contains(r#""dimensions":1024"#), "json: {json}");
119 assert!(
120 json.contains(r#""encoding_format":"float""#),
121 "json: {json}"
122 );
123 }
124
125 #[test]
127 fn array_input_serialization() {
128 let request = EmbeddingRequest {
129 input: EmbeddingInput::StringArray(vec!["Hello".to_string(), "World".to_string()]),
130 model: "text-embedding-3-small".to_string(),
131 ..Default::default()
132 };
133 let json = serde_json::to_string(&request).unwrap();
134 assert!(
135 json.contains(r#""input":["Hello","World"]"#),
136 "json: {json}"
137 );
138
139 let request = EmbeddingRequest {
140 input: EmbeddingInput::TokenArray(vec![vec![1234, 5678]]),
141 model: "text-embedding-3-small".to_string(),
142 ..Default::default()
143 };
144 let json = serde_json::to_string(&request).unwrap();
145 assert!(json.contains(r#""input":[[1234,5678]]"#), "json: {json}");
146 }
147
148 #[test]
150 fn flat_token_array_serialization() {
151 let request = EmbeddingRequest {
152 input: EmbeddingInput::Tokens(vec![1234, 5678]),
153 model: "text-embedding-3-small".to_string(),
154 ..Default::default()
155 };
156 let json = serde_json::to_string(&request).unwrap();
157 assert!(json.contains(r#""input":[1234,5678]"#), "json: {json}");
158 }
159
160 #[test]
161 fn test_build_url() {
162 let request = EmbeddingRequest::default();
163 let url = request.build_url("https://api.openai.com/v1/").unwrap();
164 assert_eq!(url, "https://api.openai.com/v1/embeddings");
165 }
166}