opendev_memory/
selector.rs1use std::collections::HashMap;
6
7use crate::embeddings::{EmbeddingCache, cosine_similarity};
8use crate::playbook::Bullet;
9
10#[derive(Debug, Clone)]
12pub struct ScoredBullet {
13 pub bullet: Bullet,
14 pub score: f64,
15 pub score_breakdown: HashMap<String, f64>,
16}
17
18pub struct BulletSelector {
25 pub weights: HashMap<String, f64>,
26 pub embedding_model: String,
27 pub cache_file: Option<String>,
28 pub embedding_cache: EmbeddingCache,
29}
30
31impl BulletSelector {
32 pub fn new(
34 weights: Option<HashMap<String, f64>>,
35 embedding_model: &str,
36 cache_file: Option<&str>,
37 ) -> Self {
38 let weights = weights.unwrap_or_else(|| {
39 let mut w = HashMap::new();
40 w.insert("effectiveness".to_string(), 0.6);
41 w.insert("recency".to_string(), 0.4);
42 w.insert("semantic".to_string(), 0.0);
43 w
44 });
45
46 let embedding_cache = cache_file
47 .and_then(|p| EmbeddingCache::load_from_file(std::path::Path::new(p)))
48 .unwrap_or_else(|| EmbeddingCache::new(embedding_model));
49
50 Self {
51 weights,
52 embedding_model: embedding_model.to_string(),
53 cache_file: cache_file.map(String::from),
54 embedding_cache,
55 }
56 }
57
58 pub fn select(&self, bullets: &[Bullet], max_count: usize, query: Option<&str>) -> Vec<Bullet> {
60 if bullets.len() <= max_count {
61 return bullets.to_vec();
62 }
63
64 let mut scored: Vec<ScoredBullet> = bullets
65 .iter()
66 .map(|b| self.score_bullet(b, query))
67 .collect();
68
69 scored.sort_by(|a, b| {
70 b.score
71 .partial_cmp(&a.score)
72 .unwrap_or(std::cmp::Ordering::Equal)
73 });
74 scored
75 .into_iter()
76 .take(max_count)
77 .map(|sb| sb.bullet)
78 .collect()
79 }
80
81 pub fn score_bullet(&self, bullet: &Bullet, query: Option<&str>) -> ScoredBullet {
83 let mut breakdown = HashMap::new();
84
85 let effectiveness = self.effectiveness_score(bullet);
86 breakdown.insert("effectiveness".to_string(), effectiveness);
87
88 let recency = self.recency_score(bullet);
89 breakdown.insert("recency".to_string(), recency);
90
91 let semantic = match query {
92 Some(q) if self.weights.get("semantic").copied().unwrap_or(0.0) > 0.0 => {
93 self.semantic_score(q, bullet)
94 }
95 _ => 0.0,
96 };
97 breakdown.insert("semantic".to_string(), semantic);
98
99 let final_score = self.weights.get("effectiveness").unwrap_or(&0.6) * effectiveness
100 + self.weights.get("recency").unwrap_or(&0.4) * recency
101 + self.weights.get("semantic").unwrap_or(&0.0) * semantic;
102
103 ScoredBullet {
104 bullet: bullet.clone(),
105 score: final_score,
106 score_breakdown: breakdown,
107 }
108 }
109
110 fn effectiveness_score(&self, bullet: &Bullet) -> f64 {
113 let total = bullet.helpful + bullet.harmful + bullet.neutral;
114 if total == 0 {
115 return 0.5;
116 }
117 let weighted =
118 bullet.helpful as f64 * 1.0 + bullet.neutral as f64 * 0.5 + bullet.harmful as f64 * 0.0;
119 weighted / total as f64
120 }
121
122 fn recency_score(&self, bullet: &Bullet) -> f64 {
125 let updated_at = bullet
126 .updated_at
127 .replace("Z", "+00:00")
128 .parse::<chrono::DateTime<chrono::Utc>>()
129 .or_else(|_| {
130 chrono::DateTime::parse_from_rfc3339(&bullet.updated_at)
131 .map(|dt| dt.with_timezone(&chrono::Utc))
132 });
133
134 match updated_at {
135 Ok(dt) => {
136 let days_old = (chrono::Utc::now() - dt).num_days().max(0) as f64;
137 let decay_rate = 0.1;
138 1.0 / (1.0 + days_old * decay_rate)
139 }
140 Err(_) => 0.5,
141 }
142 }
143
144 fn semantic_score(&self, query: &str, bullet: &Bullet) -> f64 {
146 if self.weights.get("semantic").copied().unwrap_or(0.0) <= 0.0 {
147 return 0.0;
148 }
149
150 let query_emb = self.embedding_cache.peek(query, None);
151 let bullet_emb = self.embedding_cache.peek(&bullet.content, None);
152
153 match (query_emb, bullet_emb) {
154 (Some(q), Some(b)) => {
155 let sim = cosine_similarity(q, b);
156 (sim + 1.0) / 2.0 }
158 _ => 0.5,
159 }
160 }
161
162 pub fn selection_stats(
164 &self,
165 all_bullets: &[Bullet],
166 selected: &[Bullet],
167 ) -> HashMap<String, f64> {
168 let all_scored: Vec<ScoredBullet> = all_bullets
169 .iter()
170 .map(|b| self.score_bullet(b, None))
171 .collect();
172
173 let selected_ids: std::collections::HashSet<&str> =
174 selected.iter().map(|b| b.id.as_str()).collect();
175
176 let avg_all = if all_scored.is_empty() {
177 0.0
178 } else {
179 all_scored.iter().map(|s| s.score).sum::<f64>() / all_scored.len() as f64
180 };
181
182 let selected_scores: Vec<f64> = all_scored
183 .iter()
184 .filter(|s| selected_ids.contains(s.bullet.id.as_str()))
185 .map(|s| s.score)
186 .collect();
187
188 let avg_selected = if selected_scores.is_empty() {
189 0.0
190 } else {
191 selected_scores.iter().sum::<f64>() / selected_scores.len() as f64
192 };
193
194 let mut stats = HashMap::new();
195 stats.insert("total_bullets".to_string(), all_bullets.len() as f64);
196 stats.insert("selected_bullets".to_string(), selected.len() as f64);
197 stats.insert(
198 "selection_rate".to_string(),
199 if all_bullets.is_empty() {
200 0.0
201 } else {
202 selected.len() as f64 / all_bullets.len() as f64
203 },
204 );
205 stats.insert("avg_all_score".to_string(), avg_all);
206 stats.insert("avg_selected_score".to_string(), avg_selected);
207 stats.insert("score_improvement".to_string(), avg_selected - avg_all);
208 stats
209 }
210}
211
212impl Default for BulletSelector {
213 fn default() -> Self {
214 Self::new(None, "text-embedding-3-small", None)
215 }
216}
217
218#[cfg(test)]
219#[path = "selector_tests.rs"]
220mod tests;