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