Skip to main content

llm_kernel/embedding/
openai.rs

1//! OpenAI text-embedding provider (sync, via ureq).
2
3use 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
19/// OpenAI embedding provider.
20///
21/// Uses `text-embedding-3-small` (1536-dim) by default.
22/// Swap model to `text-embedding-3-large` (3072-dim) for higher accuracy.
23///
24/// `api_key` is not exposed via `Debug` to prevent accidental logging.
25pub struct OpenAIEmbeddingClient {
26    api_key: String,
27    model: String,
28    dim: usize,
29}
30
31impl OpenAIEmbeddingClient {
32    /// Create with `text-embedding-3-small` (1536 dimensions).
33    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    /// Create with `text-embedding-3-large` (3072 dimensions).
42    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    /// Create with an explicit model name and embedding dimension.
51    ///
52    /// Use this for legacy models (`text-embedding-ada-002`), reduced-dimension
53    /// variants, or any future model not covered by [`new_small`](Self::new_small)
54    /// and [`new_large`](Self::new_large).
55    ///
56    /// # Panics
57    ///
58    /// Panics if `dim` is zero.
59    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    /// Create from environment variable `OPENAI_API_KEY`.
73    ///
74    /// Always uses `text-embedding-3-small` (1536-dim). For a different model
75    /// use [`new_with_model`](Self::new_with_model) after reading the key manually.
76    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    /// Build the `/v1/embeddings` request body.
83    ///
84    /// For Matryoshka-capable models (`text-embedding-3-*`) the configured
85    /// [`dim`](crate::embedding::EmbeddingProvider::dim) is sent as the
86    /// `dimensions` parameter, so the API returns an already-shortened and
87    /// renormalized vector. Without it the API always returns the model's
88    /// native width (1536 / 3072) and a shorter configured `dim` would
89    /// silently disagree with the vectors actually emitted.
90    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
99/// Whether `model` accepts the `dimensions` request parameter (Matryoshka
100/// shortening). First-generation models such as `text-embedding-ada-002`
101/// reject the field, so it is only sent for the v3 family.
102fn 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        // The OpenAI API does not guarantee that `data` is returned in input
175        // order; sort by `index` before zipping with `texts`.
176        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        // SAFETY: `remove_var` is unsafe in Rust 2024 because env mutation is
221        // not thread-safe. Tests run in their own process and this binary does
222        // not spawn threads that read OPENAI_API_KEY concurrently.
223        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        // Simulate an out-of-order API response (index 1 before index 0).
239        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        // ada-002 rejects `dimensions` — the field must not be sent.
272        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        // Each Korean char is 3 bytes; byte-slicing at 64 would panic.
315        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}