Skip to main content

smriti/ml/
remote_embedder.rs

1//! Remote GPU bridge embedder.
2//!
3//! Sends face crops via HTTP multipart to a Colab/Kaggle-hosted GPU
4//! inference server. Falls back to local ONNX when the bridge is
5//! unreachable or returns errors. The bridge must advertise the same
6//! face-embedding model the desktop is configured for — embeddings
7//! from different models live in incompatible metric spaces and would
8//! corrupt clustering if mixed.
9
10use std::io::Cursor;
11use std::time::{Duration, Instant};
12
13use image::codecs::jpeg::JpegEncoder;
14use image::RgbImage;
15
16use super::face_embedder::FaceEmbedding;
17
18const HEARTBEAT_INTERVAL: Duration = Duration::from_secs(30);
19const JPEG_QUALITY: u8 = 85;
20
21pub struct RemoteEmbedder {
22    client: reqwest::blocking::Client,
23    base_url: String,
24    /// Model name the desktop expects the bridge to serve. Compared to the
25    /// notebook's `/health.model` field, normalized by stripping any
26    /// trailing `.onnx`.
27    expected_model: String,
28    healthy: bool,
29    consecutive_failures: u32,
30    last_health_at: Instant,
31}
32
33impl RemoteEmbedder {
34    pub fn new(base_url: String, expected_model: String) -> Self {
35        let client = reqwest::blocking::Client::builder()
36            .timeout(Duration::from_secs(30))
37            .build()
38            .unwrap_or_default();
39        let healthy = Self::check_health(&client, &base_url, &expected_model);
40        if !healthy {
41            tracing::warn!(
42                "Remote GPU bridge at {} unhealthy or wrong model (expected {}); will use local fallback",
43                base_url,
44                expected_model
45            );
46        }
47        Self {
48            client,
49            base_url,
50            expected_model,
51            healthy,
52            consecutive_failures: 0,
53            last_health_at: Instant::now(),
54        }
55    }
56
57    pub fn is_healthy(&self) -> bool {
58        self.healthy
59    }
60
61    /// Health check. Returns true only when the bridge is reachable,
62    /// reports a GPU provider, AND the served model matches the one
63    /// the desktop is configured to use. Mismatched models silently
64    /// produce embeddings in a different metric space — refusing here
65    /// is the only safe response.
66    fn check_health(
67        client: &reqwest::blocking::Client,
68        base_url: &str,
69        expected_model: &str,
70    ) -> bool {
71        let resp = match client
72            .get(format!("{}/health", base_url))
73            .timeout(Duration::from_secs(2))
74            .send()
75        {
76            Ok(r) => r,
77            Err(_) => return false,
78        };
79        if !resp.status().is_success() {
80            return false;
81        }
82        let body: serde_json::Value = match resp.json() {
83            Ok(v) => v,
84            Err(_) => return false,
85        };
86        let has_gpu = body
87            .get("provider")
88            .and_then(|v| v.as_str())
89            .is_some_and(|s| s.contains("CUDA") || s.contains("GPU"));
90        if !has_gpu {
91            return false;
92        }
93        let served_model = body.get("model").and_then(|v| v.as_str()).unwrap_or("");
94        if !models_match(served_model, expected_model) {
95            tracing::warn!(
96                "Remote bridge model mismatch: served={:?}, expected={:?}",
97                served_model,
98                expected_model
99            );
100            return false;
101        }
102        true
103    }
104
105    pub fn embed_batch(&mut self, faces: &[RgbImage]) -> Vec<Option<FaceEmbedding>> {
106        let n = faces.len();
107        if n == 0 {
108            return Vec::new();
109        }
110
111        if !self.healthy {
112            return vec![None; n];
113        }
114
115        // Lazy heartbeat: if it's been a while since we last verified
116        // the bridge, ping /health before sending real work. Catches a
117        // bridge that died silently without us having to wait for the
118        // 3-strikes reactive path.
119        if self.last_health_at.elapsed() >= HEARTBEAT_INTERVAL {
120            let ok = Self::check_health(&self.client, &self.base_url, &self.expected_model);
121            self.last_health_at = Instant::now();
122            if !ok {
123                self.healthy = false;
124                tracing::warn!("Remote GPU bridge failed lazy heartbeat; marking unhealthy");
125                return vec![None; n];
126            }
127        }
128
129        match self.do_embed_batch(faces) {
130            Ok(embeddings) => {
131                self.consecutive_failures = 0;
132                self.last_health_at = Instant::now();
133                embeddings
134            }
135            Err(e) => {
136                self.consecutive_failures += 1;
137                tracing::warn!(
138                    "Remote embed batch failed ({}/3): {}",
139                    self.consecutive_failures,
140                    e
141                );
142                if self.consecutive_failures >= 3 {
143                    self.healthy = false;
144                    tracing::error!(
145                        "Remote GPU bridge marked unhealthy after {} consecutive failures",
146                        self.consecutive_failures
147                    );
148                }
149                vec![None; n]
150            }
151        }
152    }
153
154    pub fn embed(&mut self, face: &RgbImage) -> Option<FaceEmbedding> {
155        self.embed_batch(std::slice::from_ref(face))
156            .into_iter()
157            .next()
158            .flatten()
159    }
160
161    fn do_embed_batch(&mut self, faces: &[RgbImage]) -> Result<Vec<Option<FaceEmbedding>>, String> {
162        let mut form = reqwest::blocking::multipart::Form::new();
163        for (i, face) in faces.iter().enumerate() {
164            let mut buf: Vec<u8> = Vec::with_capacity(8 * 1024);
165            {
166                let mut encoder =
167                    JpegEncoder::new_with_quality(Cursor::new(&mut buf), JPEG_QUALITY);
168                encoder
169                    .encode(
170                        face.as_raw(),
171                        face.width(),
172                        face.height(),
173                        image::ExtendedColorType::Rgb8,
174                    )
175                    .map_err(|e| format!("JPEG encode face[{}]: {}", i, e))?;
176            }
177            let part = reqwest::blocking::multipart::Part::bytes(buf)
178                .file_name(format!("{}.jpg", i))
179                .mime_str("image/jpeg")
180                .map_err(|e| format!("mime error: {}", e))?;
181            // FastAPI's `files: list[UploadFile] = File(...)` expects
182            // repeated multipart fields with the same name.
183            form = form.part("files", part);
184        }
185
186        let resp = self
187            .client
188            .post(format!("{}/embed", self.base_url))
189            .multipart(form)
190            .send()
191            .map_err(|e| format!("POST failed: {}", e))?;
192
193        if !resp.status().is_success() {
194            return Err(format!("server returned {}", resp.status()));
195        }
196
197        let json: serde_json::Value = resp.json().map_err(|e| format!("JSON parse: {}", e))?;
198
199        let arr = json
200            .get("embeddings")
201            .and_then(|v: &serde_json::Value| v.as_array())
202            .ok_or_else(|| "missing 'embeddings' array in response".to_string())?;
203
204        let mut results: Vec<Option<FaceEmbedding>> = Vec::with_capacity(faces.len());
205        for item in arr.iter() {
206            let item: &serde_json::Value = item;
207            let vec: Vec<f32> = item
208                .as_array()
209                .ok_or_else(|| "embedding entry is not an array".to_string())?
210                .iter()
211                .map(|v: &serde_json::Value| {
212                    v.as_f64()
213                        .map(|f| f as f32)
214                        .ok_or_else(|| "non-f32 value".to_string())
215                })
216                .collect::<Result<Vec<f32>, String>>()?;
217            if vec.len() != 512 {
218                return Err(format!("expected 512-d embedding, got {}-d", vec.len()));
219            }
220            let emb = FaceEmbedding::new(ndarray::Array1::from_vec(vec));
221            results.push(Some(emb));
222        }
223
224        // If the server returned fewer embeddings than faces, pad with
225        // None. In practice they should always match, but be defensive.
226        while results.len() < faces.len() {
227            results.push(None);
228        }
229
230        Ok(results)
231    }
232}
233
234/// Normalize and compare model names. Both sides are case-insensitive
235/// and `.onnx` suffix is optional.
236fn models_match(a: &str, b: &str) -> bool {
237    fn norm(s: &str) -> String {
238        s.trim()
239            .strip_suffix(".onnx")
240            .unwrap_or(s.trim())
241            .to_ascii_lowercase()
242    }
243    !a.is_empty() && !b.is_empty() && norm(a) == norm(b)
244}
245
246#[cfg(test)]
247mod tests {
248    use super::*;
249
250    #[test]
251    fn model_match_handles_onnx_suffix_and_case() {
252        assert!(models_match(
253            "adaface_ir101_webface12m",
254            "adaface_ir101_webface12m.onnx"
255        ));
256        assert!(models_match(
257            "ADAFACE_IR101_WEBFACE12M.onnx",
258            "adaface_ir101_webface12m"
259        ));
260        assert!(!models_match("other", "adaface_ir101_webface12m"));
261        assert!(!models_match("", "adaface_ir101_webface12m"));
262    }
263}