1use std::collections::{HashMap, HashSet};
4
5use rusqlite::{params, params_from_iter, Result as SqliteResult};
6
7use crate::ml::FaceEmbedding;
8
9use super::{FaceClusterRecord, FaceRepo};
10
11impl<'a> FaceRepo<'a> {
12 pub fn refresh_all_galleries(&self) -> SqliteResult<()> {
13 let tx = self.conn.unchecked_transaction()?;
14 let mut ids = Vec::new();
15 {
16 let mut stmt = tx.prepare("SELECT id FROM face_clusters")?;
17 let rows = stmt.query_map([], |row| row.get::<_, i64>(0))?;
18 for row in rows {
19 ids.push(row?);
20 }
21 }
22
23 tx.execute("DELETE FROM person_gallery_embeddings", [])?;
24 for cluster_id in ids {
25 Self::refresh_gallery_tx(&tx, cluster_id)?;
26 }
27 tx.commit()
28 }
29
30 pub(crate) fn refresh_gallery_tx(
31 tx: &rusqlite::Transaction<'_>,
32 cluster_id: i64,
33 ) -> SqliteResult<()> {
34 let mut sticky: Vec<(i64, FaceEmbedding)> = Vec::new();
37 {
38 let mut stmt = tx.prepare(
39 r#"
40 SELECT face_id, embedding
41 FROM person_gallery_embeddings
42 WHERE cluster_id = ?1 AND source = 'user_confirmed'
43 "#,
44 )?;
45 let rows = stmt.query_map(params![cluster_id], |row| {
46 Ok((row.get::<_, i64>(0)?, row.get::<_, Vec<u8>>(1)?))
47 })?;
48 for row in rows {
49 let (face_id, bytes) = row?;
50 if let Some(emb) = FaceEmbedding::from_bytes(&bytes) {
51 sticky.push((face_id, emb));
52 }
53 }
54 }
55
56 tx.execute(
57 "DELETE FROM person_gallery_embeddings WHERE cluster_id = ?1 AND source != 'user_confirmed'",
58 params![cluster_id],
59 )?;
60
61 let sticky_ids: std::collections::HashSet<i64> = sticky.iter().map(|(id, _)| *id).collect();
62
63 let mut stmt = tx.prepare(
64 r#"
65 SELECT id, embedding, confidence
66 FROM faces
67 WHERE cluster_id = ?1
68 ORDER BY confidence DESC, id ASC
69 "#,
70 )?;
71
72 let rows = stmt.query_map(params![cluster_id], |row| {
73 Ok((
74 row.get::<_, i64>(0)?,
75 row.get::<_, Vec<u8>>(1)?,
76 row.get::<_, Option<f32>>(2)?.unwrap_or(0.0),
77 ))
78 })?;
79
80 const MAX_GALLERY: usize = 30;
81 const DIVERSITY_THRESHOLD: f32 = 0.70;
82
83 let mut auto: Vec<(i64, FaceEmbedding, f32)> = Vec::new();
85 for row in rows {
86 let (face_id, bytes, confidence) = row?;
87 if sticky_ids.contains(&face_id) {
88 continue;
89 }
90 let Some(emb) = FaceEmbedding::from_bytes(&bytes) else {
91 tracing::warn!(
92 "Corrupted embedding in refresh_gallery for face_id={}: {} bytes",
93 face_id,
94 bytes.len()
95 );
96 continue;
97 };
98
99 let total = sticky.len() + auto.len();
101 if total < MAX_GALLERY {
102 auto.push((face_id, emb, confidence));
103 continue;
104 }
105
106 let mut min_similarity = 1.0f32;
108 for (_, existing) in &sticky {
109 let sim = emb.cosine_similarity(existing);
110 if sim < min_similarity {
111 min_similarity = sim;
112 }
113 }
114 for (_, existing, _) in &auto {
115 let sim = emb.cosine_similarity(existing);
116 if sim < min_similarity {
117 min_similarity = sim;
118 }
119 }
120
121 if min_similarity < DIVERSITY_THRESHOLD {
122 let mut replace_idx = 0usize;
124 let mut replace_score = 2.0f32;
125 for (idx, (_, existing, _)) in auto.iter().enumerate() {
126 let mut avg = 0.0f32;
127 let mut cnt = 0.0f32;
128 for (_, other) in &sticky {
129 avg += existing.cosine_similarity(other);
130 cnt += 1.0;
131 }
132 for (j, (_, other, _)) in auto.iter().enumerate() {
133 if idx == j {
134 continue;
135 }
136 avg += existing.cosine_similarity(other);
137 cnt += 1.0;
138 }
139 if cnt > 0.0 {
140 avg /= cnt;
141 }
142 if avg < replace_score {
143 replace_score = avg;
144 replace_idx = idx;
145 }
146 }
147 if !auto.is_empty() {
148 auto[replace_idx] = (face_id, emb, confidence);
149 }
150 }
151 }
152
153 for (face_id, emb, quality_score) in auto {
154 tx.execute(
155 r#"
156 INSERT INTO person_gallery_embeddings (cluster_id, face_id, embedding, quality_score, source)
157 VALUES (?1, ?2, ?3, ?4, 'auto')
158 "#,
159 params![cluster_id, face_id, emb.to_bytes(), quality_score],
160 )?;
161 }
162
163 Ok(())
164 }
165
166 pub(crate) fn refresh_cluster_stats_tx(
167 tx: &rusqlite::Transaction<'_>,
168 cluster_id: i64,
169 ) -> SqliteResult<()> {
170 tx.execute(
171 r#"
172 UPDATE face_clusters SET
173 face_count = (
174 SELECT COUNT(*)
175 FROM faces f
176 JOIN photos p ON p.id = f.photo_id
177 WHERE f.cluster_id = ?1 AND p.is_trashed = FALSE
178 ),
179 photo_count = (
180 SELECT COUNT(DISTINCT photo_id)
181 FROM (
182 SELECT f.photo_id
183 FROM faces f
184 JOIN photos p ON p.id = f.photo_id
185 WHERE f.cluster_id = ?1 AND p.is_trashed = FALSE
186 UNION
187 SELECT pii.photo_id
188 FROM photo_inferred_identities pii
189 JOIN photos p ON p.id = pii.photo_id
190 WHERE pii.cluster_id = ?1 AND p.is_trashed = FALSE
191 )
192 ),
193 representative_face_id = (
194 SELECT f.id
195 FROM faces f
196 JOIN photos p ON p.id = f.photo_id
197 WHERE f.cluster_id = ?1 AND p.is_trashed = FALSE
198 ORDER BY f.confidence DESC
199 LIMIT 1
200 ),
201 updated_at = CURRENT_TIMESTAMP
202 WHERE id = ?1
203 "#,
204 params![cluster_id],
205 )?;
206
207 Ok(())
208 }
209
210 pub fn face_thumbnail_candidates(
216 &self,
217 cluster_ids: &[i64],
218 max_per_cluster: usize,
219 ) -> SqliteResult<HashMap<i64, Vec<i64>>> {
220 if cluster_ids.is_empty() || max_per_cluster == 0 {
221 return Ok(HashMap::new());
222 }
223
224 let mut out: HashMap<i64, Vec<i64>> = HashMap::new();
225 let mut seen = HashSet::new();
226 let mut unique_ids = Vec::new();
227 for id in cluster_ids {
228 if seen.insert(*id) {
229 unique_ids.push(*id);
230 }
231 }
232
233 for chunk in unique_ids.chunks(400) {
234 let placeholders = (1..=chunk.len())
235 .map(|idx| format!("?{}", idx))
236 .collect::<Vec<_>>()
237 .join(",");
238 let limit_param = chunk.len() + 1;
239 let sql = format!(
240 r#"
241 SELECT cluster_id, id
242 FROM (
243 SELECT
244 faces.cluster_id,
245 faces.id,
246 ROW_NUMBER() OVER (
247 PARTITION BY faces.cluster_id
248 ORDER BY faces.confidence DESC, faces.id ASC
249 ) AS rn
250 FROM faces
251 JOIN photos p ON p.id = faces.photo_id
252 WHERE faces.cluster_id IN ({})
253 AND p.is_trashed = FALSE
254 )
255 WHERE rn <= ?{}
256 ORDER BY cluster_id ASC, rn ASC
257 "#,
258 placeholders, limit_param
259 );
260
261 let mut values: Vec<rusqlite::types::Value> = chunk
262 .iter()
263 .map(|id| rusqlite::types::Value::from(*id))
264 .collect();
265 values.push(rusqlite::types::Value::from(max_per_cluster as i64));
266
267 let mut stmt = self.conn.prepare(&sql)?;
268 let rows = stmt.query_map(params_from_iter(values.iter()), |row| {
269 Ok((row.get::<_, i64>(0)?, row.get::<_, i64>(1)?))
270 })?;
271 for row in rows {
272 let (cluster_id, face_id) = row?;
273 out.entry(cluster_id).or_default().push(face_id);
274 }
275 }
276
277 Ok(out)
278 }
279
280 pub fn update_representative_faces(&self, updates: &[(i64, i64)]) -> SqliteResult<()> {
281 if updates.is_empty() {
282 return Ok(());
283 }
284
285 let tx = self.conn.unchecked_transaction()?;
286 {
287 let mut stmt = tx.prepare(
288 "UPDATE face_clusters SET representative_face_id = ?1, updated_at = CURRENT_TIMESTAMP WHERE id = ?2",
289 )?;
290 for (cluster_id, face_id) in updates {
291 stmt.execute(params![face_id, cluster_id])?;
292 }
293 }
294 tx.commit()
295 }
296
297 pub fn normalize_cluster_stats(&self) -> SqliteResult<()> {
299 let tx = self.conn.unchecked_transaction()?;
300
301 let mut ids = Vec::new();
302 {
303 let mut stmt = tx.prepare("SELECT id FROM face_clusters")?;
304 let rows = stmt.query_map([], |row| row.get::<_, i64>(0))?;
305 for row in rows {
306 ids.push(row?);
307 }
308 }
309
310 for cluster_id in ids {
311 Self::refresh_cluster_stats_tx(&tx, cluster_id)?;
312 }
313
314 tx.execute(
315 "DELETE FROM face_clusters WHERE face_count <= 0 AND photo_count <= 0",
316 [],
317 )?;
318
319 tx.commit()
320 }
321
322 pub fn populate_face_thumbnails(
329 &self,
330 clusters: &mut [FaceClusterRecord],
331 drive_path: &std::path::Path,
332 ) -> SqliteResult<()> {
333 let faces_dir = drive_path.join(".photovault").join("faces");
334 for cluster in clusters.iter_mut() {
335 cluster.face_thumbnail_path = None;
336
337 if let Some(face_id) = cluster.representative_face_id {
338 let crop_path = faces_dir.join(format!("{}.jpg", face_id));
339 if crop_path.exists() {
340 cluster.face_thumbnail_path =
341 Some(format!(".photovault/faces/{}.jpg", face_id));
342 continue;
343 }
344 }
345
346 let mut replacement_face_id: Option<i64> = None;
347 let mut stmt = self.conn.prepare(
348 r#"
349 SELECT f.id
350 FROM faces f
351 JOIN photos p ON p.id = f.photo_id
352 WHERE f.cluster_id = ?1
353 AND p.is_trashed = FALSE
354 ORDER BY f.confidence DESC
355 "#,
356 )?;
357
358 let mut rows = stmt.query(params![cluster.id])?;
359 while let Some(row) = rows.next()? {
360 let face_id: i64 = row.get(0)?;
361 let crop_path = faces_dir.join(format!("{}.jpg", face_id));
362 if crop_path.exists() {
363 replacement_face_id = Some(face_id);
364 cluster.face_thumbnail_path =
365 Some(format!(".photovault/faces/{}.jpg", face_id));
366 break;
367 }
368 }
369
370 drop(rows);
371 drop(stmt);
372
373 if let Some(face_id) = replacement_face_id {
374 cluster.representative_face_id = Some(face_id);
375 self.conn.execute(
376 "UPDATE face_clusters SET representative_face_id = ?1, updated_at = CURRENT_TIMESTAMP WHERE id = ?2",
377 params![face_id, cluster.id],
378 )?;
379 }
380 }
381
382 Ok(())
383 }
384}
385
386#[cfg(test)]
387mod tests {
388 use rusqlite::{params, Connection};
389
390 use super::*;
391 use crate::db::create_schema;
392
393 fn seeded_conn() -> Connection {
394 let conn = Connection::open_in_memory().unwrap();
395 create_schema(&conn).unwrap();
396 for id in 1..=3 {
397 conn.execute(
398 "INSERT INTO photos (id, file_path, file_name, file_hash, file_size)
399 VALUES (?1, ?2, ?3, ?4, 100)",
400 params![
401 id,
402 format!("photos/{id}.jpg"),
403 format!("{id}.jpg"),
404 format!("hash-{id}")
405 ],
406 )
407 .unwrap();
408 }
409 conn.execute(
410 "INSERT INTO face_clusters (id, name, face_count, photo_count)
411 VALUES (10, 'A', 2, 2), (20, 'B', 1, 1)",
412 [],
413 )
414 .unwrap();
415 for (id, cluster_id, confidence) in [(1, 10, 0.5), (2, 10, 0.9), (3, 20, 0.7)] {
416 conn.execute(
417 "INSERT INTO faces (
418 id, photo_id, bbox_x, bbox_y, bbox_width, bbox_height,
419 embedding, cluster_id, confidence, user_confirmed
420 )
421 VALUES (?1, ?1, 0.1, 0.1, 0.2, 0.2, ?2, ?3, ?4, 0)",
422 params![id, vec![id as u8; 16], cluster_id, confidence],
423 )
424 .unwrap();
425 }
426 conn
427 }
428
429 #[test]
430 fn face_thumbnail_candidates_are_batched_and_confidence_ordered() {
431 let conn = seeded_conn();
432 let repo = FaceRepo::new(&conn);
433
434 let candidates = repo.face_thumbnail_candidates(&[10, 20, 10], 2).unwrap();
435
436 assert_eq!(candidates.get(&10).unwrap(), &vec![2, 1]);
437 assert_eq!(candidates.get(&20).unwrap(), &vec![3]);
438 }
439
440 #[test]
441 fn face_thumbnail_candidates_ignore_trashed_photos() {
442 let conn = seeded_conn();
443 conn.execute("UPDATE photos SET is_trashed = TRUE WHERE id = 2", [])
444 .unwrap();
445 let repo = FaceRepo::new(&conn);
446
447 let candidates = repo.face_thumbnail_candidates(&[10], 2).unwrap();
448
449 assert_eq!(candidates.get(&10).unwrap(), &vec![1]);
450 }
451
452 #[test]
453 fn update_representative_faces_updates_all_rows_in_one_call() {
454 let conn = seeded_conn();
455 let repo = FaceRepo::new(&conn);
456
457 repo.update_representative_faces(&[(10, 2), (20, 3)])
458 .unwrap();
459
460 let reps: Vec<(i64, i64)> = conn
461 .prepare("SELECT id, representative_face_id FROM face_clusters ORDER BY id")
462 .unwrap()
463 .query_map([], |row| Ok((row.get(0)?, row.get(1)?)))
464 .unwrap()
465 .collect::<Result<_, _>>()
466 .unwrap();
467 assert_eq!(reps, vec![(10, 2), (20, 3)]);
468 }
469}