1use std::{error::Error, fmt, future::Future, pin::Pin};
2
3use serde::{Deserialize, Serialize};
4
5use super::{EmbeddingProviderKind, RemoteEmbeddingConfig};
6
7const PROVIDER_ERROR_MESSAGE_LIMIT: usize = 240;
8
9pub type EmbeddingFuture<'a, T> =
10 Pin<Box<dyn Future<Output = Result<T, EmbeddingProviderError>> + Send + 'a>>;
11
12#[derive(Debug, Clone, PartialEq, Eq)]
14pub struct EmbeddingRequest {
15 pub inputs: Vec<String>,
16 pub model: String,
17 pub dimension: u32,
18}
19
20#[derive(Debug, Clone, PartialEq)]
22pub struct EmbeddingVector {
23 pub values: Vec<f64>,
24}
25
26pub trait EmbeddingProvider: Send + Sync {
28 fn embed(&self, request: EmbeddingRequest) -> EmbeddingFuture<'_, Vec<EmbeddingVector>>;
29}
30
31#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub enum ProviderRetryClass {
34 Retryable,
35 Permanent,
36}
37
38#[derive(Debug, Clone, PartialEq, Eq)]
40pub struct EmbeddingProviderError {
41 pub retry: ProviderRetryClass,
42 pub status_code: Option<u16>,
43 pub code: String,
44 pub message: String,
45}
46
47impl fmt::Display for EmbeddingProviderError {
48 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
49 match self.status_code {
50 Some(status) => write!(formatter, "{} ({status}): {}", self.code, self.message),
51 None => write!(formatter, "{}: {}", self.code, self.message),
52 }
53 }
54}
55
56impl Error for EmbeddingProviderError {}
57
58pub fn embedding_provider(
60 config: RemoteEmbeddingConfig,
61 client: reqwest::Client,
62) -> Box<dyn EmbeddingProvider> {
63 match config.provider {
64 EmbeddingProviderKind::OpenAiCompatible => {
65 Box::new(OpenAiCompatibleEmbeddingProvider { config, client })
66 }
67 EmbeddingProviderKind::Echo => Box::new(EchoEmbeddingProvider { config }),
68 }
69}
70
71struct OpenAiCompatibleEmbeddingProvider {
72 config: RemoteEmbeddingConfig,
73 client: reqwest::Client,
74}
75
76impl EmbeddingProvider for OpenAiCompatibleEmbeddingProvider {
77 fn embed(&self, request: EmbeddingRequest) -> EmbeddingFuture<'_, Vec<EmbeddingVector>> {
78 Box::pin(async move {
79 validate_request(&request)?;
80 let url = embeddings_url(&self.config.base_url);
81 let response = self
82 .client
83 .post(url)
84 .bearer_auth(&self.config.api_key)
85 .timeout(self.config.timeout)
86 .json(&OpenAiEmbeddingRequest {
87 model: &request.model,
88 input: &request.inputs,
89 })
90 .send()
91 .await
92 .map_err(transport_error)?;
93 let status = response.status();
94 if !status.is_success() {
95 return Err(status_error(status.as_u16(), response.text().await.ok()));
96 }
97 let payload = response
98 .json::<OpenAiEmbeddingResponse>()
99 .await
100 .map_err(|error| permanent_error("invalid_response_json", error.to_string()))?;
101
102 parse_embedding_response(payload, request.inputs.len(), request.dimension)
103 })
104 }
105}
106
107struct EchoEmbeddingProvider {
108 config: RemoteEmbeddingConfig,
109}
110
111impl EmbeddingProvider for EchoEmbeddingProvider {
112 fn embed(&self, request: EmbeddingRequest) -> EmbeddingFuture<'_, Vec<EmbeddingVector>> {
113 Box::pin(async move {
114 validate_request(&request)?;
115 let dimension = usize::try_from(request.dimension).map_err(|_| {
116 permanent_error("invalid_dimension", "embedding dimension is too large")
117 })?;
118 let vectors = request
119 .inputs
120 .iter()
121 .map(|input| deterministic_vector(input, dimension))
122 .collect::<Vec<_>>();
123 let _ = &self.config;
124
125 Ok(vectors)
126 })
127 }
128}
129
130#[derive(Serialize)]
131struct OpenAiEmbeddingRequest<'a> {
132 model: &'a str,
133 input: &'a [String],
134}
135
136#[derive(Deserialize)]
137struct OpenAiEmbeddingResponse {
138 data: Vec<OpenAiEmbeddingData>,
139}
140
141#[derive(Deserialize)]
142struct OpenAiEmbeddingData {
143 embedding: Vec<f64>,
144}
145
146fn parse_embedding_response(
147 response: OpenAiEmbeddingResponse,
148 expected_count: usize,
149 expected_dimension: u32,
150) -> Result<Vec<EmbeddingVector>, EmbeddingProviderError> {
151 if response.data.len() != expected_count {
152 return Err(permanent_error(
153 "embedding_count_mismatch",
154 format!(
155 "provider returned {} embeddings for {} inputs",
156 response.data.len(),
157 expected_count
158 ),
159 ));
160 }
161 let expected_dimension = usize::try_from(expected_dimension)
162 .map_err(|_| permanent_error("invalid_dimension", "embedding dimension is too large"))?;
163 response
164 .data
165 .into_iter()
166 .map(|item| validate_vector(item.embedding, expected_dimension))
167 .collect()
168}
169
170fn validate_request(request: &EmbeddingRequest) -> Result<(), EmbeddingProviderError> {
171 if request.inputs.is_empty() {
172 return Err(permanent_error(
173 "empty_embedding_batch",
174 "embedding request must contain at least one input",
175 ));
176 }
177 if request.model.trim().is_empty() {
178 return Err(permanent_error(
179 "empty_embedding_model",
180 "embedding model must not be blank",
181 ));
182 }
183 if request.dimension == 0 {
184 return Err(permanent_error(
185 "invalid_dimension",
186 "embedding dimension must be greater than zero",
187 ));
188 }
189
190 Ok(())
191}
192
193fn validate_vector(
194 values: Vec<f64>,
195 expected_dimension: usize,
196) -> Result<EmbeddingVector, EmbeddingProviderError> {
197 if values.len() != expected_dimension {
198 return Err(permanent_error(
199 "embedding_dimension_mismatch",
200 format!(
201 "provider returned dimension {} while {} was configured",
202 values.len(),
203 expected_dimension
204 ),
205 ));
206 }
207 if values.iter().any(|value| !value.is_finite()) {
208 return Err(permanent_error(
209 "invalid_embedding_value",
210 "provider returned a non-finite embedding value",
211 ));
212 }
213
214 Ok(EmbeddingVector { values })
215}
216
217fn embeddings_url(base_url: &str) -> String {
218 let base = base_url
219 .trim()
220 .split(['?', '#'])
221 .next()
222 .unwrap_or("")
223 .trim_end_matches('/');
224 if base.ends_with("/embeddings") {
225 return base.to_owned();
226 }
227 if final_path_segment(base).is_some_and(is_api_version_segment) {
228 return format!("{base}/embeddings");
229 }
230
231 format!("{base}/v1/embeddings")
232}
233
234fn final_path_segment(url: &str) -> Option<&str> {
235 let after_authority = url.split_once("://").map_or(url, |(_, rest)| rest);
236 let path = after_authority.split_once('/')?.1;
237 let path = path.split(['?', '#']).next().unwrap_or(path);
238
239 path.rsplit('/').find(|segment| !segment.is_empty())
240}
241
242fn is_api_version_segment(segment: &str) -> bool {
243 let Some(digits) = segment
244 .strip_prefix('v')
245 .or_else(|| segment.strip_prefix('V'))
246 else {
247 return false;
248 };
249
250 !digits.is_empty() && digits.chars().all(|character| character.is_ascii_digit())
251}
252
253fn deterministic_vector(input: &str, dimension: usize) -> EmbeddingVector {
254 let mut values = vec![0.0; dimension];
255 for (index, byte) in input.bytes().enumerate() {
256 values[index % dimension] += f64::from(byte) / 255.0;
257 }
258 let norm = values.iter().map(|value| value * value).sum::<f64>().sqrt();
259 if norm > 0.0 {
260 for value in &mut values {
261 *value /= norm;
262 }
263 }
264
265 EmbeddingVector { values }
266}
267
268fn status_error(status_code: u16, body: Option<String>) -> EmbeddingProviderError {
269 let body_reports_resource_limit = body
270 .as_deref()
271 .is_some_and(provider_error_reports_resource_limit);
272 let resource_limited = matches!(status_code, 402 | 429)
273 || (status_allows_resource_limit_body(status_code) && body_reports_resource_limit);
274 let retry = if resource_limited || matches!(status_code, 408 | 500..=599) {
275 ProviderRetryClass::Retryable
276 } else {
277 ProviderRetryClass::Permanent
278 };
279
280 EmbeddingProviderError {
281 retry,
282 status_code: Some(status_code),
283 code: status_code_error_code(status_code, resource_limited).to_owned(),
284 message: body
285 .map(error_body_preview)
286 .unwrap_or_else(|| "provider request failed".to_owned()),
287 }
288}
289
290fn transport_error(error: reqwest::Error) -> EmbeddingProviderError {
291 let code = if error.is_timeout() {
292 "network_timeout"
293 } else {
294 "network_error"
295 };
296
297 EmbeddingProviderError {
298 retry: ProviderRetryClass::Retryable,
299 status_code: error.status().map(|status| status.as_u16()),
300 code: code.to_owned(),
301 message: error.to_string(),
302 }
303}
304
305fn permanent_error(code: &'static str, message: impl Into<String>) -> EmbeddingProviderError {
306 EmbeddingProviderError {
307 retry: ProviderRetryClass::Permanent,
308 status_code: None,
309 code: code.to_owned(),
310 message: message.into(),
311 }
312}
313
314fn status_code_error_code(status_code: u16, resource_limited: bool) -> &'static str {
315 if resource_limited {
316 return "rate_limited";
317 }
318
319 match status_code {
320 400 => "invalid_request",
321 401 | 403 => "auth_invalid",
322 404 => "model_or_endpoint_not_found",
323 408 => "network_timeout",
324 500..=599 => "provider_unavailable",
325 _ => "provider_http_error",
326 }
327}
328
329fn status_allows_resource_limit_body(status_code: u16) -> bool {
330 matches!(status_code, 400 | 403 | 409 | 425 | 500..=599)
331}
332
333fn provider_error_reports_resource_limit(body: &str) -> bool {
334 if let Ok(payload) = serde_json::from_str::<serde_json::Value>(body) {
335 return json_strings_report_resource_limit(&payload);
336 }
337
338 text_reports_resource_limit(body)
339}
340
341fn json_strings_report_resource_limit(value: &serde_json::Value) -> bool {
342 match value {
343 serde_json::Value::String(text) => text_reports_resource_limit(text),
344 serde_json::Value::Array(values) => values.iter().any(json_strings_report_resource_limit),
345 serde_json::Value::Object(fields) => {
346 fields.values().any(json_strings_report_resource_limit)
347 }
348 serde_json::Value::Null | serde_json::Value::Bool(_) | serde_json::Value::Number(_) => {
349 false
350 }
351 }
352}
353
354fn text_reports_resource_limit(text: &str) -> bool {
355 let normalized = text
356 .chars()
357 .map(|character| {
358 if character == '_' || character == '-' {
359 ' '
360 } else {
361 character.to_ascii_lowercase()
362 }
363 })
364 .collect::<String>();
365
366 [
367 "rate limit",
368 "too many request",
369 "insufficient quota",
370 "quota exceeded",
371 "quota exhausted",
372 "out of quota",
373 "insufficient balance",
374 "resource exhausted",
375 "no resource package",
376 "capacity exceeded",
377 "billing limit",
378 "payment required",
379 ]
380 .iter()
381 .any(|marker| normalized.contains(marker))
382}
383
384fn error_body_preview(value: String) -> String {
385 value.chars().take(PROVIDER_ERROR_MESSAGE_LIMIT).collect()
386}
387
388#[cfg(test)]
389mod tests {
390 use super::*;
391
392 #[test]
393 fn embeddings_url_accepts_base_or_endpoint() {
394 assert_eq!(
395 embeddings_url("https://example.test"),
396 "https://example.test/v1/embeddings"
397 );
398 assert_eq!(
399 embeddings_url("https://example.test/v1"),
400 "https://example.test/v1/embeddings"
401 );
402 assert_eq!(
403 embeddings_url("https://example.test/v1/embeddings"),
404 "https://example.test/v1/embeddings"
405 );
406 assert_eq!(
407 embeddings_url("https://example.test/v4"),
408 "https://example.test/v4/embeddings"
409 );
410 assert_eq!(
411 embeddings_url("https://example.test/openai/v2"),
412 "https://example.test/openai/v2/embeddings"
413 );
414 assert_eq!(
415 embeddings_url("https://example.test/openai"),
416 "https://example.test/openai/v1/embeddings"
417 );
418 assert_eq!(
419 embeddings_url("https://example.test/openai/v4/?probe=true#fragment"),
420 "https://example.test/openai/v4/embeddings"
421 );
422 assert_eq!(
423 embeddings_url("https://example.test/v4/embeddings?probe=true"),
424 "https://example.test/v4/embeddings"
425 );
426 }
427
428 #[test]
429 fn rejects_embedding_dimension_mismatch() {
430 let response = OpenAiEmbeddingResponse {
431 data: vec![OpenAiEmbeddingData {
432 embedding: vec![0.1, 0.2],
433 }],
434 };
435
436 let error = parse_embedding_response(response, 1, 3).expect_err("dimension should fail");
437
438 assert_eq!(error.code, "embedding_dimension_mismatch");
439 assert_eq!(error.retry, ProviderRetryClass::Permanent);
440 }
441
442 #[test]
443 fn classifies_rate_limit_as_retryable() {
444 let error = status_error(429, None);
445
446 assert_eq!(error.retry, ProviderRetryClass::Retryable);
447 assert_eq!(error.code, "rate_limited");
448 }
449
450 #[test]
451 fn classifies_provider_resource_limit_bodies_as_retryable() {
452 let payment_required = status_error(402, None);
453 let quota_forbidden = status_error(
454 403,
455 Some(
456 r#"{"error":{"code":"insufficient_quota","message":"Insufficient balance or no resource package."}}"#
457 .to_owned(),
458 ),
459 );
460 let invalid_request_quota = status_error(
461 400,
462 Some(
463 r#"{"error":{"type":"resource_exhausted","message":"quota exceeded"}}"#.to_owned(),
464 ),
465 );
466 let retry_after_resource_exhausted = status_error(
467 503,
468 Some(
469 r#"{"error":{"status":"RESOURCE_EXHAUSTED","message":"rate limit exceeded"}}"#
470 .to_owned(),
471 ),
472 );
473 let top_level_message_quota = status_error(
474 403,
475 Some(
476 r#"{"error":{"code":"invalid_request"},"message":"Quota exceeded for tenant"}"#
477 .to_owned(),
478 ),
479 );
480 let nested_detail_resource_exhausted = status_error(
481 500,
482 Some(
483 r#"{"error":{"code":"provider_error"},"details":[{"reason":"Resource exhausted"}]}"#
484 .to_owned(),
485 ),
486 );
487
488 assert_eq!(payment_required.retry, ProviderRetryClass::Retryable);
489 assert_eq!(payment_required.code, "rate_limited");
490 assert_eq!(quota_forbidden.retry, ProviderRetryClass::Retryable);
491 assert_eq!(quota_forbidden.code, "rate_limited");
492 assert_eq!(invalid_request_quota.retry, ProviderRetryClass::Retryable);
493 assert_eq!(invalid_request_quota.code, "rate_limited");
494 assert_eq!(
495 retry_after_resource_exhausted.retry,
496 ProviderRetryClass::Retryable
497 );
498 assert_eq!(retry_after_resource_exhausted.code, "rate_limited");
499 assert_eq!(top_level_message_quota.retry, ProviderRetryClass::Retryable);
500 assert_eq!(top_level_message_quota.code, "rate_limited");
501 assert_eq!(
502 nested_detail_resource_exhausted.retry,
503 ProviderRetryClass::Retryable
504 );
505 assert_eq!(nested_detail_resource_exhausted.code, "rate_limited");
506 }
507
508 #[test]
509 fn preserves_permanent_provider_errors_without_resource_limit_signals() {
510 let auth_forbidden = status_error(
511 403,
512 Some(r#"{"error":{"code":"invalid_api_key","message":"Invalid API key"}}"#.to_owned()),
513 );
514 let invalid_request = status_error(
515 400,
516 Some(
517 r#"{"error":{"code":"invalid_request","message":"quota field is not supported"}}"#
518 .to_owned(),
519 ),
520 );
521 let limit_key_without_limited_value = status_error(
522 400,
523 Some(r#"{"error":{"code":"invalid_request"},"rate_limit":false}"#.to_owned()),
524 );
525
526 assert_eq!(auth_forbidden.retry, ProviderRetryClass::Permanent);
527 assert_eq!(auth_forbidden.code, "auth_invalid");
528 assert_eq!(invalid_request.retry, ProviderRetryClass::Permanent);
529 assert_eq!(invalid_request.code, "invalid_request");
530 assert_eq!(
531 limit_key_without_limited_value.retry,
532 ProviderRetryClass::Permanent
533 );
534 assert_eq!(limit_key_without_limited_value.code, "invalid_request");
535 }
536
537 #[test]
538 fn classifies_provider_http_status_codes() {
539 for (status, code, retry) in [
540 (400, "invalid_request", ProviderRetryClass::Permanent),
541 (401, "auth_invalid", ProviderRetryClass::Permanent),
542 (403, "auth_invalid", ProviderRetryClass::Permanent),
543 (
544 404,
545 "model_or_endpoint_not_found",
546 ProviderRetryClass::Permanent,
547 ),
548 (408, "network_timeout", ProviderRetryClass::Retryable),
549 (500, "provider_unavailable", ProviderRetryClass::Retryable),
550 (418, "provider_http_error", ProviderRetryClass::Permanent),
551 ] {
552 let error = status_error(status, Some("x".repeat(300)));
553
554 assert_eq!(error.code, code);
555 assert_eq!(error.retry, retry);
556 assert_eq!(error.message.len(), 240);
557 }
558 }
559
560 #[tokio::test]
561 async fn openai_provider_posts_and_parses_embeddings() {
562 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
563 .await
564 .expect("listener should bind");
565 let addr = listener.local_addr().expect("local addr should load");
566 let server = tokio::spawn(async move {
567 let (stream, _) = listener.accept().await.expect("request should connect");
568 let mut buffer = vec![0; 2048];
569 let count = stream
570 .readable()
571 .await
572 .and_then(|()| stream.try_read(&mut buffer));
573 let request = String::from_utf8_lossy(&buffer[..count.expect("request should read")]);
574
575 assert!(request.starts_with("POST /v1/embeddings HTTP/1.1"));
576 assert!(request.contains("authorization: Bearer secret"));
577 assert!(request.contains("\"model\":\"text-embedding-3-small\""));
578 stream
579 .writable()
580 .await
581 .expect("stream should become writable");
582 stream
583 .try_write(
584 b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 34\r\n\r\n{\"data\":[{\"embedding\":[0.1,0.2]}]}",
585 )
586 .expect("response should write");
587 });
588 let provider = OpenAiCompatibleEmbeddingProvider {
589 config: remote_config(
590 format!("http://{addr}/v1"),
591 std::time::Duration::from_secs(5),
592 ),
593 client: reqwest::Client::new(),
594 };
595
596 let vectors = provider
597 .embed(EmbeddingRequest {
598 inputs: vec!["probe".to_owned()],
599 model: "text-embedding-3-small".to_owned(),
600 dimension: 2,
601 })
602 .await
603 .expect("provider response should parse");
604
605 assert_eq!(vectors[0].values, [0.1, 0.2]);
606 server.await.expect("server should finish");
607 }
608
609 #[tokio::test]
610 async fn echo_provider_returns_deterministic_vectors() {
611 let provider = EchoEmbeddingProvider {
612 config: remote_config("http://example.test/v1", std::time::Duration::from_secs(5)),
613 };
614
615 let vectors = provider
616 .embed(EmbeddingRequest {
617 inputs: vec!["abc".to_owned(), "abc".to_owned()],
618 model: "echo".to_owned(),
619 dimension: 4,
620 })
621 .await
622 .expect("echo provider should embed");
623
624 assert_eq!(vectors.len(), 2);
625 assert_eq!(vectors[0], vectors[1]);
626 assert_eq!(vectors[0].values.len(), 4);
627 }
628
629 #[test]
630 fn rejects_invalid_requests_and_response_values() {
631 let empty = validate_request(&EmbeddingRequest {
632 inputs: Vec::new(),
633 model: "model".to_owned(),
634 dimension: 1,
635 })
636 .expect_err("empty inputs should fail");
637 let model = validate_request(&EmbeddingRequest {
638 inputs: vec!["x".to_owned()],
639 model: " ".to_owned(),
640 dimension: 1,
641 })
642 .expect_err("blank model should fail");
643 let dimension = validate_request(&EmbeddingRequest {
644 inputs: vec!["x".to_owned()],
645 model: "model".to_owned(),
646 dimension: 0,
647 })
648 .expect_err("zero dimension should fail");
649 let invalid_value = validate_vector(vec![f64::NAN], 1).expect_err("nan values should fail");
650
651 assert_eq!(empty.code, "empty_embedding_batch");
652 assert_eq!(model.code, "empty_embedding_model");
653 assert_eq!(dimension.code, "invalid_dimension");
654 assert_eq!(invalid_value.code, "invalid_embedding_value");
655 }
656
657 #[tokio::test]
658 async fn applies_configured_embedding_timeout() {
659 let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
660 .await
661 .expect("listener should bind");
662 let addr = listener.local_addr().expect("local addr should load");
663 let server = tokio::spawn(async move {
664 let (_stream, _) = listener.accept().await.expect("request should connect");
665 tokio::time::sleep(std::time::Duration::from_secs(2)).await;
666 });
667 let provider = OpenAiCompatibleEmbeddingProvider {
668 config: RemoteEmbeddingConfig {
669 provider: EmbeddingProviderKind::OpenAiCompatible,
670 base_url: format!("http://{addr}/v1"),
671 api_key: "secret".to_owned(),
672 batch_size: 1,
673 timeout: std::time::Duration::from_millis(20),
674 max_concurrency: 1,
675 },
676 client: reqwest::Client::builder()
677 .timeout(std::time::Duration::from_secs(5))
678 .build()
679 .expect("client should build"),
680 };
681
682 let error = provider
683 .embed(EmbeddingRequest {
684 inputs: vec!["probe".to_owned()],
685 model: "text-embedding-3-small".to_owned(),
686 dimension: 3,
687 })
688 .await
689 .expect_err("provider request should use embedding timeout");
690
691 assert_eq!(error.code, "network_timeout");
692 server.abort();
693 }
694
695 fn remote_config(
696 base_url: impl Into<String>,
697 timeout: std::time::Duration,
698 ) -> RemoteEmbeddingConfig {
699 RemoteEmbeddingConfig {
700 provider: EmbeddingProviderKind::OpenAiCompatible,
701 base_url: base_url.into(),
702 api_key: "secret".to_owned(),
703 batch_size: 1,
704 timeout,
705 max_concurrency: 1,
706 }
707 }
708}