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