1use rusqlite::{params, params_from_iter, Result as SqliteResult};
4
5use crate::ml::FaceEmbedding;
6
7use super::{FaceClusterRecord, FaceDetail, FaceRepo, FaceStatus, GalleryEmbedding, ReviewItem};
8
9type FacePathRow = (i64, String, i32, f32, f32, f32, f32);
10type UnprocessedPhotoRow = (i64, String, i32, Option<i64>, String);
17
18impl<'a> FaceRepo<'a> {
19 pub fn count_pending_face_processing(&self) -> SqliteResult<i64> {
22 self.conn.query_row(
23 "SELECT COUNT(*) FROM photos WHERE faces_processed = FALSE AND is_trashed = FALSE",
24 [],
25 |row| row.get(0),
26 )
27 }
28
29 pub fn get_unclustered_faces_with_photo_embeddings(
31 &self,
32 ) -> SqliteResult<Vec<(i64, i64, FaceEmbedding)>> {
33 let mut stmt = self.conn.prepare(
34 "SELECT f.id, f.photo_id, f.embedding
35 FROM faces f
36 JOIN photos p ON p.id = f.photo_id
37 WHERE f.cluster_id IS NULL
38 AND f.user_confirmed >= 0
39 AND p.is_trashed = FALSE",
40 )?;
41
42 let rows = stmt.query_map([], |row| {
43 let id: i64 = row.get(0)?;
44 let photo_id: i64 = row.get(1)?;
45 let bytes: Vec<u8> = row.get(2)?;
46 Ok((id, photo_id, bytes))
47 })?;
48
49 let mut faces = Vec::new();
50 for row in rows {
51 let (id, photo_id, bytes) = row?;
52 match FaceEmbedding::from_bytes(&bytes) {
53 Some(emb) => faces.push((id, photo_id, emb)),
54 None => tracing::warn!(
55 "Corrupted face embedding for face_id={}: {} bytes",
56 id,
57 bytes.len()
58 ),
59 }
60 }
61
62 Ok(faces)
63 }
64
65 pub fn get_all_clusters(&self) -> SqliteResult<Vec<FaceClusterRecord>> {
67 let mut stmt = self.conn.prepare(
68 r#"
69 SELECT id, name, representative_face_id, photo_count
70 FROM face_clusters
71 ORDER BY photo_count DESC, face_count DESC
72 "#,
73 )?;
74
75 let rows = stmt.query_map([], |row| {
76 Ok(FaceClusterRecord {
77 id: row.get(0)?,
78 name: row.get(1)?,
79 representative_face_id: row.get(2)?,
80 photo_count: row.get(3)?,
81 face_thumbnail_path: None, })
83 })?;
84
85 let mut clusters = Vec::new();
86 for row in rows {
87 clusters.push(row?);
88 }
89
90 Ok(clusters)
91 }
92
93 pub fn get_gallery_embeddings(&self) -> SqliteResult<Vec<GalleryEmbedding>> {
94 let mut stmt = self.conn.prepare(
95 "SELECT cluster_id, face_id, embedding FROM person_gallery_embeddings ORDER BY cluster_id",
96 )?;
97
98 let rows = stmt.query_map([], |row| {
99 Ok((
100 row.get::<_, i64>(0)?,
101 row.get::<_, i64>(1)?,
102 row.get::<_, Vec<u8>>(2)?,
103 ))
104 })?;
105
106 let mut result = Vec::new();
107 for row in rows {
108 let (cluster_id, face_id, bytes) = row?;
109 match FaceEmbedding::from_bytes(&bytes) {
110 Some(embedding) => {
111 result.push(GalleryEmbedding {
112 cluster_id,
113 face_id,
114 embedding,
115 });
116 }
117 None => tracing::warn!(
118 "Corrupted gallery embedding for cluster_id={}, face_id={}: {} bytes",
119 cluster_id,
120 face_id,
121 bytes.len()
122 ),
123 }
124 }
125
126 Ok(result)
127 }
128
129 pub fn get_cluster_photo_ids(&self) -> SqliteResult<Vec<(i64, i64)>> {
130 let mut stmt = self.conn.prepare(
131 "SELECT DISTINCT cluster_id, photo_id FROM faces WHERE cluster_id IS NOT NULL",
132 )?;
133
134 let rows = stmt.query_map([], |row| Ok((row.get::<_, i64>(0)?, row.get::<_, i64>(1)?)))?;
135 let mut result = Vec::new();
136 for row in rows {
137 result.push(row?);
138 }
139 Ok(result)
140 }
141
142 pub fn get_people_for_photo(&self, photo_id: i64) -> SqliteResult<Vec<(i64, String)>> {
145 let mut stmt = self.conn.prepare(
146 r#"
147 SELECT DISTINCT cluster_id, name FROM (
148 SELECT fc.id AS cluster_id, COALESCE(fc.name, 'Person ' || fc.id) AS name
149 FROM faces f
150 JOIN face_clusters fc ON f.cluster_id = fc.id
151 WHERE f.photo_id = ?1
152
153 UNION
154
155 SELECT fc.id AS cluster_id, COALESCE(fc.name, 'Person ' || fc.id) AS name
156 FROM photo_inferred_identities pii
157 JOIN face_clusters fc ON pii.cluster_id = fc.id
158 WHERE pii.photo_id = ?1
159 )
160 ORDER BY name
161 "#,
162 )?;
163
164 let people = stmt
165 .query_map(params![photo_id], |row| {
166 Ok((row.get::<_, i64>(0)?, row.get::<_, String>(1)?))
167 })?
168 .collect::<SqliteResult<Vec<_>>>()?;
169
170 Ok(people)
171 }
172
173 pub fn get_photos_for_cluster(&self, cluster_id: i64) -> SqliteResult<Vec<i64>> {
175 let mut stmt = self.conn.prepare(
176 r#"
177 SELECT p.id
178 FROM photos p
179 JOIN (
180 SELECT f.photo_id
181 FROM faces f
182 WHERE f.cluster_id = ?1
183 UNION
184 SELECT pii.photo_id
185 FROM photo_inferred_identities pii
186 WHERE pii.cluster_id = ?1
187 ) matches ON matches.photo_id = p.id
188 WHERE p.is_trashed = FALSE
189 GROUP BY p.id
190 ORDER BY p.date_taken IS NULL ASC, p.date_taken DESC, p.id DESC
191 "#,
192 )?;
193
194 let rows = stmt.query_map(params![cluster_id], |row| row.get(0))?;
195
196 let mut photo_ids = Vec::new();
197 for row in rows {
198 photo_ids.push(row?);
199 }
200
201 Ok(photo_ids)
202 }
203
204 pub fn get_all_faces_with_paths(&self) -> SqliteResult<Vec<FacePathRow>> {
207 let mut stmt = self.conn.prepare(
208 r#"
209 SELECT f.id, p.file_path, COALESCE(p.orientation, 1), f.bbox_x, f.bbox_y, f.bbox_width, f.bbox_height
210 FROM faces f
211 JOIN photos p ON f.photo_id = p.id
212 "#,
213 )?;
214
215 let rows = stmt.query_map([], |row| {
216 Ok((
217 row.get::<_, i64>(0)?,
218 row.get::<_, String>(1)?,
219 row.get::<_, i32>(2)?,
220 row.get::<_, f32>(3)?,
221 row.get::<_, f32>(4)?,
222 row.get::<_, f32>(5)?,
223 row.get::<_, f32>(6)?,
224 ))
225 })?;
226
227 let mut result = Vec::new();
228 for row in rows {
229 result.push(row?);
230 }
231
232 Ok(result)
233 }
234
235 pub fn get_contextual_cluster_candidates(
237 &self,
238 photo_id: i64,
239 folder_prefix_like: Option<&str>,
240 target_ts: i64,
241 window_secs: i64,
242 ) -> SqliteResult<Vec<(i64, i64, i64, String)>> {
243 let sql_with_folder = r#"
244 SELECT DISTINCT
245 p.id,
246 f.cluster_id,
247 CAST(strftime('%s', p.date_taken) AS INTEGER) AS source_ts,
248 p.file_path
249 FROM photos p
250 JOIN faces f ON f.photo_id = p.id
251 WHERE p.id != ?1
252 AND p.is_trashed = FALSE
253 AND p.date_taken IS NOT NULL
254 AND f.cluster_id IS NOT NULL
255 AND p.file_path LIKE ?2
256 AND ABS(CAST(strftime('%s', p.date_taken) AS INTEGER) - ?3) <= ?4
257 "#;
258
259 let sql_without_folder = r#"
260 SELECT DISTINCT
261 p.id,
262 f.cluster_id,
263 CAST(strftime('%s', p.date_taken) AS INTEGER) AS source_ts,
264 p.file_path
265 FROM photos p
266 JOIN faces f ON f.photo_id = p.id
267 WHERE p.id != ?1
268 AND p.is_trashed = FALSE
269 AND p.date_taken IS NOT NULL
270 AND f.cluster_id IS NOT NULL
271 AND ABS(CAST(strftime('%s', p.date_taken) AS INTEGER) - ?2) <= ?3
272 "#;
273
274 let mut result = Vec::new();
275
276 if let Some(folder_like) = folder_prefix_like {
277 let mut stmt = self.conn.prepare(sql_with_folder)?;
278 let rows = stmt.query_map(
279 params![photo_id, folder_like, target_ts, window_secs],
280 |row| {
281 Ok((
282 row.get::<_, i64>(0)?,
283 row.get::<_, i64>(1)?,
284 row.get::<_, i64>(2)?,
285 row.get::<_, String>(3)?,
286 ))
287 },
288 )?;
289 for row in rows {
290 result.push(row?);
291 }
292 } else {
293 let mut stmt = self.conn.prepare(sql_without_folder)?;
294 let rows = stmt.query_map(params![photo_id, target_ts, window_secs], |row| {
295 Ok((
296 row.get::<_, i64>(0)?,
297 row.get::<_, i64>(1)?,
298 row.get::<_, i64>(2)?,
299 row.get::<_, String>(3)?,
300 ))
301 })?;
302 for row in rows {
303 result.push(row?);
304 }
305 }
306
307 Ok(result)
308 }
309
310 pub fn get_unprocessed_photos_with_context(&self) -> SqliteResult<Vec<UnprocessedPhotoRow>> {
312 let mut stmt = self.conn.prepare(
313 r#"
314 SELECT
315 id,
316 file_path,
317 COALESCE(orientation, 1) AS orientation,
318 CASE
319 WHEN date_taken IS NOT NULL THEN CAST(strftime('%s', date_taken) AS INTEGER)
320 ELSE NULL
321 END AS taken_ts,
322 file_hash
323 FROM photos
324 WHERE faces_processed = FALSE AND is_trashed = FALSE
325 ORDER BY date_taken DESC
326 "#,
327 )?;
328
329 let rows = stmt.query_map([], |row| {
330 Ok((
331 row.get::<_, i64>(0)?,
332 row.get::<_, String>(1)?,
333 row.get::<_, i32>(2)?,
334 row.get::<_, Option<i64>>(3)?,
335 row.get::<_, String>(4)?,
336 ))
337 })?;
338
339 let mut result = Vec::new();
340 for row in rows {
341 result.push(row?);
342 }
343 Ok(result)
344 }
345
346 pub fn get_photo_brightness(&self, photo_id: i64) -> SqliteResult<Option<f32>> {
351 self.conn
352 .query_row(
353 "SELECT brightness FROM photos WHERE id = ?1",
354 rusqlite::params![photo_id],
355 |row| row.get::<_, Option<f32>>(0),
356 )
357 .or_else(|e| match e {
358 rusqlite::Error::QueryReturnedNoRows => Ok(None),
359 other => Err(other),
360 })
361 }
362
363 pub fn review_queue_size(&self) -> SqliteResult<i64> {
365 self.conn.query_row(
366 "SELECT COUNT(*) FROM face_review_queue WHERE resolved_at IS NULL",
367 [],
368 |row| row.get(0),
369 )
370 }
371
372 pub fn get_review_queue_items(&self, limit: usize) -> SqliteResult<Vec<ReviewItem>> {
378 let mut stmt = self.conn.prepare(
379 r#"
380 SELECT
381 q.id, q.face_id,
382 q.candidate_cluster_id, c.name, c.face_count,
383 q.score
384 FROM face_review_queue q
385 JOIN faces f ON f.id = q.face_id
386 JOIN face_clusters c ON c.id = q.candidate_cluster_id
387 WHERE q.resolved_at IS NULL
388 ORDER BY COALESCE(q.ambiguity, 1.0) ASC, c.face_count DESC
389 LIMIT ?1
390 "#,
391 )?;
392
393 let rows = stmt.query_map(params![limit as i64], |row| {
394 Ok(ReviewItem {
395 queue_id: row.get(0)?,
396 face_id: row.get(1)?,
397 candidate_cluster_id: row.get(2)?,
398 candidate_cluster_name: row.get(3)?,
399 candidate_cluster_size: row.get(4)?,
400 candidate_sample_face_ids: Vec::new(),
401 score: row.get(5)?,
402 })
403 })?;
404
405 let mut items = Vec::new();
406 for row in rows {
407 items.push(row?);
408 }
409 drop(stmt);
410
411 let mut sample_stmt = self.conn.prepare(
413 r#"
414 SELECT id
415 FROM faces
416 WHERE cluster_id = ?1
417 ORDER BY confidence DESC, id ASC
418 LIMIT 4
419 "#,
420 )?;
421 for item in items.iter_mut() {
422 let ids = sample_stmt.query_map(params![item.candidate_cluster_id], |row| {
423 row.get::<_, i64>(0)
424 })?;
425 let mut collected = Vec::new();
426 for id in ids {
427 collected.push(id?);
428 }
429 item.candidate_sample_face_ids = collected;
430 }
431
432 Ok(items)
433 }
434
435 pub fn get_photo_other_clusters(
438 &self,
439 photo_id: i64,
440 exclude_face_id: i64,
441 ) -> SqliteResult<Vec<i64>> {
442 let mut stmt = self.conn.prepare(
443 "SELECT DISTINCT cluster_id FROM faces
444 WHERE photo_id = ?1 AND id != ?2 AND cluster_id IS NOT NULL",
445 )?;
446 let rows = stmt.query_map(params![photo_id, exclude_face_id], |row| {
447 row.get::<_, i64>(0)
448 })?;
449 let mut out = Vec::new();
450 for row in rows {
451 out.push(row?);
452 }
453 Ok(out)
454 }
455
456 pub fn cooccurrence_count(&self, cluster_a: i64, cluster_b: i64) -> SqliteResult<i64> {
462 if cluster_a == cluster_b {
463 return Ok(0);
464 }
465 self.conn.query_row(
466 r#"
467 SELECT COUNT(*) FROM (
468 SELECT photo_id FROM faces WHERE cluster_id = ?1
469 INTERSECT
470 SELECT photo_id FROM faces WHERE cluster_id = ?2
471 )
472 "#,
473 params![cluster_a, cluster_b],
474 |row| row.get(0),
475 )
476 }
477
478 pub fn temporal_neighbor_clusters(
485 &self,
486 photo_id: i64,
487 window_secs: i64,
488 ) -> SqliteResult<Vec<(i64, i64)>> {
489 let mut stmt = self.conn.prepare(
490 r#"
491 WITH base AS (
492 SELECT date_taken AS t
493 FROM photos
494 WHERE id = ?1 AND date_taken IS NOT NULL
495 )
496 SELECT DISTINCT f.cluster_id,
497 CAST(ABS(
498 strftime('%s', p.date_taken) - strftime('%s', (SELECT t FROM base))
499 ) AS INTEGER) AS delta_sec
500 FROM faces f
501 JOIN photos p ON p.id = f.photo_id
502 WHERE f.cluster_id IS NOT NULL
503 AND p.id != ?1
504 AND p.is_trashed = FALSE
505 AND p.date_taken IS NOT NULL
506 AND ABS(
507 strftime('%s', p.date_taken) - strftime('%s', (SELECT t FROM base))
508 ) <= ?2
509 "#,
510 )?;
511 let rows = stmt.query_map(params![photo_id, window_secs], |row| {
512 Ok((row.get::<_, i64>(0)?, row.get::<_, i64>(1)?))
513 })?;
514 let mut out = Vec::new();
515 for row in rows {
516 out.push(row?);
517 }
518 Ok(out)
519 }
520
521 pub fn get_faces_by_cluster(
523 &self,
524 cluster_id: i64,
525 status: FaceStatus,
526 cursor: Option<i64>,
527 limit: usize,
528 ) -> SqliteResult<Vec<FaceDetail>> {
529 let (where_clause, params_slice): (String, Vec<rusqlite::types::Value>) = match status {
530 FaceStatus::Confirmed => (
531 "f.cluster_id = ?1 AND f.user_confirmed = 1".to_string(),
532 vec![rusqlite::types::Value::from(cluster_id)],
533 ),
534 FaceStatus::Unconfirmed => (
535 "f.cluster_id = ?1 AND f.user_confirmed = 0".to_string(),
536 vec![rusqlite::types::Value::from(cluster_id)],
537 ),
538 FaceStatus::All => (
539 "f.cluster_id = ?1".to_string(),
540 vec![rusqlite::types::Value::from(cluster_id)],
541 ),
542 };
543
544 let mut sql = format!(
545 "SELECT f.id, f.photo_id, f.cluster_id, f.confidence, f.user_confirmed
546 FROM faces f
547 WHERE {} ",
548 where_clause
549 );
550
551 let mut params: Vec<rusqlite::types::Value> = params_slice;
552 if let Some(c) = cursor {
553 sql.push_str("AND f.id > ?2 ");
554 params.push(rusqlite::types::Value::from(c));
555 sql.push_str("ORDER BY f.id ASC LIMIT ?3");
556 } else {
557 sql.push_str("ORDER BY f.id ASC LIMIT ?2");
558 }
559 params.push(rusqlite::types::Value::from(limit as i64));
560
561 let param_refs: Vec<&dyn rusqlite::types::ToSql> = params
562 .iter()
563 .map(|v| v as &dyn rusqlite::types::ToSql)
564 .collect();
565
566 let mut stmt = self.conn.prepare(&sql)?;
567 let rows = stmt.query_map(param_refs.as_slice(), |row| {
568 Ok(FaceDetail {
569 face_id: row.get(0)?,
570 photo_id: row.get(1)?,
571 cluster_id: row.get(2)?,
572 confidence: row.get(3)?,
573 user_confirmed: row.get(4)?,
574 })
575 })?;
576
577 let mut faces = Vec::new();
578 for row in rows {
579 faces.push(row?);
580 }
581 Ok(faces)
582 }
583
584 pub fn get_unclustered_faces(
588 &self,
589 cursor: Option<i64>,
590 limit: usize,
591 ) -> SqliteResult<Vec<FaceDetail>> {
592 let mut sql = String::from(
593 "SELECT f.id, f.photo_id, f.cluster_id, f.confidence, f.user_confirmed
594 FROM faces f
595 JOIN photos p ON p.id = f.photo_id
596 WHERE f.cluster_id IS NULL
597 AND p.is_trashed = FALSE ",
598 );
599 let mut params: Vec<rusqlite::types::Value> = Vec::new();
600 if let Some(c) = cursor {
601 sql.push_str("AND f.id > ?1 ORDER BY f.id ASC LIMIT ?2");
602 params.push(rusqlite::types::Value::from(c));
603 params.push(rusqlite::types::Value::from(limit as i64));
604 } else {
605 sql.push_str("ORDER BY f.id ASC LIMIT ?1");
606 params.push(rusqlite::types::Value::from(limit as i64));
607 }
608 let param_refs: Vec<&dyn rusqlite::types::ToSql> = params
609 .iter()
610 .map(|v| v as &dyn rusqlite::types::ToSql)
611 .collect();
612
613 let mut stmt = self.conn.prepare(&sql)?;
614 let rows = stmt.query_map(param_refs.as_slice(), |row| {
615 Ok(FaceDetail {
616 face_id: row.get(0)?,
617 photo_id: row.get(1)?,
618 cluster_id: row.get(2)?,
619 confidence: row.get(3)?,
620 user_confirmed: row.get(4)?,
621 })
622 })?;
623
624 let mut faces = Vec::new();
625 for row in rows {
626 faces.push(row?);
627 }
628 Ok(faces)
629 }
630
631 pub fn count_unconfirmed_in_cluster(&self, cluster_id: i64) -> SqliteResult<i64> {
633 self.conn.query_row(
634 "SELECT COUNT(*) FROM faces WHERE cluster_id = ?1 AND user_confirmed = 0",
635 params![cluster_id],
636 |row| row.get(0),
637 )
638 }
639
640 pub fn count_unconfirmed_global(&self) -> SqliteResult<(i64, i64)> {
643 let total: i64 = self.conn.query_row(
644 "SELECT COUNT(*) FROM faces WHERE cluster_id IS NOT NULL AND user_confirmed = 0",
645 [],
646 |row| row.get(0),
647 )?;
648 let cluster_count: i64 = self.conn.query_row(
649 "SELECT COUNT(DISTINCT cluster_id) FROM faces WHERE cluster_id IS NOT NULL AND user_confirmed = 0",
650 [],
651 |row| row.get(0),
652 )?;
653 Ok((total, cluster_count))
654 }
655
656 pub fn next_unconfirmed_face_batch(&self, limit: usize) -> SqliteResult<Vec<FaceDetail>> {
660 self.next_unconfirmed_face_batch_excluding(limit, &[])
661 }
662
663 pub fn next_unconfirmed_face_batch_excluding(
669 &self,
670 limit: usize,
671 excluded_face_ids: &[i64],
672 ) -> SqliteResult<Vec<FaceDetail>> {
673 if excluded_face_ids.is_empty() {
674 return self.next_unconfirmed_face_batch_without_exclusions(limit);
675 }
676
677 let placeholders = vec!["?"; excluded_face_ids.len()].join(", ");
678 let sql = format!(
679 r#"
680 WITH next_cluster AS (
681 SELECT cluster_id
682 FROM faces
683 WHERE cluster_id IS NOT NULL
684 AND user_confirmed = 0
685 AND id NOT IN ({placeholders})
686 GROUP BY cluster_id
687 ORDER BY COUNT(*) DESC, MIN(id) ASC
688 LIMIT 1
689 )
690 SELECT f.id, f.photo_id, f.cluster_id, f.confidence, f.user_confirmed
691 FROM faces f
692 JOIN next_cluster n ON n.cluster_id = f.cluster_id
693 WHERE f.user_confirmed = 0
694 AND f.id NOT IN ({placeholders})
695 ORDER BY f.id ASC
696 LIMIT ?
697 "#
698 );
699 let mut args = Vec::with_capacity(excluded_face_ids.len() * 2 + 1);
700 args.extend_from_slice(excluded_face_ids);
701 args.extend_from_slice(excluded_face_ids);
702 args.push(limit as i64);
703
704 let mut stmt = self.conn.prepare(&sql)?;
705 let rows = stmt.query_map(params_from_iter(args), |row| {
706 Ok(FaceDetail {
707 face_id: row.get(0)?,
708 photo_id: row.get(1)?,
709 cluster_id: row.get(2)?,
710 confidence: row.get(3)?,
711 user_confirmed: row.get(4)?,
712 })
713 })?;
714
715 let mut faces = Vec::new();
716 for row in rows {
717 faces.push(row?);
718 }
719 Ok(faces)
720 }
721
722 fn next_unconfirmed_face_batch_without_exclusions(
723 &self,
724 limit: usize,
725 ) -> SqliteResult<Vec<FaceDetail>> {
726 let mut stmt = self.conn.prepare(
727 r#"
728 WITH next_cluster AS (
729 SELECT cluster_id
730 FROM faces
731 WHERE cluster_id IS NOT NULL
732 AND user_confirmed = 0
733 GROUP BY cluster_id
734 ORDER BY COUNT(*) DESC, MIN(id) ASC
735 LIMIT 1
736 )
737 SELECT f.id, f.photo_id, f.cluster_id, f.confidence, f.user_confirmed
738 FROM faces f
739 JOIN next_cluster n ON n.cluster_id = f.cluster_id
740 WHERE f.user_confirmed = 0
741 ORDER BY f.id ASC
742 LIMIT ?1
743 "#,
744 )?;
745 let rows = stmt.query_map(params![limit as i64], |row| {
746 Ok(FaceDetail {
747 face_id: row.get(0)?,
748 photo_id: row.get(1)?,
749 cluster_id: row.get(2)?,
750 confidence: row.get(3)?,
751 user_confirmed: row.get(4)?,
752 })
753 })?;
754
755 let mut faces = Vec::new();
756 for row in rows {
757 faces.push(row?);
758 }
759 Ok(faces)
760 }
761
762 pub fn get_negatives_for_face(&self, face_id: i64) -> SqliteResult<Vec<i64>> {
764 let mut stmt = self
765 .conn
766 .prepare("SELECT not_cluster_id FROM face_negatives WHERE face_id = ?1")?;
767 let rows = stmt.query_map(params![face_id], |row| row.get::<_, i64>(0))?;
768 let mut out = Vec::new();
769 for row in rows {
770 out.push(row?);
771 }
772 Ok(out)
773 }
774
775 pub fn k_similar_to_cluster(&self, cluster_id: i64, k: usize) -> SqliteResult<Vec<(i64, f32)>> {
777 if k == 0 {
778 return Ok(Vec::new());
779 }
780
781 let mut stmt = self.conn.prepare(
783 "SELECT embedding FROM person_gallery_embeddings WHERE cluster_id = ?1 AND source = 'user_confirmed'",
784 )?;
785 let rows = stmt.query_map(params![cluster_id], |row| row.get::<_, Vec<u8>>(0))?;
786 let mut embeddings: Vec<FaceEmbedding> = Vec::new();
787 for row in rows {
788 let bytes = row?;
789 if let Some(emb) = FaceEmbedding::from_bytes(&bytes) {
790 embeddings.push(emb);
791 }
792 }
793
794 if embeddings.is_empty() {
796 let mut stmt2 = self
797 .conn
798 .prepare("SELECT embedding FROM person_gallery_embeddings WHERE cluster_id = ?1")?;
799 let rows2 = stmt2.query_map(params![cluster_id], |row| row.get::<_, Vec<u8>>(0))?;
800 for row in rows2 {
801 let bytes = row?;
802 if let Some(emb) = FaceEmbedding::from_bytes(&bytes) {
803 embeddings.push(emb);
804 }
805 }
806 }
807
808 if embeddings.is_empty() {
809 return Ok(Vec::new());
810 }
811
812 let dim = embeddings[0].vector.len();
814 let mut centroid = ndarray::Array1::<f32>::zeros(dim);
815 for emb in &embeddings {
816 centroid += &emb.vector;
817 }
818 centroid /= embeddings.len() as f32;
819
820 let norm = centroid.dot(¢roid).sqrt();
822 if norm > 0.0 {
823 centroid /= norm;
824 }
825 let centroid_emb = FaceEmbedding::new(centroid);
826
827 let mut cand_stmt = self.conn.prepare(
830 r#"
831 SELECT f.id, f.embedding
832 FROM faces f
833 JOIN photos p ON p.id = f.photo_id
834 WHERE (f.cluster_id IS NULL OR (f.cluster_id IS NOT NULL AND f.user_confirmed = 0))
835 AND f.user_confirmed >= 0
836 AND p.is_trashed = FALSE
837 AND f.id NOT IN (
838 SELECT face_id FROM face_negatives WHERE not_cluster_id = ?1
839 )
840 "#,
841 )?;
842 let cand_rows = cand_stmt.query_map(params![cluster_id], |row| {
843 Ok((row.get::<_, i64>(0)?, row.get::<_, Vec<u8>>(1)?))
844 })?;
845
846 let mut scored: Vec<(i64, f32)> = Vec::new();
847 for row in cand_rows {
848 let (face_id, bytes) = row?;
849 if let Some(emb) = FaceEmbedding::from_bytes(&bytes) {
850 let sim = centroid_emb.cosine_similarity(&emb);
851 scored.push((face_id, sim));
852 }
853 }
854
855 scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
857 scored.truncate(k);
858
859 Ok(scored)
860 }
861}
862
863#[cfg(test)]
864mod tests {
865 use rusqlite::{params, Connection};
866
867 use super::*;
868 use crate::db::create_schema;
869 use ndarray::Array1;
870
871 fn embedding(seed: f32) -> Vec<u8> {
872 let mut values = vec![0.0; 512];
873 values[0] = seed;
874 values[1] = 1.0 - seed;
875 FaceEmbedding::new(Array1::from_vec(values)).to_bytes()
876 }
877
878 fn seeded_conn() -> Connection {
879 let conn = Connection::open_in_memory().unwrap();
880 create_schema(&conn).unwrap();
881 conn.execute(
882 "INSERT INTO face_clusters (id, name, face_count, photo_count) VALUES (10, 'Person', 3, 3)",
883 [],
884 )
885 .unwrap();
886
887 for id in 1..=3 {
888 conn.execute(
889 "INSERT INTO photos (id, file_path, file_name, file_hash, file_size)
890 VALUES (?1, ?2, ?3, ?4, 100)",
891 params![
892 id,
893 format!("photos/{id}.jpg"),
894 format!("{id}.jpg"),
895 format!("hash-{id}")
896 ],
897 )
898 .unwrap();
899 conn.execute(
900 "INSERT INTO faces (
901 id, photo_id, bbox_x, bbox_y, bbox_width, bbox_height,
902 embedding, cluster_id, confidence, user_confirmed
903 )
904 VALUES (?1, ?2, 0.1, 0.1, 0.2, 0.2, ?3, 10, 0.99, ?4)",
905 params![
906 id,
907 id,
908 embedding(id as f32 / 10.0),
909 if id == 3 { 1 } else { 0 }
910 ],
911 )
912 .unwrap();
913 }
914 conn
915 }
916
917 fn insert_cluster(conn: &Connection, id: i64, name: &str) {
918 conn.execute(
919 "INSERT INTO face_clusters (id, name, face_count, photo_count) VALUES (?1, ?2, 0, 0)",
920 params![id, name],
921 )
922 .unwrap();
923 }
924
925 fn insert_face(conn: &Connection, id: i64, cluster_id: i64, confirmed: i32) {
926 conn.execute(
927 "INSERT INTO photos (id, file_path, file_name, file_hash, file_size)
928 VALUES (?1, ?2, ?3, ?4, 100)",
929 params![
930 id + 100,
931 format!("photos/review-{id}.jpg"),
932 format!("review-{id}.jpg"),
933 format!("review-hash-{id}")
934 ],
935 )
936 .unwrap();
937 conn.execute(
938 "INSERT INTO faces (
939 id, photo_id, bbox_x, bbox_y, bbox_width, bbox_height,
940 embedding, cluster_id, confidence, user_confirmed
941 )
942 VALUES (?1, ?2, 0.1, 0.1, 0.2, 0.2, ?3, ?4, 0.99, ?5)",
943 params![
944 id,
945 id + 100,
946 embedding(id as f32 / 10.0),
947 cluster_id,
948 confirmed
949 ],
950 )
951 .unwrap();
952 }
953
954 #[test]
955 fn get_faces_by_cluster_first_page_uses_valid_limit_parameter() {
956 let conn = seeded_conn();
957 let repo = FaceRepo::new(&conn);
958
959 let faces = repo
960 .get_faces_by_cluster(10, FaceStatus::Unconfirmed, None, 10)
961 .unwrap();
962
963 assert_eq!(
964 faces.iter().map(|f| f.face_id).collect::<Vec<_>>(),
965 vec![1, 2]
966 );
967 }
968
969 #[test]
970 fn get_faces_by_cluster_cursor_uses_valid_limit_parameter() {
971 let conn = seeded_conn();
972 let repo = FaceRepo::new(&conn);
973
974 let faces = repo
975 .get_faces_by_cluster(10, FaceStatus::All, Some(1), 10)
976 .unwrap();
977
978 assert_eq!(
979 faces.iter().map(|f| f.face_id).collect::<Vec<_>>(),
980 vec![2, 3]
981 );
982 }
983
984 #[test]
985 fn next_unconfirmed_face_batch_uses_the_cluster_with_most_pending_faces() {
986 let conn = Connection::open_in_memory().unwrap();
987 create_schema(&conn).unwrap();
988 insert_cluster(&conn, 10, "Small");
989 insert_cluster(&conn, 20, "Large");
990 insert_face(&conn, 1, 10, 0);
991 insert_face(&conn, 2, 20, 0);
992 insert_face(&conn, 3, 20, 0);
993 insert_face(&conn, 4, 20, 1);
994 let repo = FaceRepo::new(&conn);
995
996 let faces = repo.next_unconfirmed_face_batch(10).unwrap();
997
998 assert_eq!(
999 faces.iter().map(|f| f.face_id).collect::<Vec<_>>(),
1000 vec![2, 3]
1001 );
1002 assert!(faces.iter().all(|f| f.cluster_id == Some(20)));
1003 }
1004
1005 #[test]
1006 fn next_unconfirmed_face_batch_excluding_skipped_faces_moves_to_next_cluster() {
1007 let conn = Connection::open_in_memory().unwrap();
1008 create_schema(&conn).unwrap();
1009 insert_cluster(&conn, 10, "Small");
1010 insert_cluster(&conn, 20, "Large");
1011 insert_face(&conn, 1, 10, 0);
1012 insert_face(&conn, 2, 20, 0);
1013 insert_face(&conn, 3, 20, 0);
1014 let repo = FaceRepo::new(&conn);
1015
1016 let faces = repo
1017 .next_unconfirmed_face_batch_excluding(10, &[2, 3])
1018 .unwrap();
1019
1020 assert_eq!(faces.iter().map(|f| f.face_id).collect::<Vec<_>>(), vec![1]);
1021 assert!(faces.iter().all(|f| f.cluster_id == Some(10)));
1022 }
1023
1024 #[test]
1025 fn unclustered_embedding_candidates_skip_trashed_and_hidden_faces() {
1026 let conn = Connection::open_in_memory().unwrap();
1027 create_schema(&conn).unwrap();
1028 for id in 1..=3 {
1029 conn.execute(
1030 "INSERT INTO photos (id, file_path, file_name, file_hash, file_size, is_trashed)
1031 VALUES (?1, ?2, ?3, ?4, 100, ?5)",
1032 params![
1033 id,
1034 format!("photos/{id}.jpg"),
1035 format!("{id}.jpg"),
1036 format!("hash-{id}"),
1037 if id == 2 { 1 } else { 0 }
1038 ],
1039 )
1040 .unwrap();
1041 }
1042 conn.execute(
1043 "INSERT INTO faces (
1044 id, photo_id, bbox_x, bbox_y, bbox_width, bbox_height,
1045 embedding, cluster_id, confidence, user_confirmed
1046 )
1047 VALUES
1048 (1, 1, 0.1, 0.1, 0.2, 0.2, ?1, NULL, 0.9, 0),
1049 (2, 2, 0.1, 0.1, 0.2, 0.2, ?2, NULL, 0.9, 0),
1050 (3, 3, 0.1, 0.1, 0.2, 0.2, ?3, NULL, 0.9, -1)",
1051 params![embedding(0.1), embedding(0.2), embedding(0.3)],
1052 )
1053 .unwrap();
1054
1055 let faces = FaceRepo::new(&conn)
1056 .get_unclustered_faces_with_photo_embeddings()
1057 .unwrap();
1058
1059 assert_eq!(
1060 faces.iter().map(|(id, _, _)| *id).collect::<Vec<_>>(),
1061 vec![1]
1062 );
1063 }
1064
1065 #[test]
1066 fn k_similar_skips_trashed_candidates_and_zero_k_short_circuits() {
1067 let conn = Connection::open_in_memory().unwrap();
1068 create_schema(&conn).unwrap();
1069 insert_cluster(&conn, 10, "Person");
1070 for id in 1..=3 {
1071 conn.execute(
1072 "INSERT INTO photos (id, file_path, file_name, file_hash, file_size, is_trashed)
1073 VALUES (?1, ?2, ?3, ?4, 100, ?5)",
1074 params![
1075 id,
1076 format!("photos/{id}.jpg"),
1077 format!("{id}.jpg"),
1078 format!("hash-{id}"),
1079 if id == 3 { 1 } else { 0 }
1080 ],
1081 )
1082 .unwrap();
1083 }
1084 conn.execute(
1085 "INSERT INTO faces (
1086 id, photo_id, bbox_x, bbox_y, bbox_width, bbox_height,
1087 embedding, cluster_id, confidence, user_confirmed
1088 )
1089 VALUES
1090 (1, 1, 0.1, 0.1, 0.2, 0.2, ?1, 10, 0.9, 1),
1091 (2, 2, 0.1, 0.1, 0.2, 0.2, ?2, NULL, 0.9, 0),
1092 (3, 3, 0.1, 0.1, 0.2, 0.2, ?3, NULL, 0.9, 0)",
1093 params![embedding(0.9), embedding(0.8), embedding(0.9)],
1094 )
1095 .unwrap();
1096 conn.execute(
1097 "INSERT INTO person_gallery_embeddings (cluster_id, face_id, embedding, quality_score, source)
1098 VALUES (10, 1, ?1, 0.9, 'user_confirmed')",
1099 params![embedding(0.9)],
1100 )
1101 .unwrap();
1102 let repo = FaceRepo::new(&conn);
1103
1104 assert!(repo.k_similar_to_cluster(10, 0).unwrap().is_empty());
1105 let similar = repo.k_similar_to_cluster(10, 10).unwrap();
1106
1107 assert_eq!(
1108 similar.iter().map(|(id, _)| *id).collect::<Vec<_>>(),
1109 vec![2]
1110 );
1111 }
1112}