Skip to main content

opendev_memory/
selector.rs

1//! Bullet selection logic for ACE playbook context optimization.
2//!
3//! Mirrors `opendev/core/context_engineering/memory/selector.py`.
4
5use std::collections::HashMap;
6
7use crate::embeddings::{EmbeddingCache, cosine_similarity};
8use crate::playbook::Bullet;
9
10/// Bullet with its calculated relevance score.
11#[derive(Debug, Clone)]
12pub struct ScoredBullet {
13    pub bullet: Bullet,
14    pub score: f64,
15    pub score_breakdown: HashMap<String, f64>,
16}
17
18/// Selects most relevant bullets for a given query.
19///
20/// Implements hybrid retrieval with three scoring factors:
21/// - Effectiveness: Based on helpful/harmful feedback
22/// - Recency: Prefers recently updated bullets
23/// - Semantic: Query-to-bullet similarity using embeddings
24pub 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    /// Create a new bullet selector.
33    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    /// Select top-K most relevant bullets.
59    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    /// Score a single bullet.
82    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    /// Effectiveness score based on helpful/harmful feedback.
111    /// Returns 0.0..1.0. Untested bullets get 0.5.
112    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    /// Recency score -- prefer recently updated bullets.
123    /// Returns 0.0..1.0 using exponential decay.
124    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    /// Semantic similarity score using cached embeddings.
145    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 // Normalize from [-1, 1] to [0, 1]
157            }
158            _ => 0.5,
159        }
160    }
161
162    /// Get statistics about a selection.
163    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;