lean_ctx/core/context_kernel/
feedback.rs1use 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}