Skip to main content

smriti/db/face_repo/
gallery.rs

1//! Gallery and stats management for FaceRepo.
2
3use 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        // User-confirmed gallery members are sticky: never evicted by diversity
35        // replacement. Auto-selected members are rebuilt from scratch each call.
36        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        // Auto members can't include faces already sticky.
84        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            // Seed with the top-N by confidence until we hit the cap.
100            let total = sticky.len() + auto.len();
101            if total < MAX_GALLERY {
102                auto.push((face_id, emb, confidence));
103                continue;
104            }
105
106            // Check diversity against both sticky and auto members.
107            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                // Replace the most redundant *auto* entry (sticky ones are untouchable).
123                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    /// Return fallback face thumbnail candidates for the supplied clusters.
211    ///
212    /// The rows are ordered by cluster, then by face confidence descending.
213    /// Callers can check the corresponding crop files outside the shared DB
214    /// mutex and then persist any representative replacement in one batch.
215    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    /// Recompute all cluster stats and prune empty clusters.
298    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    /// Populate face thumbnail paths on cluster records.
323    ///
324    /// Call this after `get_all_clusters()` with the drive root path. The
325    /// path written into each cluster is **relative to drive_root** (e.g.
326    /// `.photovault/faces/42.jpg`) so it round-trips through the same
327    /// frontend `thumbUrl()` helper as `photos.thumbnail_path`.
328    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}