Skip to main content

ctx/embeddings/
openai.rs

1//! OpenAI embedding provider.
2//!
3//! Uses the OpenAI API to generate embeddings via text-embedding-3-small.
4//! Implements secure HTTP with reqwest, proper timeouts, and retry logic.
5
6use 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
17/// OpenAI API endpoint
18const OPENAI_API_URL: &str = "https://api.openai.com/v1/embeddings";
19
20/// Request timeout in seconds
21const REQUEST_TIMEOUT_SECS: u64 = 30;
22
23/// Connection timeout in seconds
24const CONNECT_TIMEOUT_SECS: u64 = 10;
25
26/// Maximum retry attempts for retryable errors
27const MAX_RETRIES: u32 = 3;
28
29/// Base delay for exponential backoff (in milliseconds)
30const RETRY_BASE_DELAY_MS: u64 = 1000;
31
32/// Global runtime for sync API when not already in an async context.
33/// This avoids creating a new runtime per provider instance.
34static 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/// OpenAI embedding request body.
43#[derive(Serialize)]
44struct EmbeddingRequest<'a> {
45    input: Vec<&'a str>,
46    model: &'a str,
47    encoding_format: &'a str,
48}
49
50/// OpenAI API response.
51#[derive(Deserialize)]
52struct EmbeddingResponse {
53    data: Option<Vec<EmbeddingData>>,
54    error: Option<ApiError>,
55}
56
57/// Individual embedding in the response.
58#[derive(Deserialize)]
59struct EmbeddingData {
60    embedding: Vec<f32>,
61    #[allow(dead_code)]
62    index: usize,
63}
64
65/// API error response.
66#[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
75/// OpenAI embedding provider configuration.
76pub struct OpenAIProvider {
77    client: Client,
78    model: String,
79}
80
81impl OpenAIProvider {
82    /// Create a new OpenAI provider with the given API key.
83    pub fn new(api_key: impl Into<String>) -> Result<Self> {
84        let api_key = api_key.into();
85
86        // Build headers with authorization
87        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        // Build the HTTP client with timeouts
96        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    /// Create a provider from the OPENAI_API_KEY environment variable.
110    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    /// Set the model to use for embeddings.
116    #[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    /// Make an HTTP request to the OpenAI API with retry logic (sync version).
123    /// Uses a global runtime to avoid creating a new runtime per call.
124    /// Safe to call from sync context; will panic if called from within an async runtime.
125    fn request(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
126        // Check if we're already in an async context
127        if tokio::runtime::Handle::try_current().is_ok() {
128            // We're in an async context - this would deadlock with block_on
129            // Return an error instead of panicking
130            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    /// Make an HTTP request to the OpenAI API with retry logic (async version).
140    /// Safe to call from within an async runtime.
141    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                    // Check if error is retryable
155                    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                        // Exponential backoff: 1s, 2s, 4s
162                        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    /// Async version of embed for use in async contexts (e.g., MCP server).
176    #[allow(dead_code)] // Public API for async contexts
177    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    /// Async version of embed_batch for use in async contexts.
185    #[allow(dead_code)] // Public API for async contexts
186    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    /// Send a single request to the OpenAI API.
200    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        // Handle HTTP status codes
220        match status {
221            StatusCode::OK => {
222                // Parse successful response
223                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                // Rate limited - extract retry-after if available
232                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                // 5xx errors - retryable
243                let body = response.text().await.unwrap_or_default();
244                Err(CtxError::embedding(format!(
245                    "server error ({}): {}",
246                    status, body
247                )))
248            }
249            _ => {
250                // Other client errors
251                let body = response.text().await.unwrap_or_default();
252
253                // Try to parse as API error
254                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    /// Parse the OpenAI API response.
266    fn parse_response(&self, response: EmbeddingResponse) -> Result<Vec<Embedding>> {
267        // Check for API errors
268        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        // Extract embeddings
281        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        // OpenAI supports batch embedding up to ~8000 tokens
322        // For simplicity, we'll batch in chunks of 100
323        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        // Note: This will fail without a valid API key format, but tests the builder
343        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], // Wrong dimension
413                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    // Integration test - requires OPENAI_API_KEY
430    #[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}