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