Skip to main content

mongreldb_server/
remote_embedding.rs

1//! HTTPS embedding provider with bounded I/O and referenced secrets.
2
3use 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/// Environment-backed secret references. Configuration stores only the
25/// variable name. Values never enter Debug output or replicated schema.
26#[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}