1use std::sync::Arc;
12use std::sync::atomic::Ordering;
13
14use parking_lot::Mutex;
15
16use super::asi::AsiState;
17use super::bandit::embedding_to_features;
18use super::cascade::{self, ClassifierMode, heuristic_score};
19use super::embed_cache::TurnEmbedCache;
20use super::{ASI_WARN_LAST_SECS, MAX_ASI_TASKS, RouterProvider, RouterStrategy};
21use crate::any::AnyProvider;
22use crate::ema::EmaTracker;
23use crate::provider::LlmProvider;
24
25impl RouterProvider {
26 fn maybe_warn_asi_coherence(provider: &str, coherence: f32, threshold: f32) {
29 let now = std::time::SystemTime::now()
30 .duration_since(std::time::UNIX_EPOCH)
31 .unwrap_or(std::time::Duration::MAX)
32 .as_secs();
33 let last = ASI_WARN_LAST_SECS.load(Ordering::Relaxed);
34 if now.saturating_sub(last) >= 60
35 && ASI_WARN_LAST_SECS
36 .compare_exchange(last, now, Ordering::Relaxed, Ordering::Relaxed)
37 .is_ok()
38 {
39 tracing::warn!(
40 provider,
41 coherence,
42 threshold,
43 "asi: coherence below threshold"
44 );
45 } else {
46 tracing::trace!(
47 provider,
48 coherence,
49 threshold,
50 "asi: coherence below threshold (warn rate-limited)"
51 );
52 }
53 }
54
55 fn query_hash(query: &str) -> u64 {
57 use std::hash::{Hash as _, Hasher as _};
58 let mut h = std::collections::hash_map::DefaultHasher::new();
59 query.hash(&mut h);
60 h.finish()
61 }
62
63 #[tracing::instrument(name = "llm.router.bandit_features", skip_all)]
70 pub(crate) async fn bandit_features(&self, query: &str) -> Option<Vec<f32>> {
71 let cfg = self.bandit_config.as_ref()?;
72 let key = Self::query_hash(query);
73
74 {
76 let cache = self.bandit_embed_cache.lock();
77 if let Some(cached) = cache.get(key) {
78 return Some(cached.clone());
79 }
80 }
81
82 let provider = self.bandit_embedding_provider.as_ref()?;
83 let timeout = std::time::Duration::from_millis(cfg.embedding_timeout_ms);
84 let embed_future = provider.embed(query);
85 let embedding = match tokio::time::timeout(timeout, embed_future).await {
86 Ok(Ok(emb)) => emb,
87 Ok(Err(e)) => {
88 tracing::debug!(error = %e, "bandit: embedding failed, falling back");
89 return None;
90 }
91 Err(_) => {
92 tracing::debug!(
93 timeout_ms = cfg.embedding_timeout_ms,
94 "bandit: embedding timed out, falling back"
95 );
96 return None;
97 }
98 };
99
100 let features = embedding_to_features(&embedding, cfg.dim)?;
101
102 {
104 let mut cache = self.bandit_embed_cache.lock();
105 cache.insert(key, features.clone());
106 }
107 Some(features)
108 }
109
110 #[tracing::instrument(name = "llm.router.bandit_select_provider", skip_all)]
116 pub(crate) async fn bandit_select_provider(&self, query: &str) -> Option<AnyProvider> {
117 let Some(ref bandit_arc) = self.bandit else {
118 return self.state.providers.first().cloned();
119 };
120 let cfg = self.bandit_config.as_ref()?;
121
122 let names: Vec<String> = self
123 .state
124 .providers
125 .iter()
126 .map(|p| p.name().to_owned())
127 .collect();
128
129 if let Some(features) = self.bandit_features(query).await {
131 let raw = self
132 .state
133 .last_memory_confidence
134 .load(std::sync::atomic::Ordering::Relaxed);
135 let memory_confidence = if raw == u32::MAX {
136 None
137 } else {
138 Some(f32::from_bits(raw))
139 };
140 let selected = {
141 let state = bandit_arc.lock();
142 state.select(
143 &names,
144 &features,
145 cfg.alpha,
146 cfg.warmup_queries,
147 &|_| true,
148 cfg.cost_weight,
149 &self.state.provider_models,
150 memory_confidence,
151 cfg.memory_confidence_threshold,
152 )
153 };
154 if let Some(name) = selected {
155 tracing::debug!(
156 provider = %name,
157 strategy = "bandit",
158 memory_confidence = ?memory_confidence,
159 "selected provider"
160 );
161 return self
162 .state
163 .providers
164 .iter()
165 .find(|p| p.name() == name)
166 .cloned();
167 }
168 }
169
170 if let Some(ref thompson) = self.thompson {
172 let mut state = thompson.lock();
173 if let Some(sel) = state.select(&names) {
174 tracing::debug!(
175 provider = %sel.provider,
176 strategy = "bandit-fallback-thompson",
177 "selected provider"
178 );
179 return self
180 .state
181 .providers
182 .iter()
183 .find(|p| p.name() == sel.provider)
184 .cloned();
185 }
186 }
187
188 self.state.providers.first().cloned()
190 }
191
192 pub(crate) fn bandit_record_reward(
197 &self,
198 provider_name: &str,
199 features: &[f32],
200 quality_score: f64,
201 cost_fraction: f64,
202 ) {
203 let Some(ref bandit_arc) = self.bandit else {
204 return;
205 };
206 let Some(cfg) = &self.bandit_config else {
207 return;
208 };
209 #[allow(clippy::cast_possible_truncation)]
210 let reward = (quality_score as f32) - cfg.cost_weight * (cost_fraction as f32);
211 let reward = reward.clamp(-1.0, 1.0);
212 let mut state = bandit_arc.lock();
213 state.update(provider_name, features, reward);
214 tracing::debug!(
215 provider = provider_name,
216 reward,
217 quality = quality_score,
218 "bandit: recorded reward"
219 );
220 }
221
222 pub(crate) fn ordered_providers(&self) -> Vec<AnyProvider> {
223 match self.strategy {
224 RouterStrategy::Thompson => self.thompson_ordered_providers(),
225 RouterStrategy::Ema => self.ema_ordered_providers(),
226 RouterStrategy::Cascade | RouterStrategy::Bandit => self.state.providers.to_vec(),
230 }
231 }
232
233 pub(crate) fn embed_candidates(&self) -> Vec<AnyProvider> {
241 let mut providers = self.ordered_providers();
242 if let Some(dedicated) = self.state.dedicated_embed_provider.as_deref() {
243 providers.retain(|p| p.name() != dedicated.name());
244 providers.insert(0, dedicated.clone());
245 }
246 providers
247 }
248
249 fn ema_ordered_providers(&self) -> Vec<AnyProvider> {
250 let order = self.state.provider_order.lock();
251 let mut ordered: Vec<AnyProvider> = order
252 .iter()
253 .filter_map(|&i| self.state.providers.get(i).cloned())
254 .collect();
255
256 if let Some(ref reputation) = self.reputation
263 && let Some(ref ema) = self.ema
264 {
265 let rep = reputation.lock();
266 let w = self.reputation_weight;
267 let snap = ema.snapshot();
268 let mut scored: Vec<(usize, f64)> = ordered
269 .iter()
270 .enumerate()
271 .map(|(idx, p)| {
272 let ema_score = snap
273 .get(p.name())
274 .map_or(0.0, |s| s.success_ema - s.latency_ema_ms / 10_000.0);
275 let score = if let Some(rep_factor) = rep.ema_reputation_factor(p.name()) {
276 let adjustment = 1.0 + w * (rep_factor - 0.5) * 2.0;
278 ema_score * adjustment
279 } else {
280 ema_score
281 };
282 (idx, score)
283 })
284 .collect();
285 scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
286 let reordered: Vec<AnyProvider> = scored
287 .into_iter()
288 .filter_map(|(idx, _)| ordered.get(idx).cloned())
289 .collect();
290 ordered = reordered;
291 }
292
293 if let (Some(asi_arc), Some(asi_cfg)) = (&self.asi, &self.asi_config) {
295 let asi: parking_lot::MutexGuard<'_, AsiState> = asi_arc.lock();
296 let snap = self.ema.as_ref().map(EmaTracker::snapshot);
297 let mut scored: Vec<(usize, f64)> = ordered
298 .iter()
299 .enumerate()
300 .map(|(idx, p)| {
301 let coherence = asi.coherence(p.name());
302 if coherence < asi_cfg.coherence_threshold {
303 Self::maybe_warn_asi_coherence(
304 p.name(),
305 coherence,
306 asi_cfg.coherence_threshold,
307 );
308 }
309 let base_score = snap
310 .as_ref()
311 .and_then(|s| s.get(p.name()))
312 .map_or(0.0, |s| s.success_ema - s.latency_ema_ms / 10_000.0);
313 let multiplier = (coherence / asi_cfg.coherence_threshold).clamp(0.5, 1.0);
315 #[allow(clippy::cast_possible_truncation)]
316 let adjusted = base_score * f64::from(multiplier);
317 (idx, adjusted)
318 })
319 .collect();
320 scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
321 let reordered: Vec<AnyProvider> = scored
322 .into_iter()
323 .filter_map(|(idx, _)| ordered.get(idx).cloned())
324 .collect();
325 ordered = reordered;
326 }
327
328 if let Some(first) = ordered.first() {
329 tracing::debug!(
330 provider = %first.name(),
331 strategy = "ema",
332 "selected provider"
333 );
334 }
335 ordered
336 }
337
338 fn thompson_ordered_providers(&self) -> Vec<AnyProvider> {
339 let Some(ref thompson) = self.thompson else {
340 return self.state.providers.to_vec();
341 };
342 let mut state = thompson.lock();
343 let names: Vec<String> = self
344 .state
345 .providers
346 .iter()
347 .map(|p| p.name().to_owned())
348 .collect();
349
350 let has_reputation = self.reputation.is_some();
353 let has_asi = self.asi.is_some() && self.asi_config.is_some();
354
355 let selected = if has_reputation || has_asi {
356 let rep_guard = self.reputation.as_ref().map(|r| r.lock());
358 let asi_guard: Option<parking_lot::MutexGuard<'_, AsiState>> =
359 self.asi.as_ref().map(|a| a.lock());
360 let w = self.reputation_weight;
361
362 let overrides: std::collections::HashMap<String, (f64, f64)> = names
363 .iter()
364 .map(|name| {
365 let base = state.get_distribution(name);
366 let (alpha, mut beta) = if let Some(ref rep) = rep_guard {
368 rep.shift_thompson_priors(name, base.alpha, base.beta, w)
369 } else {
370 (base.alpha, base.beta)
371 };
372 if let (Some(asi), Some(asi_cfg)) = (&asi_guard, &self.asi_config) {
374 let coherence = asi.coherence(name);
375 if coherence < asi_cfg.coherence_threshold {
376 Self::maybe_warn_asi_coherence(
377 name.as_str(),
378 coherence,
379 asi_cfg.coherence_threshold,
380 );
381 let deficit = asi_cfg.coherence_threshold - coherence;
382 let penalty = f64::from(asi_cfg.penalty_weight * deficit);
383 beta += penalty;
384 }
385 }
386 (name.clone(), (alpha, beta))
387 })
388 .collect();
389
390 drop(rep_guard);
391 drop(asi_guard);
392 state.select_with_priors(&names, &overrides)
393 } else {
394 state.select(&names)
395 };
396
397 if let Some(ref sel) = selected {
398 tracing::debug!(
399 provider = %sel.provider,
400 strategy = "thompson",
401 mode = if sel.exploit { "exploit" } else { "explore" },
402 alpha = sel.alpha,
403 beta = sel.beta,
404 "selected provider"
405 );
406 }
407 let mut ordered = self.state.providers.to_vec();
409 if let Some(ref sel) = selected
410 && let Some(pos) = ordered.iter().position(|p| p.name() == sel.provider)
411 {
412 ordered.swap(0, pos);
413 }
414 ordered
415 }
416
417 pub(crate) fn record_availability(&self, provider_name: &str, success: bool, latency_ms: u64) {
423 match self.strategy {
424 RouterStrategy::Thompson => {
425 if let Some(ref thompson) = self.thompson {
426 let mut state = thompson.lock();
427 state.update(provider_name, success);
428 }
429 }
430 RouterStrategy::Ema => {
431 self.ema_record(provider_name, success, latency_ms);
432 }
433 RouterStrategy::Cascade | RouterStrategy::Bandit => {
434 }
437 }
438 }
439
440 fn ema_record(&self, provider_name: &str, success: bool, latency_ms: u64) {
441 let Some(ref ema) = self.ema else {
442 return;
443 };
444 ema.record(provider_name, success, latency_ms);
445 let current_names: Vec<String> = self
446 .state
447 .providers
448 .iter()
449 .map(|p| p.name().to_owned())
450 .collect();
451 if let Some(new_order_names) = ema.maybe_reorder(¤t_names) {
452 let name_to_idx: std::collections::HashMap<&str, usize> = self
453 .state
454 .providers
455 .iter()
456 .enumerate()
457 .map(|(i, p)| (p.name(), i))
458 .collect();
459 let new_order: Vec<usize> = new_order_names
460 .iter()
461 .filter_map(|n| name_to_idx.get(n.as_str()).copied())
462 .collect();
463 let mut order = self.state.provider_order.lock();
464 *order = new_order;
465 }
466 }
467 pub(crate) fn evaluate_heuristic(response: &str, threshold: f64) -> cascade::QualityVerdict {
469 let mut verdict = heuristic_score(response);
470 verdict.should_escalate = verdict.score < threshold;
471 verdict
472 }
473
474 #[tracing::instrument(name = "llm.router.evaluate_quality", skip_all)]
479 pub(crate) async fn evaluate_quality(
480 response: &str,
481 threshold: f64,
482 mode: ClassifierMode,
483 summary_provider: Option<&dyn crate::provider_dyn::LlmProviderDyn>,
484 judge_timeout_ms: u64,
485 ) -> cascade::QualityVerdict {
486 if mode == ClassifierMode::Judge {
487 if let Some(judge) = summary_provider {
488 match cascade::judge_score(
489 judge,
490 response,
491 std::time::Duration::from_millis(judge_timeout_ms),
492 )
493 .await
494 {
495 Some(score) => {
496 let should_escalate = score < threshold;
497 tracing::debug!(
498 score,
499 threshold,
500 should_escalate,
501 "cascade: judge scored response"
502 );
503 return cascade::QualityVerdict {
504 score,
505 should_escalate,
506 reason: format!("judge score: {score:.2}"),
507 };
508 }
509 None => {
510 tracing::warn!("cascade: judge call failed, falling back to heuristic");
511 }
512 }
513 } else {
514 tracing::warn!(
515 "cascade: classifier_mode=judge but no summary_provider configured, \
516 using heuristic"
517 );
518 }
519 }
520 Self::evaluate_heuristic(response, threshold)
521 }
522 #[tracing::instrument(name = "llm.router.embed_cached", skip_all)]
528 pub(crate) async fn embed_cached(
529 &self,
530 text: &str,
531 cache: &Mutex<TurnEmbedCache>,
532 ) -> Result<Vec<f32>, crate::error::LlmError> {
533 self.state.embed_call_count.fetch_add(1, Ordering::Relaxed);
534 if let Some(emb) = cache.lock().get(text) {
535 self.state.embed_cache_hits.fetch_add(1, Ordering::Relaxed);
536 return Ok(emb.clone());
537 }
538 let emb = self.embed(text).await?;
539 cache.lock().insert(text, emb.clone());
540 Ok(emb)
541 }
542
543 #[must_use]
545 pub fn embed_cache_metrics(&self) -> (u64, u64) {
546 (
547 self.state.embed_call_count.load(Ordering::Relaxed),
548 self.state.embed_cache_hits.load(Ordering::Relaxed),
549 )
550 }
551
552 pub(crate) fn spawn_asi_update(
564 &self,
565 provider: &str,
566 response: String,
567 turn_id: u64,
568 precomputed_embedding: Option<Vec<f32>>,
569 ) {
570 let prev = self.state.asi_last_turn.swap(turn_id, Ordering::AcqRel);
574 if prev == turn_id {
575 return;
576 }
577
578 let Some(ref asi_arc) = self.asi else { return };
579 let Some(ref asi_cfg) = self.asi_config else {
580 return;
581 };
582
583 let mut tasks = self.asi_tasks.lock();
584 while tasks.try_join_next().is_some() {}
586 if tasks.len() >= MAX_ASI_TASKS {
587 tracing::debug!("asi: task limit reached, skipping coherence update");
588 return;
589 }
590
591 let asi = Arc::clone(asi_arc);
592 let router = self.clone();
593 let window_size = asi_cfg.window;
594 let provider_name = provider.to_owned();
595 let embed_timeout_ms = self.embed_timeout_ms;
596 tasks.spawn(async move {
597 let emb = if let Some(e) = precomputed_embedding {
598 e
599 } else {
600 let embed_fut = router.embed(&response);
601 let embed_result = if embed_timeout_ms > 0 {
602 let timeout = std::time::Duration::from_millis(embed_timeout_ms);
603 if let Ok(r) = tokio::time::timeout(timeout, embed_fut).await {
604 r
605 } else {
606 tracing::debug!(
607 provider = provider_name,
608 timeout_ms = embed_timeout_ms,
609 "asi: embed timed out, skipping coherence update"
610 );
611 return;
612 }
613 } else {
614 embed_fut.await
615 };
616 match embed_result {
617 Ok(e) => e,
618 Err(err) => {
619 tracing::debug!(
620 provider = provider_name,
621 error = %err,
622 "asi: embed failed, skipping coherence update"
623 );
624 return;
625 }
626 }
627 };
628 let mut state = asi.lock();
629 state.push_embedding(&provider_name, emb, window_size);
630 });
631 }
632}