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