Skip to main content

zeph_llm/router/
select.rs

1// SPDX-FileCopyrightText: 2026 Andrei G <bug-ops>
2// SPDX-License-Identifier: MIT OR Apache-2.0
3
4//! Provider selection and ordering for [`RouterProvider`].
5//!
6//! Implements the per-strategy ordering (`ema_ordered_providers`,
7//! `thompson_ordered_providers`), bandit feature extraction and selection, availability
8//! recording, cascade quality evaluation, and the per-turn embedding cache plus ASI
9//! coherence update spawning.
10
11use 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    /// Emit a rate-limited warn (once per 60 s) when a provider's ASI coherence drops below
27    /// threshold. Falls back to a trace-level message while the rate limit is active.
28    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    /// Hash a query string to a `u64` cache key.
56    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    /// Fetch or compute the feature vector for `query` using the bandit embedding provider.
64    ///
65    /// Returns `None` if:
66    /// - No embedding provider is configured.
67    /// - The embedding call exceeds `embedding_timeout_ms`.
68    /// - The embedding is shorter than `dim` or is all-zero.
69    #[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        // Check cache first (no async needed).
75        {
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        // Insert into cache.
103        {
104            let mut cache = self.bandit_embed_cache.lock();
105            cache.insert(key, features.clone());
106        }
107        Some(features)
108    }
109
110    /// Select a provider using `LinUCB` bandit, with Thompson fallback on cold start / missing features.
111    ///
112    /// Falls through to Thompson or first available provider when bandit cannot decide.
113    /// Budget enforcement via global `CostTracker` is handled at the caller level.
114    /// Per-provider budget fractions are intentionally NOT implemented (scope creep, see #2230).
115    #[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        // Try LinUCB selection with feature vector.
130        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        // Fallback: Thompson sampling.
171        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        // Last resort: first provider.
189        self.state.providers.first().cloned()
190    }
191
192    /// Record the bandit reward for a completed request.
193    ///
194    /// `quality_score`: heuristic quality in [0, 1] from `heuristic_score()`.
195    /// `cost_fraction`: `request_cost_cents / max_daily_cents` (0 when budget is unlimited).
196    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            // Cascade/Bandit: sync path used only for debug_request_json(); hot paths use
227            // dedicated async selection methods. For Cascade, providers are sorted at
228            // construction time.
229            RouterStrategy::Cascade | RouterStrategy::Bandit => self.state.providers.to_vec(),
230        }
231    }
232
233    /// Candidate providers for `embed`/`embed_batch`, with the dedicated `embed = true`
234    /// provider (if configured via [`crate::router::RouterProvider::with_embed_provider`])
235    /// moved to the front.
236    ///
237    /// This keeps the existing `supports_embeddings()` fallback loop intact while ensuring
238    /// a provider explicitly configured for embeddings is tried before any provider that
239    /// merely reports `supports_embeddings() == true` (#5859).
240    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        // CRIT-2 fix: apply reputation as a multiplicative adjustment to the EMA score,
257        // not an additive term. This avoids unbounded score inflation.
258        //
259        // Adjustment formula: ema_score * (1 + weight * (rep_factor - 0.5) * 2)
260        // where rep_factor in [0,1]: 0.5 = neutral, >0.5 = positive, <0.5 = negative.
261        // CRIT-1 fix: reputation factor is sampled per-provider (each has its own Beta mean).
262        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                        // Multiplicative blend: neutral at rep_factor=0.5, range ±weight.
277                        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        // ASI: re-score by down-weighting providers with low coherence.
294        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                    // Multiply EMA score by coherence multiplier clamped to [0.5, 1.0].
314                    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        // Compute per-provider prior overrides: start from base Beta distribution, apply
351        // reputation shift (CRIT-3), then apply ASI coherence penalty.
352        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            // Build overrides by composing reputation and ASI adjustments.
357            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                    // Apply reputation prior shift.
367                    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                    // Apply ASI coherence penalty: shift beta by penalty_weight * deficit.
373                    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        // Put selected provider first, keep rest in original order.
408        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    /// Record availability outcome (network success/failure) for EMA or Thompson.
418    ///
419    /// For cascade routing, quality outcomes are tracked separately in `CascadeState`.
420    /// Only availability outcomes (API up/down) are recorded here to avoid corrupting
421    /// Thompson/EMA distributions with quality-based failures (HIGH-01).
422    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                // Cascade does not use Thompson/EMA for ordering; no-op.
435                // Bandit tracks rewards separately via bandit_record_reward().
436            }
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(&current_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    /// Evaluate quality with heuristics only.
468    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    /// Evaluate quality using the configured classifier mode.
475    ///
476    /// For `ClassifierMode::Judge`, calls the summary provider and falls back to heuristic
477    /// on any error or timeout. For `ClassifierMode::Heuristic`, evaluates synchronously.
478    #[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    /// Embed `text` with per-turn caching.
523    ///
524    /// Checks `cache` before calling the underlying provider. On a cache hit, increments
525    /// `embed_cache_hits`; on a miss, embeds via `self.embed()` and populates the cache.
526    /// Either way, `embed_call_count` is incremented for observability.
527    #[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    /// Return session-level embedding cache metrics: `(total_calls, cache_hits)`.
544    #[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    /// Spawn a background task to update the ASI window for `provider`.
553    ///
554    /// Fire-and-forget: routing is not blocked on the embed call. If the embed fails,
555    /// the ASI window is not updated (no penalty for embed failure).
556    ///
557    /// `turn_id` is used to debounce: at most one ASI update fires per turn even when
558    /// `chat()` is called N times concurrently (e.g., tool schema fetches). Subsequent
559    /// calls within the same turn are silently dropped.
560    ///
561    /// `precomputed_embedding` — when `Some`, skips the embed call entirely (reuse from
562    /// quality gate). When `None`, embeds `response` inline in the spawned task.
563    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        // Debounce: swap in turn_id; if the previous value equals turn_id, another call
571        // already claimed this turn → drop silently. `swap` is atomic so exactly one
572        // concurrent caller wins the "first for this turn" race.
573        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        // Drain finished tasks so completed handles don't count toward the cap.
585        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}