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)]
462#[path = "provider_tests.rs"]
463mod tests;