1use super::retrieval::RetrievalHit;
20
21#[derive(Debug, Clone, Copy)]
25pub struct ResolverWeights {
26 pub embedding: f32,
28 pub cooccurrence: f32,
30 pub temporal: f32,
32}
33
34impl Default for ResolverWeights {
35 fn default() -> Self {
36 Self {
37 embedding: 1.0,
38 cooccurrence: 0.30,
39 temporal: 0.50,
40 }
41 }
42}
43
44#[derive(Debug, Clone)]
46pub struct ResolverContext {
47 pub photo_other_clusters: Vec<i64>,
49 pub cooccurrence_scores: std::collections::HashMap<i64, i64>,
52 pub temporal_neighbor_clusters: std::collections::HashSet<i64>,
54}
55
56pub fn rerank(
61 hits: &[RetrievalHit],
62 ctx: &ResolverContext,
63 weights: ResolverWeights,
64) -> Vec<RetrievalHit> {
65 if hits.is_empty() {
66 return Vec::new();
67 }
68
69 let cooccurrence_bonus = |count: i64| -> f32 {
72 if count <= 0 {
73 0.0
74 } else {
75 ((count as f32 + 1.0).ln() / 5.0).min(1.0)
76 }
77 };
78
79 let mut adjusted: Vec<RetrievalHit> = hits
80 .iter()
81 .map(|h| {
82 let mut score = h.score * weights.embedding;
83
84 if weights.cooccurrence > 0.0 && !ctx.photo_other_clusters.is_empty() {
85 let raw = ctx
86 .cooccurrence_scores
87 .get(&h.cluster_id)
88 .copied()
89 .unwrap_or(0);
90 score += weights.cooccurrence * cooccurrence_bonus(raw);
91 }
92
93 if weights.temporal > 0.0 && ctx.temporal_neighbor_clusters.contains(&h.cluster_id) {
94 score += weights.temporal * 0.1;
95 }
96
97 RetrievalHit {
98 cluster_id: h.cluster_id,
99 score,
100 best_match_face_id: h.best_match_face_id,
101 num_matches: h.num_matches,
102 }
103 })
104 .collect();
105
106 adjusted.sort_by(|a, b| {
107 b.score
108 .partial_cmp(&a.score)
109 .unwrap_or(std::cmp::Ordering::Equal)
110 });
111 adjusted
112}
113
114#[cfg(test)]
115mod tests {
116 use super::*;
117 use std::collections::{HashMap, HashSet};
118
119 fn hit(cluster_id: i64, score: f32) -> RetrievalHit {
120 RetrievalHit {
121 cluster_id,
122 score,
123 best_match_face_id: 0,
124 num_matches: 1,
125 }
126 }
127
128 #[test]
129 fn no_context_preserves_order() {
130 let hits = vec![hit(1, 0.60), hit(2, 0.55)];
131 let ctx = ResolverContext {
132 photo_other_clusters: Vec::new(),
133 cooccurrence_scores: HashMap::new(),
134 temporal_neighbor_clusters: HashSet::new(),
135 };
136 let out = rerank(&hits, &ctx, ResolverWeights::default());
137 assert_eq!(out[0].cluster_id, 1);
138 assert_eq!(out[1].cluster_id, 2);
139 }
140
141 #[test]
142 fn cooccurrence_can_flip_close_pair() {
143 let hits = vec![hit(1, 0.52), hit(2, 0.50)];
146 let mut co = HashMap::new();
147 co.insert(2i64, 50i64); let ctx = ResolverContext {
149 photo_other_clusters: vec![99],
150 cooccurrence_scores: co,
151 temporal_neighbor_clusters: HashSet::new(),
152 };
153 let out = rerank(&hits, &ctx, ResolverWeights::default());
154 assert_eq!(out[0].cluster_id, 2);
155 }
156
157 #[test]
158 fn temporal_boost_applies() {
159 let hits = vec![hit(1, 0.50), hit(2, 0.49)];
160 let mut neighbors = HashSet::new();
161 neighbors.insert(2i64);
162 let ctx = ResolverContext {
163 photo_other_clusters: Vec::new(),
164 cooccurrence_scores: HashMap::new(),
165 temporal_neighbor_clusters: neighbors,
166 };
167 let out = rerank(&hits, &ctx, ResolverWeights::default());
168 assert_eq!(out[0].cluster_id, 2);
169 }
170}