1use std::collections::HashMap;
9use std::sync::{Arc, Mutex};
10
11#[derive(Debug, Clone)]
13pub struct Feedback {
14 pub request_id: String,
16 pub response_id: String,
18 pub rating: f64,
20 pub latency_ms: u64,
22 pub tokens_used: u64,
24 pub timestamp: u64,
26}
27
28#[derive(Debug, Clone)]
32pub struct RewardSignal {
33 pub quality_score: f64,
35 pub cost_penalty: f64,
37 pub latency_penalty: f64,
39 pub reward: f64,
41}
42
43impl RewardSignal {
44 pub fn from_feedback(feedback: &Feedback, max_tokens: u64, max_latency_ms: u64) -> Self {
49 let quality_score = feedback.rating.clamp(0.0, 1.0);
50 let cost_penalty = if max_tokens == 0 {
51 0.0
52 } else {
53 (feedback.tokens_used as f64 / max_tokens as f64).clamp(0.0, 1.0)
54 };
55 let latency_penalty = if max_latency_ms == 0 {
56 0.0
57 } else {
58 (feedback.latency_ms as f64 / max_latency_ms as f64).clamp(0.0, 1.0)
59 };
60 let reward = quality_score * 0.5 - cost_penalty * 0.3 - latency_penalty * 0.2;
61 Self {
62 quality_score,
63 cost_penalty,
64 latency_penalty,
65 reward,
66 }
67 }
68}
69
70#[derive(Debug, Clone)]
72pub struct FeedbackStore {
73 inner: Arc<Mutex<FeedbackStoreInner>>,
74}
75
76#[derive(Debug)]
77struct FeedbackStoreInner {
78 capacity: usize,
79 data: HashMap<String, Vec<Feedback>>,
81 heads: HashMap<String, usize>,
83}
84
85impl FeedbackStore {
86 pub fn new(capacity: usize) -> Self {
88 Self {
89 inner: Arc::new(Mutex::new(FeedbackStoreInner {
90 capacity,
91 data: HashMap::new(),
92 heads: HashMap::new(),
93 })),
94 }
95 }
96
97 pub fn push(&self, variant: &str, feedback: Feedback) {
99 let mut guard = self.inner.lock().unwrap_or_else(|e| e.into_inner());
100 let cap = guard.capacity;
101 let key = variant.to_string();
102 let guard = &mut *guard;
104 let head = guard.heads.entry(key.clone()).or_insert(0);
105 let buf = guard.data.entry(key).or_default();
106 if buf.len() < cap {
107 buf.push(feedback);
108 } else {
109 buf[*head] = feedback;
110 *head = (*head + 1) % cap;
111 }
112 }
113
114 pub fn get(&self, variant: &str) -> Vec<Feedback> {
116 let guard = self.inner.lock().unwrap_or_else(|e| e.into_inner());
117 guard.data.get(variant).cloned().unwrap_or_default()
118 }
119
120 pub fn variants(&self) -> Vec<String> {
122 let guard = self.inner.lock().unwrap_or_else(|e| e.into_inner());
123 guard.data.keys().cloned().collect()
124 }
125}
126
127#[derive(Debug, Clone)]
132pub struct RewardModel {
133 inner: Arc<Mutex<RewardModelInner>>,
134}
135
136#[derive(Debug)]
137struct RewardModelInner {
138 decay: f64,
140 averages: HashMap<(String, String), (f64, Vec<f64>)>,
142}
143
144impl RewardModel {
145 pub fn new(decay: f64) -> Self {
150 let decay = decay.clamp(1e-6, 1.0 - 1e-6);
151 Self {
152 inner: Arc::new(Mutex::new(RewardModelInner {
153 decay,
154 averages: HashMap::new(),
155 })),
156 }
157 }
158
159 pub fn update(&self, model: &str, template: &str, reward: f64) {
161 let mut guard = self.inner.lock().unwrap_or_else(|e| e.into_inner());
162 let decay = guard.decay;
163 let key = (model.to_string(), template.to_string());
164 let entry = guard.averages.entry(key).or_insert((reward, Vec::new()));
165 entry.0 = decay * entry.0 + (1.0 - decay) * reward;
166 entry.1.push(reward);
167 }
168
169 pub fn average(&self, model: &str, template: &str) -> Option<f64> {
171 let guard = self.inner.lock().unwrap_or_else(|e| e.into_inner());
172 guard
173 .averages
174 .get(&(model.to_string(), template.to_string()))
175 .map(|(avg, _)| *avg)
176 }
177
178 pub fn history(&self, model: &str, template: &str) -> Vec<f64> {
180 let guard = self.inner.lock().unwrap_or_else(|e| e.into_inner());
181 guard
182 .averages
183 .get(&(model.to_string(), template.to_string()))
184 .map(|(_, h)| h.clone())
185 .unwrap_or_default()
186 }
187}
188
189#[derive(Debug, Clone)]
191pub struct FeedbackLoop {
192 store: FeedbackStore,
193 reward_model: RewardModel,
194 max_tokens: u64,
196 max_latency_ms: u64,
198}
199
200impl FeedbackLoop {
201 pub fn new(ring_capacity: usize, decay: f64, max_tokens: u64, max_latency_ms: u64) -> Self {
208 Self {
209 store: FeedbackStore::new(ring_capacity),
210 reward_model: RewardModel::new(decay),
211 max_tokens,
212 max_latency_ms,
213 }
214 }
215
216 pub fn record(&self, feedback: Feedback) {
223 let variant = feedback
224 .request_id
225 .split(':')
226 .next()
227 .unwrap_or(&feedback.request_id)
228 .to_string();
229 let signal = RewardSignal::from_feedback(&feedback, self.max_tokens, self.max_latency_ms);
230 self.store.push(&variant, feedback);
231 self.reward_model.update(&variant, "default", signal.reward);
232 }
233
234 pub fn best_variant(&self, model_options: &[String]) -> Option<String> {
238 model_options
239 .iter()
240 .filter_map(|m| {
241 self.reward_model
242 .average(m, "default")
243 .map(|avg| (m.clone(), avg))
244 })
245 .max_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
246 .map(|(m, _)| m)
247 }
248
249 pub fn reward_history(&self, variant: &str) -> Vec<f64> {
251 self.reward_model.history(variant, "default")
252 }
253
254 pub fn convergence_score(&self) -> f64 {
259 let variants = self.store.variants();
260 if variants.is_empty() {
261 return 0.0;
262 }
263 let averages: Vec<f64> = variants
264 .iter()
265 .filter_map(|v| self.reward_model.average(v, "default"))
266 .collect();
267 if averages.len() < 2 {
268 return 0.0;
269 }
270 let mean = averages.iter().sum::<f64>() / averages.len() as f64;
271 let variance = averages.iter().map(|x| (x - mean).powi(2)).sum::<f64>()
272 / averages.len() as f64;
273 variance
274 }
275}
276
277#[cfg(test)]
278mod tests {
279 use super::*;
280
281 fn make_feedback(request_id: &str, rating: f64, latency_ms: u64, tokens_used: u64) -> Feedback {
282 Feedback {
283 request_id: request_id.to_string(),
284 response_id: format!("resp-{}", request_id),
285 rating,
286 latency_ms,
287 tokens_used,
288 timestamp: 1_000_000,
289 }
290 }
291
292 #[test]
293 fn reward_signal_formula() {
294 let fb = make_feedback("m", 1.0, 500, 1000);
295 let sig = RewardSignal::from_feedback(&fb, 2000, 2000);
296 let expected = 1.0 * 0.5 - 0.5 * 0.3 - 0.25 * 0.2;
298 assert!((sig.reward - expected).abs() < 1e-9);
299 }
300
301 #[test]
302 fn reward_signal_zero_max() {
303 let fb = make_feedback("m", 0.8, 100, 100);
304 let sig = RewardSignal::from_feedback(&fb, 0, 0);
305 let expected = 0.8 * 0.5;
307 assert!((sig.reward - expected).abs() < 1e-9);
308 }
309
310 #[test]
311 fn feedback_store_ring_buffer() {
312 let store = FeedbackStore::new(3);
313 for i in 0u64..5 {
314 store.push("m", make_feedback("m", 0.5, i * 10, i * 100));
315 }
316 assert_eq!(store.get("m").len(), 3);
318 }
319
320 #[test]
321 fn reward_model_weighted_average() {
322 let model = RewardModel::new(0.9);
323 model.update("gpt4", "default", 1.0);
324 model.update("gpt4", "default", 0.0);
325 let avg = model.average("gpt4", "default").expect("should have avg");
326 assert!((avg - 0.9).abs() < 1e-9);
328 }
329
330 #[test]
331 fn reward_model_history_length() {
332 let model = RewardModel::new(0.8);
333 model.update("claude", "default", 0.5);
334 model.update("claude", "default", 0.7);
335 model.update("claude", "default", 0.6);
336 assert_eq!(model.history("claude", "default").len(), 3);
337 }
338
339 #[test]
340 fn feedback_loop_best_variant() {
341 let fl = FeedbackLoop::new(10, 0.5, 10_000, 5_000);
342 for _ in 0..5 {
344 fl.record(make_feedback("good_model:req1", 1.0, 100, 100));
345 }
346 for _ in 0..5 {
348 fl.record(make_feedback("bad_model:req1", 0.1, 4000, 9000));
349 }
350 let options = vec!["good_model".to_string(), "bad_model".to_string()];
351 let best = fl.best_variant(&options).expect("should find best");
352 assert_eq!(best, "good_model");
353 }
354
355 #[test]
356 fn feedback_loop_reward_history() {
357 let fl = FeedbackLoop::new(10, 0.5, 10_000, 5_000);
358 fl.record(make_feedback("modelA:r1", 0.8, 200, 500));
359 fl.record(make_feedback("modelA:r2", 0.9, 300, 600));
360 let h = fl.reward_history("modelA");
361 assert_eq!(h.len(), 2);
362 }
363
364 #[test]
365 fn feedback_loop_convergence_score_single_variant() {
366 let fl = FeedbackLoop::new(10, 0.5, 10_000, 5_000);
367 fl.record(make_feedback("only:req", 0.5, 500, 500));
368 assert_eq!(fl.convergence_score(), 0.0);
370 }
371
372 #[test]
373 fn feedback_loop_convergence_score_two_variants() {
374 let fl = FeedbackLoop::new(10, 0.5, 10_000, 5_000);
375 fl.record(make_feedback("alpha:req", 1.0, 100, 100));
376 fl.record(make_feedback("beta:req", 0.0, 100, 100));
377 assert!(fl.convergence_score() > 0.0);
379 }
380
381 #[test]
382 fn best_variant_returns_none_when_no_data() {
383 let fl = FeedbackLoop::new(10, 0.5, 10_000, 5_000);
384 let options = vec!["x".to_string(), "y".to_string()];
385 assert!(fl.best_variant(&options).is_none());
386 }
387}