1use std::collections::HashMap;
4use std::sync::atomic::{AtomicU64, Ordering};
5use std::sync::{Arc, Mutex, OnceLock};
6
7use crate::core::ocla::routing_quality::{RoutingDecision, RoutingOutcome, RoutingQualityTracker};
8
9const MAX_PENDING_DECISIONS: usize = 1_000;
10type PendingDecisions = HashMap<String, RoutingDecision>;
11static NEXT_DECISION_ID: AtomicU64 = AtomicU64::new(1);
12const FALLBACK_PROBE_INTERVAL: u64 = 20;
13
14static GLOBAL_FEEDBACK: OnceLock<RoutingFeedback> = OnceLock::new();
15
16pub fn global_feedback() -> &'static RoutingFeedback {
18 GLOBAL_FEEDBACK.get_or_init(RoutingFeedback::new)
19}
20
21#[derive(Clone, Debug)]
23pub struct RoutingFeedback {
24 tracker: Arc<Mutex<RoutingQualityTracker>>,
25 pending_decisions: Arc<Mutex<PendingDecisions>>,
26 fallback_checks: Arc<AtomicU64>,
27}
28
29impl RoutingFeedback {
30 pub fn new() -> Self {
32 Self {
33 tracker: Arc::new(Mutex::new(RoutingQualityTracker::new())),
34 pending_decisions: Arc::new(Mutex::new(HashMap::new())),
35 fallback_checks: Arc::new(AtomicU64::new(1)),
36 }
37 }
38
39 pub fn record_decision(&self, original: &str, routed: &str, reason: &str) -> String {
41 let decision_id = format!("route-{}", NEXT_DECISION_ID.fetch_add(1, Ordering::Relaxed));
42 let decision = RoutingDecision {
43 decision_id: decision_id.clone(),
44 original_model: original.to_string(),
45 routed_model: routed.to_string(),
46 reason: reason.to_string(),
47 timestamp: chrono::Utc::now().to_rfc3339(),
48 };
49 let mut pending = self
50 .pending_decisions
51 .lock()
52 .expect("routing feedback pending decision mutex poisoned");
53 pending.insert(decision_id.clone(), decision);
54 while pending.len() > MAX_PENDING_DECISIONS {
55 let Some(evicted_key) = pending
56 .iter()
57 .min_by_key(|(_, decision)| &decision.timestamp)
58 .map(|(key, _)| key.clone())
59 else {
60 break;
61 };
62 pending.remove(&evicted_key);
63 }
64 decision_id
65 }
66
67 pub fn record_outcome_for_decision(
69 &self,
70 decision_id: &str,
71 quality: Option<f64>,
72 tokens_saved: u64,
73 latency_delta_ms: i64,
74 ) {
75 let decision = {
76 let mut pending = self
77 .pending_decisions
78 .lock()
79 .expect("routing feedback pending decision mutex poisoned");
80 pending.remove(decision_id)
81 };
82 if let Some(decision) = decision {
83 self.tracker
84 .lock()
85 .expect("routing feedback tracker mutex poisoned")
86 .record(RoutingOutcome {
87 decision,
88 quality_score: quality,
89 tokens_saved,
90 latency_delta_ms,
91 });
92 }
93 }
94
95 pub fn record_outcome(
97 &self,
98 original: &str,
99 routed: &str,
100 quality: Option<f64>,
101 tokens_saved: u64,
102 latency_delta_ms: i64,
103 ) {
104 let decision = RoutingDecision {
105 decision_id: format!(
106 "unmatched-{}",
107 NEXT_DECISION_ID.fetch_add(1, Ordering::Relaxed)
108 ),
109 original_model: original.to_string(),
110 routed_model: routed.to_string(),
111 reason: "outcome without recorded decision".to_string(),
112 timestamp: chrono::Utc::now().to_rfc3339(),
113 };
114
115 self.tracker
116 .lock()
117 .expect("routing feedback tracker mutex poisoned")
118 .record(RoutingOutcome {
119 decision,
120 quality_score: quality,
121 tokens_saved,
122 latency_delta_ms,
123 });
124 }
125
126 pub fn should_use_fallback(&self) -> bool {
128 let should_fallback = self
129 .tracker
130 .lock()
131 .expect("routing feedback tracker mutex poisoned")
132 .should_fallback();
133 should_fallback
134 && !self
135 .fallback_checks
136 .fetch_add(1, Ordering::Relaxed)
137 .is_multiple_of(FALLBACK_PROBE_INTERVAL)
138 }
139
140 pub fn stats(&self) -> (f64, f64) {
142 let tracker = self
143 .tracker
144 .lock()
145 .expect("routing feedback tracker mutex poisoned");
146 (tracker.success_rate(), tracker.average_savings())
147 }
148}
149
150impl Default for RoutingFeedback {
151 fn default() -> Self {
152 Self::new()
153 }
154}
155
156#[cfg(test)]
157mod tests {
158 use super::*;
159
160 #[test]
161 fn records_decision_until_matching_outcome() {
162 let feedback = RoutingFeedback::new();
163
164 let decision_id = feedback.record_decision("expensive", "fast", "token budget");
165 assert_eq!(feedback.stats(), (0.0, 0.0));
166 assert!(!feedback.should_use_fallback());
167
168 let pending = feedback
169 .pending_decisions
170 .lock()
171 .expect("test pending decision mutex poisoned");
172 let decision = &pending[&decision_id];
173 assert_eq!(decision.reason, "token budget");
174 }
175
176 #[test]
177 fn records_successful_outcome_and_statistics() {
178 let feedback = RoutingFeedback::new();
179
180 let decision_id = feedback.record_decision("expensive", "fast", "token budget");
181 feedback.record_outcome_for_decision(&decision_id, Some(0.95), 120, -10);
182
183 assert_eq!(feedback.stats(), (1.0, 120.0));
184 assert!(!feedback.should_use_fallback());
185 }
186
187 #[test]
188 fn poor_outcome_triggers_fallback() {
189 let feedback = RoutingFeedback::new();
190
191 for _ in 0..20 {
192 feedback.record_outcome("expensive", "fast", Some(0.4), 20, 15);
193 }
194
195 assert_eq!(feedback.stats(), (0.0, 20.0));
196 assert!(feedback.should_use_fallback());
197 }
198
199 #[test]
200 fn concurrent_same_pair_outcomes_match_by_decision_id() {
201 let feedback = RoutingFeedback::new();
202 let first = feedback.record_decision("expensive", "fast", "first");
203 let second = feedback.record_decision("expensive", "fast", "second");
204
205 feedback.record_outcome_for_decision(&second, Some(1.0), 20, 0);
206
207 let pending = feedback
208 .pending_decisions
209 .lock()
210 .expect("test pending decision mutex poisoned");
211 assert!(pending.contains_key(&first));
212 assert!(!pending.contains_key(&second));
213 assert_eq!(feedback.stats(), (1.0, 20.0));
214 }
215
216 #[test]
217 fn fallback_allows_periodic_recovery_probe() {
218 let feedback = RoutingFeedback::new();
219 for _ in 0..20 {
220 feedback.record_outcome("expensive", "fast", Some(0.4), 20, 15);
221 }
222
223 let mut allowed_probe = false;
224 for _ in 0..FALLBACK_PROBE_INTERVAL {
225 allowed_probe |= !feedback.should_use_fallback();
226 }
227
228 assert!(allowed_probe);
229 }
230}