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
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}