Skip to main content

lean_ctx/core/context_kernel/
feedback.rs

1//! Outcome feedback collection for Context Kernel learning.
2
3use std::collections::HashMap;
4use std::fs;
5use std::io::Write;
6use std::path::PathBuf;
7use std::time::{SystemTime, UNIX_EPOCH};
8
9use super::types::{ContextReceiptV1, ReceiptOutcome};
10
11const LEARNING_RATE: f64 = 0.1;
12const DEFAULT_PROVIDER_WEIGHT: f64 = 1.0;
13
14#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
15pub struct FeedbackEntry {
16    pub plan_id: String,
17    pub outcome: String,
18    pub provider_scores: HashMap<String, f64>,
19    pub timestamp_epoch: u64,
20}
21
22pub struct FeedbackCollector {
23    log_path: PathBuf,
24    provider_weights: HashMap<String, f64>,
25}
26
27impl FeedbackCollector {
28    pub fn new(log_path: PathBuf) -> Self {
29        Self {
30            log_path,
31            provider_weights: HashMap::new(),
32        }
33    }
34
35    pub fn default_for_project(_project_root: &str) -> Self {
36        let cache_root = std::env::var_os("HOME")
37            .map_or_else(|| PathBuf::from("."), PathBuf::from)
38            .join(".cache")
39            .join("lean-ctx")
40            .join("kernel");
41        Self::new(cache_root.join("feedback.jsonl"))
42    }
43
44    pub fn record_outcome(&mut self, receipt: &ContextReceiptV1) {
45        let score = outcome_score(&receipt.outcome);
46        let provider_scores: HashMap<String, f64> = receipt
47            .feedback_attribution
48            .keys()
49            .map(|provider| (provider.clone(), score))
50            .collect();
51
52        for (provider, provider_score) in &provider_scores {
53            update_weight(&mut self.provider_weights, provider, *provider_score);
54        }
55
56        let entry = FeedbackEntry {
57            plan_id: receipt.plan_id.clone(),
58            outcome: outcome_name(&receipt.outcome).to_owned(),
59            provider_scores,
60            timestamp_epoch: SystemTime::now()
61                .duration_since(UNIX_EPOCH)
62                .map_or(0, |duration| duration.as_secs()),
63        };
64
65        let Some(parent) = self.log_path.parent() else {
66            return;
67        };
68        if fs::create_dir_all(parent).is_err() {
69            return;
70        }
71        let Ok(serialized) = serde_json::to_string(&entry) else {
72            return;
73        };
74        if let Ok(mut file) = fs::OpenOptions::new()
75            .create(true)
76            .append(true)
77            .open(&self.log_path)
78        {
79            let _ = writeln!(file, "{serialized}");
80        }
81    }
82
83    pub fn provider_weight(&self, provider: &str) -> f64 {
84        self.provider_weights
85            .get(provider)
86            .copied()
87            .unwrap_or(DEFAULT_PROVIDER_WEIGHT)
88    }
89
90    pub fn load_weights(&mut self) {
91        self.provider_weights.clear();
92        let Ok(contents) = fs::read_to_string(&self.log_path) else {
93            return;
94        };
95
96        for line in contents.lines() {
97            if let Ok(entry) = serde_json::from_str::<FeedbackEntry>(line) {
98                for (provider, score) in entry.provider_scores {
99                    update_weight(&mut self.provider_weights, &provider, score);
100                }
101            }
102        }
103    }
104}
105
106pub fn record_kernel_feedback(project_root: &str, receipt: &ContextReceiptV1) {
107    let mut collector = FeedbackCollector::default_for_project(project_root);
108    collector.load_weights();
109    collector.record_outcome(receipt);
110}
111
112fn update_weight(weights: &mut HashMap<String, f64>, provider: &str, score: f64) {
113    let current = weights
114        .get(provider)
115        .copied()
116        .unwrap_or(DEFAULT_PROVIDER_WEIGHT);
117    weights.insert(
118        provider.to_owned(),
119        current * (1.0 - LEARNING_RATE) + score * LEARNING_RATE,
120    );
121}
122
123fn outcome_score(outcome: &ReceiptOutcome) -> f64 {
124    match outcome {
125        ReceiptOutcome::Accepted => 1.0,
126        ReceiptOutcome::Partial => 0.5,
127        ReceiptOutcome::Unknown => 0.2,
128        ReceiptOutcome::Rejected => 0.0,
129    }
130}
131
132fn outcome_name(outcome: &ReceiptOutcome) -> &'static str {
133    match outcome {
134        ReceiptOutcome::Accepted => "accepted",
135        ReceiptOutcome::Partial => "partial",
136        ReceiptOutcome::Unknown => "unknown",
137        ReceiptOutcome::Rejected => "rejected",
138    }
139}
140
141#[cfg(test)]
142mod tests {
143    use std::collections::HashMap;
144    use std::fs;
145    use std::path::PathBuf;
146    use std::sync::atomic::{AtomicUsize, Ordering};
147
148    use super::FeedbackCollector;
149    use crate::core::context_kernel::types::{ContextReceiptV1, ReceiptOutcome};
150
151    static NEXT_TEST_PATH: AtomicUsize = AtomicUsize::new(0);
152
153    fn test_log_path(test_name: &str) -> PathBuf {
154        let sequence = NEXT_TEST_PATH.fetch_add(1, Ordering::Relaxed);
155        std::env::temp_dir().join(format!(
156            "lean-ctx-feedback-{}-{test_name}-{sequence}.jsonl",
157            std::process::id()
158        ))
159    }
160
161    fn receipt(outcome: ReceiptOutcome) -> ContextReceiptV1 {
162        ContextReceiptV1 {
163            receipt_id: "receipt-1".to_owned(),
164            plan_id: "plan-1".to_owned(),
165            delivered_tokens: 100,
166            cache_hits: 0,
167            cache_misses: 0,
168            outcome,
169            quality_signals: Vec::new(),
170            feedback_attribution: HashMap::from([("files".to_owned(), 1.0)]),
171        }
172    }
173
174    #[test]
175    fn accepted_outcome_increases_provider_weight() {
176        let path = test_log_path("accepted");
177        let mut collector = FeedbackCollector::new(path.clone());
178        collector.provider_weights.insert("files".to_owned(), 0.5);
179
180        collector.record_outcome(&receipt(ReceiptOutcome::Accepted));
181
182        assert!(collector.provider_weight("files") > 0.5);
183        let _ = fs::remove_file(path);
184    }
185
186    #[test]
187    fn rejected_outcome_decreases_provider_weight() {
188        let path = test_log_path("rejected");
189        let mut collector = FeedbackCollector::new(path.clone());
190
191        collector.record_outcome(&receipt(ReceiptOutcome::Rejected));
192
193        assert!(collector.provider_weight("files") < 1.0);
194        let _ = fs::remove_file(path);
195    }
196
197    #[test]
198    fn feedback_persists_to_file() {
199        let path = test_log_path("persists");
200        let mut collector = FeedbackCollector::new(path.clone());
201        collector.record_outcome(&receipt(ReceiptOutcome::Rejected));
202
203        let mut restored = FeedbackCollector::new(path.clone());
204        restored.load_weights();
205
206        assert!((restored.provider_weight("files") - 0.9).abs() < f64::EPSILON);
207        assert!(fs::read_to_string(&path).is_ok());
208        let _ = fs::remove_file(path);
209    }
210}