Skip to main content

rskit_llm_openai/
embedding.rs

1//! OpenAI-compatible embedding provider using rskit-httpclient.
2
3use async_trait::async_trait;
4use rskit_ai::semconv;
5use rskit_embedding::{EmbedInput, EmbedRequest, EmbedResponse, Embedding, Provider};
6use rskit_errors::{AppError, AppResult, ErrorCode};
7use rskit_httpclient::{Auth, HttpClient, HttpClientConfig, Request};
8use rskit_observability::set_span_attribute;
9use rskit_resilience::Policy;
10use serde::{Deserialize, Serialize};
11use tracing::{Instrument, debug};
12
13use super::PROVIDER_ID;
14use super::config::Config;
15
16/// OpenAI-compatible embedding provider backed by rskit-httpclient.
17pub struct EmbeddingProvider {
18    client: HttpClient,
19    model: String,
20    dimensions: Option<usize>,
21    policy: Option<Policy>,
22}
23
24impl EmbeddingProvider {
25    /// Create a new embedding provider from an `OpenAI` [`Config`].
26    pub fn new(cfg: &Config) -> AppResult<Self> {
27        let http_cfg = HttpClientConfig::new()
28            .with_base_url(&cfg.base_url)
29            .with_auth(Auth::bearer_secret(cfg.api_key.clone()));
30
31        let client = HttpClient::new(http_cfg)?;
32
33        Ok(Self {
34            client,
35            model: cfg.embedding_model.clone(),
36            dimensions: cfg.embedding_dimensions,
37            policy: None,
38        })
39    }
40
41    /// Inject a resilience policy for outbound embedding requests.
42    #[must_use]
43    pub fn with_policy(mut self, policy: Policy) -> Self {
44        self.policy = Some(policy);
45        self
46    }
47}
48
49#[derive(Serialize)]
50struct EmbeddingRequest {
51    model: String,
52    input: Vec<String>,
53    #[serde(skip_serializing_if = "Option::is_none")]
54    dimensions: Option<usize>,
55}
56
57#[derive(Deserialize)]
58struct EmbeddingResponse {
59    data: Vec<EmbeddingData>,
60    #[serde(default)]
61    usage: Option<EmbeddingUsage>,
62}
63
64#[derive(Deserialize)]
65struct EmbeddingData {
66    embedding: Vec<f32>,
67}
68
69#[derive(Deserialize)]
70struct EmbeddingUsage {
71    prompt_tokens: u64,
72    #[serde(default)]
73    total_tokens: u64,
74}
75
76#[async_trait]
77impl Provider for EmbeddingProvider {
78    async fn embed(&self, req: EmbedRequest) -> AppResult<EmbedResponse> {
79        let mut response_model = req.model.clone();
80        if response_model.name.is_empty() {
81            response_model.name.clone_from(&self.model);
82        }
83        let model = response_model.name.clone();
84
85        let span = embedding_span(&model, req.inputs.len());
86        async move {
87            let texts = req
88                .inputs
89                .iter()
90                .map(|input| match input {
91                    EmbedInput::Text(text) => Ok(text.clone()),
92                    _ => Err(AppError::new(
93                        ErrorCode::InvalidInput,
94                        "OpenAI embedding adapter currently accepts text inputs only",
95                    )),
96                })
97                .collect::<AppResult<Vec<_>>>()?;
98
99            if texts.is_empty() {
100                return Ok(EmbedResponse {
101                    embeddings: Vec::new(),
102                    model: response_model,
103                    usage: rskit_ai::Usage::default(),
104                });
105            }
106
107            let body = EmbeddingRequest {
108                model: model.clone(),
109                input: texts,
110                dimensions: self.dimensions,
111            };
112
113            debug!(model = %model, count = body.input.len(), "requesting embeddings");
114
115            let request = Request::post("/embeddings")
116                .json_body(&body)
117                .map_err(|e| AppError::internal(e).context("build embedding request"))?;
118
119            let policy = self.policy.clone();
120            let response = if let Some(policy) = policy {
121                let request = request.clone();
122                policy
123                    .execute(|| {
124                        let request = request.clone();
125                        async move {
126                            let resp = self.client.send(request).await?;
127                            if !resp.is_success() {
128                                let status = resp.status_u16();
129                                let text = resp.text_or_diagnostic();
130                                return Err(AppError::new(
131                                    ErrorCode::ExternalService,
132                                    format!("embedding API returned HTTP {status}"),
133                                )
134                                .with_detail("status", status.to_string())
135                                .with_detail("body", text));
136                            }
137                            Ok(resp)
138                        }
139                    })
140                    .await?
141            } else {
142                let resp = self.client.send(request).await?;
143                if !resp.is_success() {
144                    let status = resp.status_u16();
145                    let text = resp.text_or_diagnostic();
146                    return Err(AppError::new(
147                        ErrorCode::ExternalService,
148                        format!("embedding API returned HTTP {status}"),
149                    )
150                    .with_detail("status", status.to_string())
151                    .with_detail("body", text));
152                }
153                resp
154            };
155
156            let result: EmbeddingResponse = response
157                .json()
158                .map_err(|e| AppError::internal(e).context("parse embedding response"))?;
159
160            let usage = result
161                .usage
162                .map(|u| rskit_ai::Usage {
163                    input_tokens: u.prompt_tokens,
164                    output_tokens: u.total_tokens.saturating_sub(u.prompt_tokens),
165                    ..Default::default()
166                })
167                .unwrap_or_default();
168
169            Ok(EmbedResponse {
170                embeddings: result
171                    .data
172                    .into_iter()
173                    .enumerate()
174                    .map(|(index, data)| Embedding::new(data.embedding, index))
175                    .collect(),
176                model: response_model,
177                usage,
178            })
179        }
180        .instrument(span)
181        .await
182    }
183
184    async fn embed_batch(&self, reqs: Vec<EmbedRequest>) -> AppResult<Vec<EmbedResponse>> {
185        let mut responses = Vec::with_capacity(reqs.len());
186        for req in reqs {
187            responses.push(self.embed(req).await?);
188        }
189        Ok(responses)
190    }
191}
192
193fn embedding_span(model: &str, input_count: usize) -> tracing::Span {
194    let span = tracing::info_span!(
195        "embedding.embed",
196        "gen_ai.system" = PROVIDER_ID,
197        "gen_ai.operation.name" = semconv::Operation::Embedding.as_str(),
198        "gen_ai.request.model" = %model,
199        "embedding.input_count" = input_count,
200    );
201    set_span_attribute(&span, semconv::SYSTEM, PROVIDER_ID);
202    set_span_attribute(
203        &span,
204        semconv::OPERATION_NAME,
205        semconv::Operation::Embedding.as_str(),
206    );
207    set_span_attribute(&span, semconv::REQUEST_MODEL, model);
208    span
209}
210
211impl rskit_provider::Provider for EmbeddingProvider {
212    fn name(&self) -> &'static str {
213        "openai_embedding"
214    }
215}
216
217#[async_trait]
218impl rskit_provider::RequestResponse<EmbedRequest, EmbedResponse> for EmbeddingProvider {
219    async fn execute(&self, input: EmbedRequest) -> AppResult<EmbedResponse> {
220        self.embed(input).await
221    }
222}
223
224#[cfg(test)]
225mod tests {
226    use super::*;
227    use rskit_embedding::EmbedAsset;
228    use rskit_resilience::{ConstantBackoff, RetryPolicy};
229    use std::sync::Arc;
230    use std::sync::atomic::{AtomicUsize, Ordering};
231    use std::time::Duration;
232    use tokio::io::{AsyncReadExt, AsyncWriteExt};
233    use tokio::net::TcpListener;
234
235    #[test]
236    fn provider_constructs_with_config() {
237        let cfg = Config {
238            api_key: rskit_util::SecretString::new("sk-test"),
239            base_url: "https://api.openai.com/v1".into(),
240            model: "gpt-4o".into(),
241            embedding_model: "text-embedding-3-small".into(),
242            embedding_dimensions: Some(1536),
243        };
244        let provider = EmbeddingProvider::new(&cfg).unwrap();
245        assert_eq!(provider.dimensions, Some(1536));
246    }
247
248    #[test]
249    fn embedding_request_omits_dimensions_when_unset() {
250        let body = EmbeddingRequest {
251            model: "text-embedding-ada-002".into(),
252            input: vec!["hello".into()],
253            dimensions: None,
254        };
255
256        let json = serde_json::to_value(body).unwrap();
257        assert!(json.get("dimensions").is_none());
258    }
259
260    #[test]
261    fn embedding_request_includes_dimensions_when_set() {
262        let body = EmbeddingRequest {
263            model: "text-embedding-3-small".into(),
264            input: vec!["hello".into()],
265            dimensions: Some(768),
266        };
267
268        let json = serde_json::to_value(body).unwrap();
269        assert_eq!(json["dimensions"], 768);
270    }
271
272    #[tokio::test]
273    async fn embed_returns_empty_response_without_http_for_empty_inputs() {
274        let provider = EmbeddingProvider::new(&config(None)).unwrap();
275
276        let response = provider.embed(request(Vec::new())).await.unwrap();
277
278        assert!(response.embeddings.is_empty());
279        assert_eq!(response.model.name, "text-embedding-3-small");
280    }
281
282    #[tokio::test]
283    async fn embed_rejects_non_text_inputs_before_http() {
284        let provider = EmbeddingProvider::new(&config(None)).unwrap();
285
286        let err = provider
287            .embed(request(vec![EmbedInput::Image(EmbedAsset::Url(
288                "https://example.test/image.png".into(),
289            ))]))
290            .await
291            .unwrap_err();
292
293        assert_eq!(err.code(), ErrorCode::InvalidInput);
294    }
295
296    #[tokio::test]
297    async fn embed_batch_and_request_response_forward_to_embed() {
298        let (base_url, server) = spawn_response_server(vec![
299            (
300                200,
301                r#"{"data":[{"embedding":[0.1]}],"usage":{"prompt_tokens":1,"total_tokens":2}}"#,
302            ),
303            (
304                200,
305                r#"{"data":[{"embedding":[0.2]}],"usage":{"prompt_tokens":2,"total_tokens":3}}"#,
306            ),
307        ])
308        .await;
309        let provider = EmbeddingProvider::new(&config(Some(base_url))).unwrap();
310
311        assert_eq!(
312            rskit_provider::Provider::name(&provider),
313            "openai_embedding"
314        );
315        let via_trait = rskit_provider::RequestResponse::execute(
316            &provider,
317            request(vec![EmbedInput::Text("one".into())]),
318        )
319        .await
320        .unwrap();
321        let batch = provider
322            .embed_batch(vec![request(vec![EmbedInput::Text("two".into())])])
323            .await
324            .unwrap();
325
326        assert_eq!(via_trait.embeddings[0].vector, vec![0.1]);
327        assert_eq!(batch[0].embeddings[0].vector, vec![0.2]);
328        server.await.unwrap();
329    }
330
331    #[tokio::test]
332    async fn embed_maps_http_errors_without_policy() {
333        let (base_url, server) = spawn_response_server(vec![(503, "try later")]).await;
334        let provider = EmbeddingProvider::new(&config(Some(base_url))).unwrap();
335
336        let err = provider
337            .embed(request(vec![EmbedInput::Text("hello".into())]))
338            .await
339            .unwrap_err();
340
341        assert_eq!(err.code(), ErrorCode::ExternalService);
342        assert!(err.message().contains("embedding API returned HTTP 503"));
343        server.await.unwrap();
344    }
345
346    #[tokio::test]
347    async fn provider_retries_with_policy() {
348        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
349        let address = listener.local_addr().unwrap();
350        let attempts = Arc::new(AtomicUsize::new(0));
351        let attempts_in_server = attempts.clone();
352
353        let server = tokio::spawn(async move {
354            for _ in 0..2 {
355                let (mut socket, _) = listener.accept().await.unwrap();
356                let attempts = attempts_in_server.clone();
357                tokio::spawn(async move {
358                    let mut buffer = [0_u8; 2048];
359                    let _ = socket.read(&mut buffer).await;
360                    let attempt = attempts.fetch_add(1, Ordering::SeqCst);
361                    if attempt == 0 {
362                        socket
363                            .write_all(
364                                b"HTTP/1.1 500 Internal Server Error\r\ncontent-length: 12\r\nconnection: close\r\n\r\nretry later",
365                            )
366                            .await
367                            .unwrap();
368                    } else {
369                        let body = r#"{"data":[{"embedding":[0.1,0.2,0.3]}],"usage":{"prompt_tokens":2,"total_tokens":2}}"#;
370                        let response = format!(
371                            "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
372                            body.len(),
373                            body
374                        );
375                        socket.write_all(response.as_bytes()).await.unwrap();
376                    }
377                    socket.shutdown().await.unwrap();
378                });
379            }
380        });
381
382        let cfg = Config {
383            api_key: rskit_util::SecretString::new("sk-test"),
384            base_url: format!("http://{address}"),
385            model: "gpt-4o".into(),
386            embedding_model: "text-embedding-3-small".into(),
387            embedding_dimensions: Some(3),
388        };
389        let provider = EmbeddingProvider::new(&cfg).unwrap().with_policy(
390            Policy::new().with_retry(
391                RetryPolicy::fast()
392                    .with_constant_backoff(ConstantBackoff::new(Duration::from_millis(1)))
393                    .with_jitter(false),
394            ),
395        );
396        let response = provider
397            .embed(EmbedRequest {
398                model: rskit_ai::Model {
399                    name: String::new(),
400                    provider: rskit_ai::Provider::OpenAI,
401                    version: None,
402                    capabilities: rskit_ai::Capabilities::default(),
403                },
404                inputs: vec![EmbedInput::Text("retry".into())],
405                options: rskit_embedding::EmbeddingOptions::default(),
406            })
407            .await
408            .unwrap();
409
410        assert_eq!(attempts.load(Ordering::SeqCst), 2);
411        assert_eq!(response.embeddings.len(), 1);
412        server.await.unwrap();
413    }
414
415    fn config(base_url: Option<String>) -> Config {
416        Config {
417            api_key: rskit_util::SecretString::new("sk-test"),
418            base_url: base_url.unwrap_or_else(|| "https://api.openai.com/v1".into()),
419            model: "gpt-4o".into(),
420            embedding_model: "text-embedding-3-small".into(),
421            embedding_dimensions: Some(3),
422        }
423    }
424
425    fn request(inputs: Vec<EmbedInput>) -> EmbedRequest {
426        EmbedRequest {
427            model: rskit_ai::Model {
428                name: String::new(),
429                provider: rskit_ai::Provider::OpenAI,
430                version: None,
431                capabilities: rskit_ai::Capabilities::default(),
432            },
433            inputs,
434            options: rskit_embedding::EmbeddingOptions::default(),
435        }
436    }
437
438    async fn spawn_response_server(
439        responses: Vec<(u16, &'static str)>,
440    ) -> (String, tokio::task::JoinHandle<()>) {
441        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
442        let address = listener.local_addr().unwrap();
443        let server = tokio::spawn(async move {
444            for (status, body) in responses {
445                let (mut socket, _) = listener.accept().await.unwrap();
446                tokio::spawn(async move {
447                    let mut buffer = [0_u8; 2048];
448                    let _ = socket.read(&mut buffer).await;
449                    let reason = if status >= 400 { "Error" } else { "OK" };
450                    let response = format!(
451                        "HTTP/1.1 {status} {reason}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
452                        body.len()
453                    );
454                    socket.write_all(response.as_bytes()).await.unwrap();
455                    socket.shutdown().await.unwrap();
456                });
457            }
458        });
459        (format!("http://{address}"), server)
460    }
461}