1use 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
16pub struct EmbeddingProvider {
18 client: HttpClient,
19 model: String,
20 dimensions: Option<usize>,
21 policy: Option<Policy>,
22}
23
24impl EmbeddingProvider {
25 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 #[must_use]
43 pub fn with_policy(mut self, policy: Policy) -> Self {
44 self.policy = Some(policy);
45 self
46 }
47}
48
49pub fn embedding_provider(config: &Config) -> AppResult<EmbeddingProvider> {
51 EmbeddingProvider::new(config)
52}
53
54pub 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}