Skip to main content

relay_knowledge/retrieval/
provider.rs

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/// Text inputs sent to a remote embedding provider.
17#[derive(Debug, Clone, PartialEq, Eq)]
18pub struct EmbeddingRequest {
19    pub inputs: Vec<String>,
20    pub model: String,
21    pub dimension: u32,
22}
23
24/// One normalized embedding vector returned by a provider.
25#[derive(Debug, Clone, PartialEq)]
26pub struct EmbeddingVector {
27    pub values: Vec<f64>,
28}
29
30/// Provider-neutral remote embedding contract.
31pub trait EmbeddingProvider: Send + Sync {
32    fn embed(&self, request: EmbeddingRequest) -> EmbeddingFuture<'_, Vec<EmbeddingVector>>;
33}
34
35/// Retry category for remote provider failures.
36#[derive(Debug, Clone, Copy, PartialEq, Eq)]
37pub enum ProviderRetryClass {
38    Retryable,
39    Permanent,
40}
41
42/// Provider error safe for diagnostics.
43#[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
62/// Builds the configured remote embedding provider.
63pub 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
70/// Builds the configured remote embedding provider with outbound QoS admission.
71pub 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}