Skip to main content

tokio_prompt_orchestrator/
feedback_loop.rs

1//! Reinforcement-learning-style feedback collector and reward modeler.
2//!
3//! This module records per-request feedback, computes composite reward signals,
4//! and maintains a running weighted average reward per (model, prompt_template)
5//! pair. The reward weight decays exponentially with time so that recent
6//! observations matter more than old ones.
7
8use std::collections::HashMap;
9use std::sync::{Arc, Mutex};
10
11/// A single feedback record attached to one request/response pair.
12#[derive(Debug, Clone)]
13pub struct Feedback {
14    /// Identifier of the originating request.
15    pub request_id: String,
16    /// Identifier of the response that was evaluated.
17    pub response_id: String,
18    /// Human or automated quality rating in [0, 1].
19    pub rating: f64,
20    /// End-to-end latency experienced by the caller, in milliseconds.
21    pub latency_ms: u64,
22    /// Total tokens consumed by the request/response pair.
23    pub tokens_used: u64,
24    /// Unix epoch seconds at which the feedback was recorded.
25    pub timestamp: u64,
26}
27
28/// Composite reward derived from a single [`Feedback`] record.
29///
30/// Formula: `quality_score * 0.5 - cost_penalty * 0.3 - latency_penalty * 0.2`
31#[derive(Debug, Clone)]
32pub struct RewardSignal {
33    /// Raw quality score in [0, 1] (equal to `feedback.rating`).
34    pub quality_score: f64,
35    /// Normalised cost penalty proportional to tokens used (range [0, 1]).
36    pub cost_penalty: f64,
37    /// Normalised latency penalty (range [0, 1]).
38    pub latency_penalty: f64,
39    /// Final composite reward value.
40    pub reward: f64,
41}
42
43impl RewardSignal {
44    /// Compute a [`RewardSignal`] from a raw [`Feedback`] record.
45    ///
46    /// `max_tokens` and `max_latency_ms` are used to normalise the penalties
47    /// into [0, 1].  If either is zero the corresponding penalty is 0.
48    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/// Thread-safe ring buffer of the last `capacity` feedbacks per model variant.
71#[derive(Debug, Clone)]
72pub struct FeedbackStore {
73    inner: Arc<Mutex<FeedbackStoreInner>>,
74}
75
76#[derive(Debug)]
77struct FeedbackStoreInner {
78    capacity: usize,
79    /// variant -> circular buffer of feedbacks
80    data: HashMap<String, Vec<Feedback>>,
81    /// variant -> write index (next position to overwrite)
82    heads: HashMap<String, usize>,
83}
84
85impl FeedbackStore {
86    /// Create a new store that retains the last `capacity` feedbacks per variant.
87    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    /// Append a feedback for `variant`.  Overwrites the oldest entry when full.
98    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        // Ensure both maps have an entry.
103        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    /// Return a snapshot of all feedbacks for `variant` in insertion order.
115    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    /// Return all known variant names.
121    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/// Running weighted-average reward per `(model, prompt_template)` key.
128///
129/// Each new reward is blended with the running average using an exponential
130/// decay factor: `avg = decay * avg + (1 - decay) * new_reward`.
131#[derive(Debug, Clone)]
132pub struct RewardModel {
133    inner: Arc<Mutex<RewardModelInner>>,
134}
135
136#[derive(Debug)]
137struct RewardModelInner {
138    /// Exponential smoothing factor in (0, 1).  Closer to 1 = slower decay.
139    decay: f64,
140    /// (model, template) -> (weighted_avg, history)
141    averages: HashMap<(String, String), (f64, Vec<f64>)>,
142}
143
144impl RewardModel {
145    /// Create a new model with the given exponential decay factor.
146    ///
147    /// `decay` must be in (0, 1).  A value of 0.9 means older observations
148    /// retain 90% of their weight after each update.
149    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    /// Update the running average for `(model, template)` with a new reward.
160    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    /// Return the current weighted-average reward for `(model, template)`.
170    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    /// Return the full chronological reward history for `(model, template)`.
179    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/// Closed-loop feedback collector that ties together the store and reward model.
190#[derive(Debug, Clone)]
191pub struct FeedbackLoop {
192    store: FeedbackStore,
193    reward_model: RewardModel,
194    /// Normalisation ceiling for token counts.
195    max_tokens: u64,
196    /// Normalisation ceiling for latency.
197    max_latency_ms: u64,
198}
199
200impl FeedbackLoop {
201    /// Create a new `FeedbackLoop`.
202    ///
203    /// - `ring_capacity`: number of feedbacks retained per variant in the ring buffer.
204    /// - `decay`: exponential smoothing factor for the reward model (0 < decay < 1).
205    /// - `max_tokens`: token ceiling used when normalising cost penalties.
206    /// - `max_latency_ms`: latency ceiling used when normalising latency penalties.
207    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    /// Record a feedback observation.
217    ///
218    /// The variant key is derived from `request_id` as the model name prefix
219    /// (everything before the first `':'`).  Callers should embed the model
220    /// name in `request_id` as `"model:rest-of-id"`.  If no `':'` is found the
221    /// whole `request_id` is used as the variant key.
222    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    /// Return the variant (from `model_options`) with the highest average reward.
235    ///
236    /// Returns `None` if none of the provided options have any recorded rewards.
237    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    /// Return the chronological reward history for a given variant.
250    pub fn reward_history(&self, variant: &str) -> Vec<f64> {
251        self.reward_model.history(variant, "default")
252    }
253
254    /// Compute a convergence score as the variance of the most recent rewards
255    /// across all known variants.
256    ///
257    /// A low value (close to 0) indicates that rewards have stabilised.
258    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        // quality=1.0, cost_penalty=0.5, latency_penalty=0.25
297        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        // penalties should both be 0
306        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        // Should retain exactly 3 entries
317        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        // After two updates: first avg=1.0, second avg=0.9*1.0 + 0.1*0.0 = 0.9
327        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        // record high-reward feedbacks under "good_model"
343        for _ in 0..5 {
344            fl.record(make_feedback("good_model:req1", 1.0, 100, 100));
345        }
346        // record low-reward feedbacks under "bad_model"
347        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        // Single variant — variance is 0
369        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        // Two very different variants → variance > 0
378        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}