1use std::sync::OnceLock;
7use std::time::Duration;
8
9use reqwest::header::{HeaderMap, HeaderValue, AUTHORIZATION, CONTENT_TYPE};
10use reqwest::{Client, StatusCode};
11use serde::{Deserialize, Serialize};
12use tokio::runtime::Runtime;
13
14use super::{Embedding, EmbeddingProvider, OPENAI_EMBEDDING_DIM};
15use crate::error::{CtxError, Result};
16
17const OPENAI_API_URL: &str = "https://api.openai.com/v1/embeddings";
19
20const REQUEST_TIMEOUT_SECS: u64 = 30;
22
23const CONNECT_TIMEOUT_SECS: u64 = 10;
25
26const MAX_RETRIES: u32 = 3;
28
29const RETRY_BASE_DELAY_MS: u64 = 1000;
31
32static GLOBAL_RUNTIME: OnceLock<Runtime> = OnceLock::new();
35
36fn get_or_create_runtime() -> &'static Runtime {
37 GLOBAL_RUNTIME.get_or_init(|| {
38 Runtime::new().expect("Failed to create global tokio runtime for OpenAI provider")
39 })
40}
41
42#[derive(Serialize)]
44struct EmbeddingRequest<'a> {
45 input: Vec<&'a str>,
46 model: &'a str,
47 encoding_format: &'a str,
48}
49
50#[derive(Deserialize)]
52struct EmbeddingResponse {
53 data: Option<Vec<EmbeddingData>>,
54 error: Option<ApiError>,
55}
56
57#[derive(Deserialize)]
59struct EmbeddingData {
60 embedding: Vec<f32>,
61 #[allow(dead_code)]
62 index: usize,
63}
64
65#[derive(Deserialize)]
67struct ApiError {
68 message: String,
69 #[allow(dead_code)]
70 r#type: Option<String>,
71 #[allow(dead_code)]
72 code: Option<String>,
73}
74
75pub struct OpenAIProvider {
77 client: Client,
78 model: String,
79}
80
81impl OpenAIProvider {
82 pub fn new(api_key: impl Into<String>) -> Result<Self> {
84 let api_key = api_key.into();
85
86 let mut headers = HeaderMap::new();
88 headers.insert(
89 AUTHORIZATION,
90 HeaderValue::from_str(&format!("Bearer {}", api_key))
91 .map_err(|e| CtxError::embedding(format!("Invalid API key format: {}", e)))?,
92 );
93 headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
94
95 let client = Client::builder()
97 .timeout(Duration::from_secs(REQUEST_TIMEOUT_SECS))
98 .connect_timeout(Duration::from_secs(CONNECT_TIMEOUT_SECS))
99 .default_headers(headers)
100 .build()
101 .map_err(|e| CtxError::embedding(format!("Failed to build HTTP client: {}", e)))?;
102
103 Ok(Self {
104 client,
105 model: "text-embedding-3-small".to_string(),
106 })
107 }
108
109 pub fn from_env() -> Result<Self> {
111 let api_key = std::env::var("OPENAI_API_KEY").map_err(|_| CtxError::InvalidApiKey)?;
112 Self::new(api_key)
113 }
114
115 #[allow(dead_code)]
117 pub fn with_model(mut self, model: impl Into<String>) -> Self {
118 self.model = model.into();
119 self
120 }
121
122 fn request(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
126 if tokio::runtime::Handle::try_current().is_ok() {
128 return Err(CtxError::embedding(
131 "Cannot call sync embed() from async context. Use embed_async() instead.",
132 ));
133 }
134
135 let runtime = get_or_create_runtime();
136 runtime.block_on(self.request_async(texts))
137 }
138
139 pub async fn request_async(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
142 let request_body = EmbeddingRequest {
143 input: texts.to_vec(),
144 model: &self.model,
145 encoding_format: "float",
146 };
147
148 let mut last_error = None;
149
150 for attempt in 0..MAX_RETRIES {
151 match self.send_request(&request_body).await {
152 Ok(embeddings) => return Ok(embeddings),
153 Err(e) => {
154 let should_retry = matches!(
156 &e,
157 CtxError::RateLimited(_) | CtxError::Embedding(_)
158 ) || matches!(&e, CtxError::Embedding(msg) if msg.contains("server error"));
159
160 if should_retry && attempt < MAX_RETRIES - 1 {
161 let delay = RETRY_BASE_DELAY_MS * (1 << attempt);
163 tokio::time::sleep(Duration::from_millis(delay)).await;
164 last_error = Some(e);
165 continue;
166 }
167 return Err(e);
168 }
169 }
170 }
171
172 Err(last_error.unwrap_or_else(|| CtxError::embedding("Max retries exceeded")))
173 }
174
175 #[allow(dead_code)] pub async fn embed_async(&self, text: &str) -> Result<Embedding> {
178 let mut results = self.request_async(&[text]).await?;
179 results
180 .pop()
181 .ok_or_else(|| CtxError::embedding("Empty response"))
182 }
183
184 #[allow(dead_code)] pub async fn embed_batch_async(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
187 const BATCH_SIZE: usize = 100;
188
189 let mut all_embeddings = Vec::with_capacity(texts.len());
190
191 for chunk in texts.chunks(BATCH_SIZE) {
192 let embeddings = self.request_async(chunk).await?;
193 all_embeddings.extend(embeddings);
194 }
195
196 Ok(all_embeddings)
197 }
198
199 async fn send_request(&self, request_body: &EmbeddingRequest<'_>) -> Result<Vec<Embedding>> {
201 let response = self
202 .client
203 .post(OPENAI_API_URL)
204 .json(request_body)
205 .send()
206 .await
207 .map_err(|e| {
208 if e.is_timeout() {
209 CtxError::embedding(format!("Request timed out: {}", e))
210 } else if e.is_connect() {
211 CtxError::embedding(format!("Connection failed: {}", e))
212 } else {
213 CtxError::embedding(e.to_string())
214 }
215 })?;
216
217 let status = response.status();
218
219 match status {
221 StatusCode::OK => {
222 let api_response: EmbeddingResponse = response
224 .json()
225 .await
226 .map_err(|e| CtxError::embedding(format!("Failed to parse response: {}", e)))?;
227
228 self.parse_response(api_response)
229 }
230 StatusCode::TOO_MANY_REQUESTS => {
231 let retry_after = response
233 .headers()
234 .get("retry-after")
235 .and_then(|v| v.to_str().ok())
236 .and_then(|s| s.parse::<u64>().ok())
237 .unwrap_or(60);
238 Err(CtxError::RateLimited(retry_after))
239 }
240 StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN => Err(CtxError::InvalidApiKey),
241 s if s.is_server_error() => {
242 let body = response.text().await.unwrap_or_default();
244 Err(CtxError::embedding(format!(
245 "server error ({}): {}",
246 status, body
247 )))
248 }
249 _ => {
250 let body = response.text().await.unwrap_or_default();
252
253 if let Ok(error_response) = serde_json::from_str::<EmbeddingResponse>(&body) {
255 if let Some(error) = error_response.error {
256 return Err(CtxError::embedding(error.message));
257 }
258 }
259
260 Err(CtxError::embedding(format!("HTTP {}: {}", status, body)))
261 }
262 }
263 }
264
265 fn parse_response(&self, response: EmbeddingResponse) -> Result<Vec<Embedding>> {
267 if let Some(error) = response.error {
269 if error.message.contains("rate limit") {
270 return Err(CtxError::RateLimited(60));
271 }
272 if error.message.contains("invalid api key")
273 || error.message.contains("Incorrect API key")
274 {
275 return Err(CtxError::InvalidApiKey);
276 }
277 return Err(CtxError::embedding(error.message));
278 }
279
280 let data = response
282 .data
283 .ok_or_else(|| CtxError::embedding("No data in response"))?;
284
285 let mut embeddings = Vec::with_capacity(data.len());
286
287 for item in data {
288 let vector = item.embedding;
289
290 if vector.len() != OPENAI_EMBEDDING_DIM {
291 return Err(CtxError::DimensionMismatch {
292 expected: OPENAI_EMBEDDING_DIM,
293 actual: vector.len(),
294 });
295 }
296
297 embeddings.push(Embedding::new(vector));
298 }
299
300 Ok(embeddings)
301 }
302}
303
304impl EmbeddingProvider for OpenAIProvider {
305 fn name(&self) -> &str {
306 "openai"
307 }
308
309 fn dimension(&self) -> usize {
310 OPENAI_EMBEDDING_DIM
311 }
312
313 fn embed(&self, text: &str) -> Result<Embedding> {
314 let mut results = self.request(&[text])?;
315 results
316 .pop()
317 .ok_or_else(|| CtxError::embedding("Empty response"))
318 }
319
320 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
321 const BATCH_SIZE: usize = 100;
324
325 let mut all_embeddings = Vec::with_capacity(texts.len());
326
327 for chunk in texts.chunks(BATCH_SIZE) {
328 let embeddings = self.request(chunk)?;
329 all_embeddings.extend(embeddings);
330 }
331
332 Ok(all_embeddings)
333 }
334}
335
336#[cfg(test)]
337mod tests {
338 use super::*;
339
340 #[test]
341 fn test_provider_creation() {
342 let result = OpenAIProvider::new("test-key-12345");
344 assert!(result.is_ok());
345
346 let provider = result.unwrap();
347 assert_eq!(provider.name(), "openai");
348 assert_eq!(provider.dimension(), OPENAI_EMBEDDING_DIM);
349 }
350
351 #[test]
352 fn test_parse_response_success() {
353 let provider = OpenAIProvider::new("test-key").unwrap();
354
355 let response = EmbeddingResponse {
356 data: Some(vec![EmbeddingData {
357 embedding: vec![0.0; OPENAI_EMBEDDING_DIM],
358 index: 0,
359 }]),
360 error: None,
361 };
362
363 let result = provider.parse_response(response);
364 assert!(result.is_ok());
365 let embeddings = result.unwrap();
366 assert_eq!(embeddings.len(), 1);
367 assert_eq!(embeddings[0].vector.len(), OPENAI_EMBEDDING_DIM);
368 }
369
370 #[test]
371 fn test_parse_response_error() {
372 let provider = OpenAIProvider::new("test-key").unwrap();
373
374 let response = EmbeddingResponse {
375 data: None,
376 error: Some(ApiError {
377 message: "Incorrect API key provided".to_string(),
378 r#type: None,
379 code: None,
380 }),
381 };
382
383 let result = provider.parse_response(response);
384 assert!(result.is_err());
385 assert!(matches!(result.unwrap_err(), CtxError::InvalidApiKey));
386 }
387
388 #[test]
389 fn test_parse_response_rate_limited() {
390 let provider = OpenAIProvider::new("test-key").unwrap();
391
392 let response = EmbeddingResponse {
393 data: None,
394 error: Some(ApiError {
395 message: "rate limit exceeded".to_string(),
396 r#type: None,
397 code: None,
398 }),
399 };
400
401 let result = provider.parse_response(response);
402 assert!(result.is_err());
403 assert!(matches!(result.unwrap_err(), CtxError::RateLimited(_)));
404 }
405
406 #[test]
407 fn test_parse_response_dimension_mismatch() {
408 let provider = OpenAIProvider::new("test-key").unwrap();
409
410 let response = EmbeddingResponse {
411 data: Some(vec![EmbeddingData {
412 embedding: vec![0.0; 100], index: 0,
414 }]),
415 error: None,
416 };
417
418 let result = provider.parse_response(response);
419 assert!(result.is_err());
420 assert!(matches!(
421 result.unwrap_err(),
422 CtxError::DimensionMismatch {
423 expected: 1536,
424 actual: 100
425 }
426 ));
427 }
428
429 #[test]
431 #[ignore]
432 fn test_embed_real() {
433 let provider = OpenAIProvider::from_env().expect("OPENAI_API_KEY not set");
434 let embedding = provider.embed("Hello, world!").expect("Embedding failed");
435 assert_eq!(embedding.dim(), OPENAI_EMBEDDING_DIM);
436 }
437
438 #[test]
439 #[ignore]
440 fn test_embed_batch_real() {
441 let provider = OpenAIProvider::from_env().expect("OPENAI_API_KEY not set");
442 let texts = vec!["Hello", "World", "Test"];
443 let embeddings = provider
444 .embed_batch(&texts)
445 .expect("Batch embedding failed");
446 assert_eq!(embeddings.len(), 3);
447 for emb in &embeddings {
448 assert_eq!(emb.dim(), OPENAI_EMBEDDING_DIM);
449 }
450 }
451
452 #[tokio::test]
453 #[ignore]
454 async fn test_embed_async_real() {
455 let provider = OpenAIProvider::from_env().expect("OPENAI_API_KEY not set");
456 let embedding = provider
457 .embed_async("Hello, world!")
458 .await
459 .expect("Async embedding failed");
460 assert_eq!(embedding.dim(), OPENAI_EMBEDDING_DIM);
461 }
462}