1use serde::Deserialize;
4
5use crate::embedding::types::{EmbeddingProvider, EmbeddingResult};
6use crate::error::{KernelError, Result};
7
8#[derive(Deserialize)]
9struct EmbeddingData {
10 embedding: Vec<f32>,
11 index: usize,
12}
13
14#[derive(Deserialize)]
15struct EmbeddingResponse {
16 data: Vec<EmbeddingData>,
17}
18
19pub struct OpenAIEmbeddingClient {
26 api_key: String,
27 model: String,
28 dim: usize,
29}
30
31impl OpenAIEmbeddingClient {
32 pub fn new_small(api_key: impl Into<String>) -> Self {
34 Self {
35 api_key: api_key.into(),
36 model: "text-embedding-3-small".into(),
37 dim: 1536,
38 }
39 }
40
41 pub fn new_large(api_key: impl Into<String>) -> Self {
43 Self {
44 api_key: api_key.into(),
45 model: "text-embedding-3-large".into(),
46 dim: 3072,
47 }
48 }
49
50 pub fn new_with_model(
60 api_key: impl Into<String>,
61 model: impl Into<String>,
62 dim: usize,
63 ) -> Self {
64 assert!(dim > 0, "dim must be non-zero");
65 Self {
66 api_key: api_key.into(),
67 model: model.into(),
68 dim,
69 }
70 }
71
72 pub fn from_env() -> Result<Self> {
77 let key = std::env::var("OPENAI_API_KEY")
78 .map_err(|_| KernelError::Embedding("OPENAI_API_KEY not set".into()))?;
79 Ok(Self::new_small(key))
80 }
81
82 fn request_body(&self, input: serde_json::Value) -> serde_json::Value {
91 let mut body = serde_json::json!({ "model": self.model, "input": input });
92 if supports_dimensions(&self.model) {
93 body["dimensions"] = serde_json::json!(self.dim);
94 }
95 body
96 }
97}
98
99fn supports_dimensions(model: &str) -> bool {
103 model.starts_with("text-embedding-3-")
104}
105
106use super::types::text_preview;
107
108impl EmbeddingProvider for OpenAIEmbeddingClient {
109 fn dim(&self) -> usize {
110 self.dim
111 }
112
113 fn name(&self) -> &str {
114 &self.model
115 }
116
117 fn embed(&self, text: &str) -> Result<EmbeddingResult> {
118 let config = ureq::config::Config::builder()
119 .timeout_global(Some(std::time::Duration::from_secs(30)))
120 .build();
121 let agent = ureq::Agent::new_with_config(config);
122
123 let body = self.request_body(serde_json::json!(text));
124
125 let mut resp = agent
126 .post("https://api.openai.com/v1/embeddings")
127 .header("Authorization", format!("Bearer {}", self.api_key))
128 .header("Content-Type", "application/json")
129 .send_json(body)
130 .map_err(KernelError::embedding)?;
131
132 let payload: EmbeddingResponse = resp
133 .body_mut()
134 .read_json()
135 .map_err(KernelError::embedding)?;
136
137 let vector = payload
138 .data
139 .into_iter()
140 .next()
141 .ok_or_else(|| KernelError::Embedding("empty embedding response".into()))?
142 .embedding;
143
144 Ok(EmbeddingResult {
145 vector,
146 text_preview: text_preview(text),
147 })
148 }
149
150 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<EmbeddingResult>> {
151 if texts.is_empty() {
152 return Ok(vec![]);
153 }
154
155 let config = ureq::config::Config::builder()
156 .timeout_global(Some(std::time::Duration::from_secs(60)))
157 .build();
158 let agent = ureq::Agent::new_with_config(config);
159
160 let body = self.request_body(serde_json::json!(texts));
161
162 let mut resp = agent
163 .post("https://api.openai.com/v1/embeddings")
164 .header("Authorization", format!("Bearer {}", self.api_key))
165 .header("Content-Type", "application/json")
166 .send_json(body)
167 .map_err(KernelError::embedding)?;
168
169 let payload: EmbeddingResponse = resp
170 .body_mut()
171 .read_json()
172 .map_err(KernelError::embedding)?;
173
174 let mut data = payload.data;
177 data.sort_unstable_by_key(|d| d.index);
178
179 if data.len() != texts.len() {
180 return Err(KernelError::Embedding(format!(
181 "API returned {} embeddings for {} inputs",
182 data.len(),
183 texts.len()
184 )));
185 }
186
187 let results = data
188 .into_iter()
189 .zip(texts.iter())
190 .map(|(item, &text)| EmbeddingResult {
191 vector: item.embedding,
192 text_preview: text_preview(text),
193 })
194 .collect();
195
196 Ok(results)
197 }
198}
199
200#[cfg(test)]
201mod tests {
202 use super::*;
203
204 #[test]
205 fn small_client_has_correct_dim() {
206 let client = OpenAIEmbeddingClient::new_small("test-key");
207 assert_eq!(client.dim(), 1536);
208 assert_eq!(client.name(), "text-embedding-3-small");
209 }
210
211 #[test]
212 fn large_client_has_correct_dim() {
213 let client = OpenAIEmbeddingClient::new_large("test-key");
214 assert_eq!(client.dim(), 3072);
215 assert_eq!(client.name(), "text-embedding-3-large");
216 }
217
218 #[test]
219 fn from_env_fails_without_key() {
220 unsafe { std::env::remove_var("OPENAI_API_KEY") };
224 assert!(OpenAIEmbeddingClient::from_env().is_err());
225 }
226
227 #[test]
228 fn parse_embedding_response() {
229 let raw = r#"{"object":"list","data":[{"object":"embedding","embedding":[0.1,-0.2,0.3],"index":0}],"model":"text-embedding-3-small","usage":{"prompt_tokens":5,"total_tokens":5}}"#;
230 let payload: EmbeddingResponse = serde_json::from_str(raw).unwrap();
231 assert_eq!(payload.data.len(), 1);
232 assert_eq!(payload.data[0].embedding, vec![0.1f32, -0.2, 0.3]);
233 assert_eq!(payload.data[0].index, 0);
234 }
235
236 #[test]
237 fn embed_batch_reorders_by_index() {
238 let raw = r#"{"object":"list","data":[{"object":"embedding","embedding":[0.2],"index":1},{"object":"embedding","embedding":[0.1],"index":0}],"model":"text-embedding-3-small","usage":{"prompt_tokens":2,"total_tokens":2}}"#;
240 let mut payload: EmbeddingResponse = serde_json::from_str(raw).unwrap();
241 payload.data.sort_unstable_by_key(|d| d.index);
242 assert_eq!(payload.data[0].embedding, vec![0.1f32]);
243 assert_eq!(payload.data[1].embedding, vec![0.2f32]);
244 }
245
246 #[test]
247 fn new_with_model_sets_name_and_dim() {
248 let client = OpenAIEmbeddingClient::new_with_model("key", "text-embedding-ada-002", 1536);
249 assert_eq!(client.dim(), 1536);
250 assert_eq!(client.name(), "text-embedding-ada-002");
251 }
252
253 #[test]
254 fn new_with_model_custom_dim() {
255 let client = OpenAIEmbeddingClient::new_with_model("key", "text-embedding-3-small", 512);
256 assert_eq!(client.dim(), 512);
257 assert_eq!(client.name(), "text-embedding-3-small");
258 }
259
260 #[test]
261 fn request_body_sends_dimensions_for_v3() {
262 let client = OpenAIEmbeddingClient::new_with_model("key", "text-embedding-3-small", 512);
263 let body = client.request_body(serde_json::json!("hello"));
264 assert_eq!(body["model"], "text-embedding-3-small");
265 assert_eq!(body["input"], "hello");
266 assert_eq!(body["dimensions"], serde_json::json!(512));
267 }
268
269 #[test]
270 fn request_body_omits_dimensions_for_first_gen() {
271 let client = OpenAIEmbeddingClient::new_with_model("key", "text-embedding-ada-002", 1536);
273 let body = client.request_body(serde_json::json!("hello"));
274 assert!(body.get("dimensions").is_none());
275 }
276
277 #[test]
278 fn request_body_carries_batch_input() {
279 let client = OpenAIEmbeddingClient::new_with_model("key", "text-embedding-3-large", 256);
280 let body = client.request_body(serde_json::json!(["a", "b"]));
281 assert_eq!(body["input"], serde_json::json!(["a", "b"]));
282 assert_eq!(body["dimensions"], serde_json::json!(256));
283 }
284
285 #[test]
286 fn supports_dimensions_matches_v3_family_only() {
287 assert!(supports_dimensions("text-embedding-3-small"));
288 assert!(supports_dimensions("text-embedding-3-large"));
289 assert!(!supports_dimensions("text-embedding-ada-002"));
290 assert!(!supports_dimensions("some-other-model"));
291 }
292
293 #[test]
294 #[should_panic(expected = "dim must be non-zero")]
295 fn new_with_model_zero_dim_panics() {
296 OpenAIEmbeddingClient::new_with_model("key", "text-embedding-ada-002", 0);
297 }
298
299 #[test]
300 fn preview_ascii_truncated() {
301 let long = "a".repeat(100);
302 let preview = text_preview(&long);
303 assert!(preview.ends_with('…'));
304 assert_eq!(preview.chars().filter(|&c| c != '…').count(), 64);
305 }
306
307 #[test]
308 fn preview_short_not_truncated() {
309 assert_eq!(text_preview("hello"), "hello");
310 }
311
312 #[test]
313 fn preview_multibyte_no_panic() {
314 let korean = "안녕하세요".repeat(20);
316 let preview = text_preview(&korean);
317 assert!(preview.ends_with('…'));
318 assert_eq!(preview.chars().filter(|&c| c != '…').count(), 64);
319 }
320}