smriti/ml/
remote_embedder.rs1use 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 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 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 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 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 while results.len() < faces.len() {
227 results.push(None);
228 }
229
230 Ok(results)
231 }
232}
233
234fn 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}