1use std::collections::BTreeSet;
4use std::sync::atomic::{AtomicU8, Ordering};
5use std::sync::Arc;
6use std::time::Duration;
7
8use futures::StreamExt;
9use mongreldb_core::{
10 EmbeddingError, EmbeddingFuture, EmbeddingNormalization, EmbeddingProvider, EmbeddingRequest,
11 EmbeddingResponse, ProviderExecutionMode, ProviderHealth,
12};
13use serde::{Deserialize, Serialize};
14use zeroize::Zeroizing;
15
16const HEALTH_READY: u8 = 0;
17const HEALTH_DEGRADED: u8 = 1;
18const HEALTH_UNAVAILABLE: u8 = 2;
19
20pub trait EmbeddingSecretResolver: Send + Sync {
21 fn resolve(&self, reference: &str) -> Result<Zeroizing<String>, EmbeddingError>;
22}
23
24#[derive(Debug, Default)]
27pub struct EnvironmentSecretResolver;
28
29impl EmbeddingSecretResolver for EnvironmentSecretResolver {
30 fn resolve(&self, reference: &str) -> Result<Zeroizing<String>, EmbeddingError> {
31 if reference.is_empty()
32 || !reference
33 .bytes()
34 .all(|byte| byte.is_ascii_uppercase() || byte.is_ascii_digit() || byte == b'_')
35 {
36 return Err(provider_error("invalid environment secret reference"));
37 }
38 std::env::var(reference)
39 .map(Zeroizing::new)
40 .map_err(|_| provider_error("embedding secret reference is unavailable"))
41 }
42}
43
44#[derive(Clone)]
45pub struct RemoteEmbeddingConfig {
46 pub provider_id: String,
47 pub model_id: String,
48 pub model_version: String,
49 pub preprocessing_version: String,
50 pub dimension: u32,
51 pub normalization: EmbeddingNormalization,
52 pub endpoint: reqwest::Url,
53 pub allowed_hosts: BTreeSet<String>,
54 pub secret_reference: String,
55 pub tenant: String,
56 pub timeout: Duration,
57 pub max_retries: usize,
58 pub max_response_bytes: usize,
59}
60
61impl std::fmt::Debug for RemoteEmbeddingConfig {
62 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
63 formatter
64 .debug_struct("RemoteEmbeddingConfig")
65 .field("provider_id", &self.provider_id)
66 .field("model_id", &self.model_id)
67 .field("model_version", &self.model_version)
68 .field("endpoint", &self.endpoint)
69 .field("secret_reference", &"<redacted>")
70 .field("tenant", &self.tenant)
71 .finish_non_exhaustive()
72 }
73}
74
75#[derive(Clone)]
76pub struct RemoteEmbeddingProvider {
77 config: RemoteEmbeddingConfig,
78 secrets: Arc<dyn EmbeddingSecretResolver>,
79 client: reqwest::Client,
80 health: Arc<AtomicU8>,
81}
82
83impl RemoteEmbeddingProvider {
84 pub fn new(
85 config: RemoteEmbeddingConfig,
86 secrets: Arc<dyn EmbeddingSecretResolver>,
87 ) -> Result<Self, EmbeddingError> {
88 let host = config
89 .endpoint
90 .host_str()
91 .ok_or_else(|| provider_error("embedding endpoint has no host"))?;
92 if config.endpoint.scheme() != "https" {
93 return Err(provider_error("embedding endpoint must use HTTPS"));
94 }
95 if !config.allowed_hosts.contains(host) {
96 return Err(provider_error("embedding endpoint host is not allowlisted"));
97 }
98 if !config.endpoint.username().is_empty() || config.endpoint.password().is_some() {
99 return Err(provider_error(
100 "embedding endpoint must not contain credentials",
101 ));
102 }
103 if config.tenant.is_empty()
104 || config.timeout.is_zero()
105 || config.max_response_bytes == 0
106 || config.dimension == 0
107 {
108 return Err(provider_error("invalid remote embedding limits"));
109 }
110 let client = reqwest::Client::builder()
111 .https_only(true)
112 .redirect(reqwest::redirect::Policy::none())
113 .timeout(config.timeout)
114 .build()
115 .map_err(|error| provider_error(error.to_string()))?;
116 Ok(Self {
117 config,
118 secrets,
119 client,
120 health: Arc::new(AtomicU8::new(HEALTH_READY)),
121 })
122 }
123
124 async fn request(
125 &self,
126 texts: &[&str],
127 trace_id: &str,
128 ) -> Result<EmbeddingResponse, EmbeddingError> {
129 let secret = self.secrets.resolve(&self.config.secret_reference)?;
130 let request = RemoteRequest {
131 model: &self.config.model_id,
132 input: texts,
133 };
134 let mut attempt = 0;
135 loop {
136 let result = self
137 .client
138 .post(self.config.endpoint.clone())
139 .bearer_auth(secret.as_str())
140 .header("x-mongreldb-tenant", &self.config.tenant)
141 .header("x-mongreldb-trace-id", trace_id)
142 .json(&request)
143 .send()
144 .await;
145 match result {
146 Ok(response)
147 if response.status().is_server_error() && attempt < self.config.max_retries =>
148 {
149 attempt += 1;
150 continue;
151 }
152 Err(error)
153 if (error.is_connect() || error.is_timeout())
154 && attempt < self.config.max_retries =>
155 {
156 attempt += 1;
157 continue;
158 }
159 Ok(response) => return self.decode(response).await,
160 Err(error) => {
161 self.health.store(HEALTH_UNAVAILABLE, Ordering::Release);
162 return Err(provider_error(error.to_string()));
163 }
164 }
165 }
166 }
167
168 async fn decode(
169 &self,
170 response: reqwest::Response,
171 ) -> Result<EmbeddingResponse, EmbeddingError> {
172 if !response.status().is_success() {
173 self.health.store(HEALTH_DEGRADED, Ordering::Release);
174 return Err(provider_error(format!(
175 "remote provider returned HTTP {}",
176 response.status()
177 )));
178 }
179 if response
180 .content_length()
181 .is_some_and(|length| length > self.config.max_response_bytes as u64)
182 {
183 return Err(provider_error("remote embedding response is too large"));
184 }
185 let mut bytes = Vec::new();
186 let mut stream = response.bytes_stream();
187 while let Some(chunk) = stream.next().await {
188 let chunk = chunk.map_err(|error| provider_error(error.to_string()))?;
189 if bytes.len().saturating_add(chunk.len()) > self.config.max_response_bytes {
190 return Err(provider_error("remote embedding response is too large"));
191 }
192 bytes.extend_from_slice(&chunk);
193 }
194 let response: RemoteResponse =
195 serde_json::from_slice(&bytes).map_err(|error| provider_error(error.to_string()))?;
196 self.health.store(HEALTH_READY, Ordering::Release);
197 Ok(EmbeddingResponse {
198 vectors: match response {
199 RemoteResponse::Vectors { vectors } => vectors,
200 RemoteResponse::Data { data } => {
201 data.into_iter().map(|item| item.embedding).collect()
202 }
203 },
204 })
205 }
206}
207
208impl EmbeddingProvider for RemoteEmbeddingProvider {
209 fn provider_id(&self) -> &str {
210 &self.config.provider_id
211 }
212
213 fn model_id(&self) -> &str {
214 &self.config.model_id
215 }
216
217 fn model_version(&self) -> &str {
218 &self.config.model_version
219 }
220
221 fn dimension(&self) -> u32 {
222 self.config.dimension
223 }
224
225 fn normalization(&self) -> EmbeddingNormalization {
226 self.config.normalization
227 }
228
229 fn preprocessing_version(&self) -> &str {
230 &self.config.preprocessing_version
231 }
232
233 fn execution_mode(&self) -> ProviderExecutionMode {
234 ProviderExecutionMode::Remote
235 }
236
237 fn health(&self) -> ProviderHealth {
238 match self.health.load(Ordering::Acquire) {
239 HEALTH_READY => ProviderHealth::Ready,
240 HEALTH_DEGRADED => ProviderHealth::Degraded,
241 _ => ProviderHealth::Unavailable,
242 }
243 }
244
245 fn embed(&self, _request: EmbeddingRequest<'_>) -> Result<EmbeddingResponse, EmbeddingError> {
246 Err(provider_error(
247 "remote embedding provider requires asynchronous execution",
248 ))
249 }
250
251 fn embed_async<'a>(&'a self, request: EmbeddingRequest<'a>) -> EmbeddingFuture<'a> {
252 Box::pin(async move { self.request(request.texts, request.trace_id).await })
253 }
254}
255
256#[derive(Serialize)]
257struct RemoteRequest<'a> {
258 model: &'a str,
259 input: &'a [&'a str],
260}
261
262#[derive(Deserialize)]
263#[serde(untagged)]
264enum RemoteResponse {
265 Vectors { vectors: Vec<Vec<f32>> },
266 Data { data: Vec<RemoteEmbedding> },
267}
268
269#[derive(Deserialize)]
270struct RemoteEmbedding {
271 embedding: Vec<f32>,
272}
273
274fn provider_error(message: impl Into<String>) -> EmbeddingError {
275 EmbeddingError::ProviderFailed {
276 provider: "remote".into(),
277 message: message.into(),
278 }
279}
280
281#[cfg(test)]
282mod tests {
283 use super::*;
284
285 #[test]
286 fn endpoint_must_be_https_and_allowlisted() {
287 let config = |endpoint: &str, allowed_hosts: &[&str]| RemoteEmbeddingConfig {
288 provider_id: "remote".into(),
289 model_id: "model".into(),
290 model_version: "1".into(),
291 preprocessing_version: "1".into(),
292 dimension: 2,
293 normalization: EmbeddingNormalization::None,
294 endpoint: endpoint.parse().unwrap(),
295 allowed_hosts: allowed_hosts.iter().map(|host| (*host).into()).collect(),
296 secret_reference: "EMBEDDING_TOKEN".into(),
297 tenant: "tenant-a".into(),
298 timeout: Duration::from_secs(1),
299 max_retries: 1,
300 max_response_bytes: 1024,
301 };
302 let secrets: Arc<dyn EmbeddingSecretResolver> = Arc::new(EnvironmentSecretResolver);
303 assert!(RemoteEmbeddingProvider::new(
304 config("http://provider.example/embed", &["provider.example"]),
305 Arc::clone(&secrets),
306 )
307 .is_err());
308 assert!(RemoteEmbeddingProvider::new(
309 config("https://provider.example/embed", &["other.example"]),
310 secrets,
311 )
312 .is_err());
313 }
314}