Skip to main content

turnframe_provider/
router.rs

1//! Provider selection (spec §20.4, §20.6, ADR-008).
2//!
3//! Routing answers one question: *which configured provider-model profiles may
4//! serve this stage, and in what order?* [`PolicyRouter`] answers it in three
5//! passes, and the order of the passes is the safety property:
6//!
7//! 1. **Capability fit.** A profile that cannot satisfy the stage's
8//!    [`CapabilityRequirements`] is out. This pass runs first and is never
9//!    relaxed by the ones that follow.
10//! 2. **Tenant policy.** Allowlist, denylist, data residency, cost ceiling and
11//!    the sensitivity-to-provider mapping of [`RoutingPolicy`].
12//! 3. **Preference and health.** Surviving candidates are ordered by the
13//!    policy's preference list, then by pool declaration order, with healthy
14//!    profiles ahead of degraded ones.
15//!
16//! # No silent downgrade
17//!
18//! When nothing survives, [`select`](ProviderRouter::select) returns a
19//! [`RoutingError`] that names every candidate it considered and why each was
20//! rejected — never an empty list a caller might proceed past, and never a
21//! weaker profile substituted for a strong one. If the structured-output
22//! requirement is what went unmet,
23//! [`RoutingError::structured_output_unmet`] says so, and the runtime's only
24//! correct responses are to route elsewhere or to reject the operation (spec §0
25//! rule 9). Lowering the requirement is not one of them.
26//!
27//! # Health is availability, never safety
28//!
29//! [`ProviderPool`] tracks consecutive failures per profile against an
30//! injectable [`Clock`], so a test can drive a cooldown without sleeping. A
31//! degraded profile is ordered last and dropped when a healthy alternative
32//! exists — but it is never dropped when it is the only thing that fits, since
33//! refusing to call a provider that might work is a self-inflicted outage, not
34//! a safety control.
35
36use std::collections::{BTreeMap, BTreeSet};
37use std::fmt;
38use std::sync::{Arc, Mutex, PoisonError};
39use std::time::Duration;
40
41use chrono::{DateTime, Utc};
42use turnframe_core::read::DataSensitivity;
43
44use crate::capabilities::{
45    CapabilityMismatch, CapabilityRequirements, MicroCents, ModelProfile,
46    StructuredOutputCapability,
47};
48use crate::error::RetryClass;
49use crate::ids::{ModelRef, ProviderKey};
50use crate::provider::ModelProvider;
51use crate::purpose::ModelPurpose;
52
53/// Source of the current time, so health windows are testable.
54pub trait Clock: Send + Sync + fmt::Debug {
55    /// The current instant.
56    fn now(&self) -> DateTime<Utc>;
57}
58
59/// The wall clock.
60#[derive(Debug, Clone, Copy, Default)]
61pub struct SystemClock;
62
63impl Clock for SystemClock {
64    fn now(&self) -> DateTime<Utc> {
65        Utc::now()
66    }
67}
68
69/// When a profile is considered degraded, and for how long.
70#[derive(Debug, Clone, Copy, PartialEq, Eq)]
71pub struct HealthPolicy {
72    /// Consecutive failures that mark a profile degraded.
73    pub failure_threshold: u32,
74    /// How long it stays degraded after the threshold is crossed.
75    pub cooldown: Duration,
76}
77
78impl HealthPolicy {
79    /// Three consecutive failures, thirty seconds of cooldown.
80    pub const DEFAULT: Self = Self {
81        failure_threshold: 3,
82        cooldown: Duration::from_secs(30),
83    };
84}
85
86impl Default for HealthPolicy {
87    fn default() -> Self {
88        Self::DEFAULT
89    }
90}
91
92/// A snapshot of one profile's recent record.
93#[derive(Debug, Clone, PartialEq, Eq, Default)]
94pub struct ProviderHealth {
95    /// Failures since the last success.
96    pub consecutive_failures: u32,
97    /// When the profile becomes eligible again, when it is degraded.
98    pub degraded_until: Option<DateTime<Utc>>,
99    /// Successful calls recorded.
100    pub successes: u64,
101    /// Failed calls recorded.
102    pub failures: u64,
103}
104
105impl ProviderHealth {
106    /// Returns `true` when the profile is not in a cooldown window at `now`.
107    #[must_use]
108    pub fn is_healthy_at(&self, now: DateTime<Utc>) -> bool {
109        self.degraded_until.is_none_or(|until| now >= until)
110    }
111}
112
113/// A profile the router is offering, with the provider that serves it.
114#[derive(Clone)]
115pub struct ProviderCandidate {
116    /// The adapter to call.
117    pub provider: Arc<dyn ModelProvider>,
118    /// Its routing profile.
119    pub profile: ModelProfile,
120    /// Whether the pool considers it healthy right now.
121    pub healthy: bool,
122}
123
124impl ProviderCandidate {
125    /// The provider-model pair.
126    #[must_use]
127    pub fn reference(&self) -> ModelRef {
128        self.profile.reference()
129    }
130}
131
132impl fmt::Debug for ProviderCandidate {
133    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
134        f.debug_struct("ProviderCandidate")
135            .field("model", &self.reference().to_string())
136            .field("healthy", &self.healthy)
137            .finish_non_exhaustive()
138    }
139}
140
141/// Tenant and deployment constraints on routing (spec §20.6, §25.5).
142///
143/// Every filter is opt-in except the sensitivity rule, which fails closed for
144/// [`Confidential`](DataSensitivity::Confidential) and
145/// [`Restricted`](DataSensitivity::Restricted): sending regulated data to a
146/// provider nobody declared is exactly the mistake this field exists to
147/// prevent.
148#[derive(Debug, Clone)]
149pub struct RoutingPolicy {
150    /// When set, only these providers may be used.
151    pub allowlist: Option<BTreeSet<ProviderKey>>,
152    /// Providers that may never be used, whatever the allowlist says.
153    pub denylist: BTreeSet<ProviderKey>,
154    /// Ceiling on the profile's higher per-million price. A profile with an
155    /// unknown price does not pass a ceiling.
156    pub max_cost_per_million: Option<MicroCents>,
157    /// When set, only profiles declaring one of these regions may be used. A
158    /// profile with no declared region does not pass a residency requirement.
159    pub allowed_regions: Option<BTreeSet<String>>,
160    /// How sensitive the payload of this call is.
161    pub sensitivity: DataSensitivity,
162    /// Which providers may see data at each sensitivity level.
163    ///
164    /// An absent entry means "no restriction" for
165    /// [`Public`](DataSensitivity::Public) and
166    /// [`Internal`](DataSensitivity::Internal), and "nothing is allowed" for
167    /// [`Confidential`](DataSensitivity::Confidential) and
168    /// [`Restricted`](DataSensitivity::Restricted).
169    pub sensitivity_allowlist: BTreeMap<DataSensitivity, BTreeSet<ProviderKey>>,
170    /// Profiles to try first, in this order. Anything not listed keeps the
171    /// pool's declaration order, after the listed ones.
172    pub preference: Vec<ModelRef>,
173    /// Whether degraded profiles are dropped when a healthy one remains.
174    /// Defaults to `true` through [`RoutingPolicy::new`].
175    pub skip_degraded: bool,
176    /// When set, only profiles carrying this tag may be used. No profile carrying
177    /// it is an error, never a quiet fallback to an untagged one.
178    pub required_tag: Option<String>,
179}
180
181impl Default for RoutingPolicy {
182    /// Same as [`RoutingPolicy::new`].
183    fn default() -> Self {
184        Self::new()
185    }
186}
187
188impl RoutingPolicy {
189    /// A policy with no restriction beyond capability fit, treating the payload
190    /// as [`Internal`](DataSensitivity::Internal) and skipping degraded
191    /// profiles when a healthy one remains.
192    #[must_use]
193    pub fn new() -> Self {
194        Self {
195            allowlist: None,
196            denylist: BTreeSet::new(),
197            max_cost_per_million: None,
198            allowed_regions: None,
199            sensitivity: DataSensitivity::Internal,
200            sensitivity_allowlist: BTreeMap::new(),
201            preference: Vec::new(),
202            skip_degraded: true,
203            required_tag: None,
204        }
205    }
206
207    /// Admits only profiles carrying `tag`.
208    #[must_use]
209    pub fn with_required_tag(mut self, tag: impl Into<String>) -> Self {
210        self.required_tag = Some(tag.into());
211        self
212    }
213
214    /// Restricts routing to these providers.
215    #[must_use]
216    pub fn with_allowlist<I: IntoIterator<Item = ProviderKey>>(mut self, providers: I) -> Self {
217        self.allowlist = Some(providers.into_iter().collect());
218        self
219    }
220
221    /// Forbids these providers.
222    #[must_use]
223    pub fn with_denylist<I: IntoIterator<Item = ProviderKey>>(mut self, providers: I) -> Self {
224        self.denylist = providers.into_iter().collect();
225        self
226    }
227
228    /// Sets the cost ceiling.
229    #[must_use]
230    pub fn with_max_cost(mut self, ceiling: MicroCents) -> Self {
231        self.max_cost_per_million = Some(ceiling);
232        self
233    }
234
235    /// Restricts routing to these regions.
236    #[must_use]
237    pub fn with_regions<I: IntoIterator<Item = String>>(mut self, regions: I) -> Self {
238        self.allowed_regions = Some(regions.into_iter().collect());
239        self
240    }
241
242    /// Declares the payload's sensitivity.
243    #[must_use]
244    pub fn with_sensitivity(mut self, sensitivity: DataSensitivity) -> Self {
245        self.sensitivity = sensitivity;
246        self
247    }
248
249    /// Declares which providers may see data at `level`.
250    #[must_use]
251    pub fn allowing<I: IntoIterator<Item = ProviderKey>>(
252        mut self,
253        level: DataSensitivity,
254        providers: I,
255    ) -> Self {
256        self.sensitivity_allowlist
257            .insert(level, providers.into_iter().collect());
258        self
259    }
260
261    /// Sets the preference order.
262    #[must_use]
263    pub fn preferring<I: IntoIterator<Item = ModelRef>>(mut self, order: I) -> Self {
264        self.preference = order.into_iter().collect();
265        self
266    }
267
268    /// Checks one profile against the policy filters.
269    fn admits(&self, profile: &ModelProfile) -> Result<(), RejectionReason> {
270        if let Some(tag) = &self.required_tag
271            && !profile.tags.contains(tag)
272        {
273            return Err(RejectionReason::MissingTag { tag: tag.clone() });
274        }
275        if self.denylist.contains(&profile.provider) {
276            return Err(RejectionReason::Denylisted);
277        }
278        if let Some(allowlist) = &self.allowlist
279            && !allowlist.contains(&profile.provider)
280        {
281            return Err(RejectionReason::NotAllowlisted);
282        }
283        match self.sensitivity_allowlist.get(&self.sensitivity) {
284            Some(allowed) if !allowed.contains(&profile.provider) => {
285                return Err(RejectionReason::Sensitivity {
286                    level: self.sensitivity,
287                });
288            }
289            None if self.sensitivity >= DataSensitivity::Confidential => {
290                return Err(RejectionReason::Sensitivity {
291                    level: self.sensitivity,
292                });
293            }
294            _ => {}
295        }
296        if let Some(regions) = &self.allowed_regions {
297            let admitted = profile
298                .region
299                .as_ref()
300                .is_some_and(|region| regions.contains(region));
301            if !admitted {
302                return Err(RejectionReason::Region {
303                    declared: profile.region.clone(),
304                });
305            }
306        }
307        if let Some(ceiling) = self.max_cost_per_million {
308            let cost = profile.max_cost_per_million();
309            if !cost.is_some_and(|cost| cost <= ceiling) {
310                return Err(RejectionReason::CostCeiling {
311                    declared: cost,
312                    ceiling,
313                });
314            }
315        }
316        Ok(())
317    }
318
319    /// Index of `model` in the preference list, or the end.
320    fn preference_rank(&self, model: &ModelRef) -> usize {
321        self.preference
322            .iter()
323            .position(|preferred| preferred == model)
324            .unwrap_or(usize::MAX)
325    }
326}
327
328/// Why one profile did not become a candidate.
329#[derive(Debug, Clone, PartialEq, Eq)]
330#[non_exhaustive]
331pub enum RejectionReason {
332    /// It cannot do what the stage needs.
333    Capability(CapabilityMismatch),
334    /// The call asked for a tag the profile does not carry.
335    MissingTag {
336        /// The tag required.
337        tag: String,
338    },
339    /// An allowlist is in force and does not name it.
340    NotAllowlisted,
341    /// A denylist names it.
342    Denylisted,
343    /// Its region is not among the allowed ones.
344    Region {
345        /// What the profile declares, when it declares anything.
346        declared: Option<String>,
347    },
348    /// It is too expensive, or its price is unknown while a ceiling is set.
349    CostCeiling {
350        /// The profile's higher per-million price, when known.
351        declared: Option<MicroCents>,
352        /// The ceiling in force.
353        ceiling: MicroCents,
354    },
355    /// It is not allowed to see data at this sensitivity.
356    Sensitivity {
357        /// The level of the payload.
358        level: DataSensitivity,
359    },
360    /// It is in a health cooldown and a healthy alternative existed.
361    Degraded,
362}
363
364impl RejectionReason {
365    /// Stable snake-case label.
366    #[must_use]
367    pub const fn as_str(&self) -> &'static str {
368        match self {
369            Self::Capability(_) => "capability",
370            Self::MissingTag { .. } => "missing_tag",
371            Self::NotAllowlisted => "not_allowlisted",
372            Self::Denylisted => "denylisted",
373            Self::Region { .. } => "region",
374            Self::CostCeiling { .. } => "cost_ceiling",
375            Self::Sensitivity { .. } => "sensitivity",
376            Self::Degraded => "degraded",
377        }
378    }
379}
380
381impl fmt::Display for RejectionReason {
382    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
383        match self {
384            Self::Capability(mismatch) => write!(f, "{mismatch}"),
385            Self::Region { declared } => match declared {
386                Some(region) => write!(f, "region({region})"),
387                None => f.write_str("region(undeclared)"),
388            },
389            Self::CostCeiling { declared, ceiling } => match declared {
390                Some(cost) => write!(f, "cost_ceiling({cost} > {ceiling})"),
391                None => write!(f, "cost_ceiling(unknown price, ceiling {ceiling})"),
392            },
393            Self::Sensitivity { level } => write!(f, "sensitivity({level:?})"),
394            Self::MissingTag { tag } => write!(f, "missing_tag({tag})"),
395            other => f.write_str(other.as_str()),
396        }
397    }
398}
399
400/// One profile the router looked at and turned down.
401#[derive(Debug, Clone, PartialEq, Eq)]
402pub struct CandidateRejection {
403    /// Which profile.
404    pub model: ModelRef,
405    /// Why.
406    pub reason: RejectionReason,
407}
408
409impl fmt::Display for CandidateRejection {
410    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
411        write!(f, "{}: {}", self.model, self.reason)
412    }
413}
414
415/// Routing produced no candidate.
416#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
417#[non_exhaustive]
418pub enum RoutingError {
419    /// The pool holds no profile at all.
420    #[error("no provider is configured for {purpose}")]
421    NoProvidersConfigured {
422        /// The stage that needed one.
423        purpose: ModelPurpose,
424    },
425    /// Every configured profile was rejected.
426    #[error("no provider can serve {purpose}: {}", DisplayRejections(.rejections))]
427    NoCandidate {
428        /// The stage that needed one.
429        purpose: ModelPurpose,
430        /// The structured-output transports the stage required, when it
431        /// required any. Kept separate so the caller can report exactly what
432        /// could not be met without re-deriving it.
433        required_structured_output: Vec<StructuredOutputCapability>,
434        /// Every profile considered, with its reason, in pool order.
435        rejections: Vec<CandidateRejection>,
436    },
437}
438
439impl RoutingError {
440    /// Returns `true` when at least one profile failed on the structured-output
441    /// requirement — the case spec §0 rule 9 forbids resolving by downgrading.
442    #[must_use]
443    pub fn structured_output_unmet(&self) -> bool {
444        match self {
445            Self::NoProvidersConfigured { .. } => false,
446            Self::NoCandidate { rejections, .. } => rejections.iter().any(|rejection| {
447                matches!(&rejection.reason, RejectionReason::Capability(mismatch)
448                    if mismatch.structured_output_unmet())
449            }),
450        }
451    }
452
453    /// The stage that could not be routed.
454    #[must_use]
455    pub const fn purpose(&self) -> ModelPurpose {
456        match self {
457            Self::NoProvidersConfigured { purpose } | Self::NoCandidate { purpose, .. } => *purpose,
458        }
459    }
460}
461
462/// Renders a rejection list for [`RoutingError`]'s `Display`.
463struct DisplayRejections<'a>(&'a [CandidateRejection]);
464
465impl fmt::Display for DisplayRejections<'_> {
466    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
467        for (index, rejection) in self.0.iter().enumerate() {
468            if index > 0 {
469                f.write_str("; ")?;
470            }
471            write!(f, "{rejection}")?;
472        }
473        Ok(())
474    }
475}
476
477/// Selects the providers that may serve a stage (spec §20.6).
478pub trait ProviderRouter: Send + Sync {
479    /// Returns the candidates for `purpose`, best first.
480    ///
481    /// # Errors
482    ///
483    /// Returns [`RoutingError`] when nothing qualifies. An empty `Ok` list is
484    /// never returned: a caller must not be able to skip past "nothing fits".
485    fn select(
486        &self,
487        purpose: ModelPurpose,
488        requirements: &CapabilityRequirements,
489        policy: &RoutingPolicy,
490    ) -> Result<Vec<ProviderCandidate>, RoutingError>;
491}
492
493/// One configured profile inside a [`ProviderPool`].
494struct PoolEntry {
495    provider: Arc<dyn ModelProvider>,
496    profile: ModelProfile,
497}
498
499/// A pool could not be built.
500#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
501#[non_exhaustive]
502pub enum PoolError {
503    /// Two entries claim the same provider-model pair.
504    #[error("profile {model} is configured twice")]
505    DuplicateProfile {
506        /// The repeated pair.
507        model: ModelRef,
508    },
509}
510
511/// The configured profiles and their health.
512///
513/// Build one with [`ProviderPool::builder`] and share it behind an [`Arc`]; the
514/// health map is behind a mutex so recording an outcome needs only `&self`.
515pub struct ProviderPool {
516    entries: Vec<PoolEntry>,
517    health: Mutex<BTreeMap<ModelRef, ProviderHealth>>,
518    clock: Arc<dyn Clock>,
519    policy: HealthPolicy,
520}
521
522impl ProviderPool {
523    /// Starts building a pool.
524    #[must_use]
525    pub fn builder() -> ProviderPoolBuilder {
526        ProviderPoolBuilder::new()
527    }
528
529    /// How many profiles are configured.
530    #[must_use]
531    pub fn len(&self) -> usize {
532        self.entries.len()
533    }
534
535    /// Returns `true` when no profile is configured.
536    #[must_use]
537    pub fn is_empty(&self) -> bool {
538        self.entries.is_empty()
539    }
540
541    /// Every configured profile, in declaration order.
542    #[must_use]
543    pub fn profiles(&self) -> Vec<ModelProfile> {
544        self.entries
545            .iter()
546            .map(|entry| entry.profile.clone())
547            .collect()
548    }
549
550    /// The adapter serving `model`, when configured.
551    #[must_use]
552    pub fn provider(&self, model: &ModelRef) -> Option<Arc<dyn ModelProvider>> {
553        self.entries
554            .iter()
555            .find(|entry| entry.profile.reference() == *model)
556            .map(|entry| Arc::clone(&entry.provider))
557    }
558
559    /// The current time according to the injected clock.
560    #[must_use]
561    pub fn now(&self) -> DateTime<Utc> {
562        self.clock.now()
563    }
564
565    /// The recorded health of `model`.
566    #[must_use]
567    pub fn health(&self, model: &ModelRef) -> ProviderHealth {
568        self.health
569            .lock()
570            .unwrap_or_else(PoisonError::into_inner)
571            .get(model)
572            .cloned()
573            .unwrap_or_default()
574    }
575
576    /// Returns `true` when `model` is outside a cooldown window.
577    #[must_use]
578    pub fn is_healthy(&self, model: &ModelRef) -> bool {
579        self.health(model).is_healthy_at(self.clock.now())
580    }
581
582    /// Records a successful call, clearing any cooldown.
583    pub fn record_success(&self, model: &ModelRef) {
584        let mut health = self.health.lock().unwrap_or_else(PoisonError::into_inner);
585        let entry = health.entry(model.clone()).or_default();
586        entry.consecutive_failures = 0;
587        entry.degraded_until = None;
588        entry.successes = entry.successes.saturating_add(1);
589    }
590
591    /// Records a failed call.
592    ///
593    /// Only classes that say something about the provider's availability count:
594    /// [`Retry`](RetryClass::Retry), [`RetryAfter`](RetryClass::RetryAfter) and
595    /// [`Fallback`](RetryClass::Fallback). A [`Fatal`](RetryClass::Fatal)
596    /// outcome — a refusal, a content filter, a context overflow — means the
597    /// provider answered exactly as asked, so it does not degrade its health.
598    pub fn record_failure(&self, model: &ModelRef, class: RetryClass) {
599        if class == RetryClass::Fatal {
600            return;
601        }
602        let now = self.clock.now();
603        let mut health = self.health.lock().unwrap_or_else(PoisonError::into_inner);
604        let entry = health.entry(model.clone()).or_default();
605        entry.consecutive_failures = entry.consecutive_failures.saturating_add(1);
606        entry.failures = entry.failures.saturating_add(1);
607        if entry.consecutive_failures >= self.policy.failure_threshold {
608            entry.degraded_until = Some(
609                now + chrono::Duration::from_std(self.policy.cooldown)
610                    .unwrap_or_else(|_| chrono::Duration::seconds(30)),
611            );
612        }
613    }
614
615    /// Forgets every recorded outcome.
616    pub fn reset_health(&self) {
617        self.health
618            .lock()
619            .unwrap_or_else(PoisonError::into_inner)
620            .clear();
621    }
622}
623
624impl fmt::Debug for ProviderPool {
625    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
626        let models: Vec<String> = self
627            .entries
628            .iter()
629            .map(|entry| entry.profile.reference().to_string())
630            .collect();
631        f.debug_struct("ProviderPool")
632            .field("profiles", &models)
633            .field("health_policy", &self.policy)
634            .finish_non_exhaustive()
635    }
636}
637
638/// Builds a [`ProviderPool`].
639pub struct ProviderPoolBuilder {
640    entries: Vec<PoolEntry>,
641    clock: Arc<dyn Clock>,
642    policy: HealthPolicy,
643}
644
645impl ProviderPoolBuilder {
646    /// An empty builder using the system clock and the default health policy.
647    #[must_use]
648    pub fn new() -> Self {
649        Self {
650            entries: Vec::new(),
651            clock: Arc::new(SystemClock),
652            policy: HealthPolicy::DEFAULT,
653        }
654    }
655
656    /// Adds a provider, taking its routing profile from
657    /// [`ModelProvider::profile`].
658    #[must_use]
659    pub fn provider(self, provider: Arc<dyn ModelProvider>) -> Self {
660        let profile = provider.profile();
661        self.provider_with_profile(provider, profile)
662    }
663
664    /// Registers a provider with extra tags on its profile, which is how a deployment names
665    /// its tiers (`small`, `large`) for [`RoutingPolicy::with_required_tag`].
666    #[must_use]
667    pub fn provider_tagged<I, T>(self, provider: Arc<dyn ModelProvider>, tags: I) -> Self
668    where
669        I: IntoIterator<Item = T>,
670        T: Into<String>,
671    {
672        let mut profile = provider.profile();
673        for tag in tags {
674            let tag = tag.into();
675            if !profile.tags.contains(&tag) {
676                profile.tags.push(tag);
677            }
678        }
679        self.provider_with_profile(provider, profile)
680    }
681
682    /// Adds a provider with an explicit routing profile, for a deployment that
683    /// knows a price or a region the adapter does not.
684    ///
685    /// The profile's declared capabilities are the ones routing trusts, so a
686    /// deployment may narrow them; widening them past what the adapter reports
687    /// is how a silent downgrade gets built by hand.
688    #[must_use]
689    pub fn provider_with_profile(
690        mut self,
691        provider: Arc<dyn ModelProvider>,
692        profile: ModelProfile,
693    ) -> Self {
694        self.entries.push(PoolEntry { provider, profile });
695        self
696    }
697
698    /// Injects a clock, so health windows can be driven deterministically.
699    #[must_use]
700    pub fn clock<C: Clock + 'static>(mut self, clock: Arc<C>) -> Self {
701        self.clock = clock;
702        self
703    }
704
705    /// Sets the health policy.
706    #[must_use]
707    pub fn health_policy(mut self, policy: HealthPolicy) -> Self {
708        self.policy = policy;
709        self
710    }
711
712    /// Builds the pool.
713    ///
714    /// # Errors
715    ///
716    /// Returns [`PoolError::DuplicateProfile`] when two entries claim the same
717    /// provider-model pair, because health and routing key on that pair and a
718    /// duplicate would make both ambiguous.
719    pub fn build(self) -> Result<ProviderPool, PoolError> {
720        let mut seen = BTreeSet::new();
721        for entry in &self.entries {
722            let reference = entry.profile.reference();
723            if !seen.insert(reference.clone()) {
724                return Err(PoolError::DuplicateProfile { model: reference });
725            }
726        }
727        Ok(ProviderPool {
728            entries: self.entries,
729            health: Mutex::new(BTreeMap::new()),
730            clock: self.clock,
731            policy: self.policy,
732        })
733    }
734}
735
736impl Default for ProviderPoolBuilder {
737    fn default() -> Self {
738        Self::new()
739    }
740}
741
742impl fmt::Debug for ProviderPoolBuilder {
743    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
744        f.debug_struct("ProviderPoolBuilder")
745            .field("entries", &self.entries.len())
746            .field("health_policy", &self.policy)
747            .finish_non_exhaustive()
748    }
749}
750
751/// The router the library ships (spec §20.6).
752///
753/// ```
754/// use std::sync::Arc;
755/// use turnframe_provider::prelude::*;
756/// use turnframe_provider::testing::StaticProvider;
757///
758/// let strong = StaticProvider::new("openai", "gpt-4o").with_capabilities(
759///     ProviderCapabilities::minimal()
760///         .with_structured_output(StructuredOutputCapability::NativeJsonSchema),
761/// );
762/// let weak = StaticProvider::new("local", "llama").with_capabilities(
763///     ProviderCapabilities::minimal()
764///         .with_structured_output(StructuredOutputCapability::PromptOnly),
765/// );
766///
767/// let pool = ProviderPool::builder()
768///     .provider(Arc::new(weak))
769///     .provider(Arc::new(strong))
770///     .build()?;
771/// let router = PolicyRouter::new(Arc::new(pool));
772///
773/// // An understanding task admits only the strong profile.
774/// let requirements = ModelPurpose::Extract.requirements();
775/// let candidates =
776///     router.select(ModelPurpose::Extract, &requirements, &RoutingPolicy::new())?;
777/// assert_eq!(candidates.len(), 1);
778/// assert_eq!(candidates[0].reference().to_string(), "openai/gpt-4o");
779/// # Ok::<(), Box<dyn std::error::Error>>(())
780/// ```
781pub struct PolicyRouter {
782    pool: Arc<ProviderPool>,
783}
784
785impl PolicyRouter {
786    /// Routes over `pool`.
787    #[must_use]
788    pub fn new(pool: Arc<ProviderPool>) -> Self {
789        Self { pool }
790    }
791
792    /// The pool being routed over.
793    #[must_use]
794    pub fn pool(&self) -> &Arc<ProviderPool> {
795        &self.pool
796    }
797
798    /// Selects for a concrete request, folding in the requirements its content
799    /// implies (tools, vision) on top of the purpose's own.
800    ///
801    /// # Errors
802    ///
803    /// Returns [`RoutingError`] when nothing qualifies.
804    pub fn select_for_request(
805        &self,
806        request: &crate::request::ModelRequest,
807        policy: &RoutingPolicy,
808    ) -> Result<Vec<ProviderCandidate>, RoutingError> {
809        let requirements = request.requirements();
810        self.select(request.purpose, &requirements, policy)
811    }
812}
813
814impl ProviderRouter for PolicyRouter {
815    fn select(
816        &self,
817        purpose: ModelPurpose,
818        requirements: &CapabilityRequirements,
819        policy: &RoutingPolicy,
820    ) -> Result<Vec<ProviderCandidate>, RoutingError> {
821        if self.pool.is_empty() {
822            return Err(RoutingError::NoProvidersConfigured { purpose });
823        }
824        let now = self.pool.now();
825        let mut rejections = Vec::new();
826        let mut admitted: Vec<ProviderCandidate> = Vec::new();
827
828        for entry in &self.pool.entries {
829            let model = entry.profile.reference();
830            // Pass 1: capability fit, before anything else may relax it.
831            if let Err(mismatch) = requirements.satisfied_by(&entry.profile.capabilities) {
832                rejections.push(CandidateRejection {
833                    model,
834                    reason: RejectionReason::Capability(mismatch),
835                });
836                continue;
837            }
838            // Pass 2: tenant policy.
839            if let Err(reason) = policy.admits(&entry.profile) {
840                rejections.push(CandidateRejection { model, reason });
841                continue;
842            }
843            let healthy = self.pool.health(&model).is_healthy_at(now);
844            admitted.push(ProviderCandidate {
845                provider: Arc::clone(&entry.provider),
846                profile: entry.profile.clone(),
847                healthy,
848            });
849        }
850
851        // Pass 3: health, then preference, then declaration order.
852        if policy.skip_degraded && admitted.iter().any(|candidate| candidate.healthy) {
853            admitted.retain(|candidate| {
854                if candidate.healthy {
855                    return true;
856                }
857                rejections.push(CandidateRejection {
858                    model: candidate.reference(),
859                    reason: RejectionReason::Degraded,
860                });
861                false
862            });
863        }
864        if admitted.is_empty() {
865            return Err(RoutingError::NoCandidate {
866                purpose,
867                required_structured_output: requirements.structured_output.clone(),
868                rejections,
869            });
870        }
871        // `sort_by_key` is stable, so profiles with equal rank keep pool order.
872        admitted.sort_by_key(|candidate| {
873            (
874                usize::from(!candidate.healthy),
875                policy.preference_rank(&candidate.reference()),
876            )
877        });
878        Ok(admitted)
879    }
880}
881
882impl fmt::Debug for PolicyRouter {
883    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
884        f.debug_struct("PolicyRouter")
885            .field("pool", &self.pool)
886            .finish()
887    }
888}
889
890#[cfg(test)]
891mod tests {
892    use super::*;
893    use crate::capabilities::ProviderCapabilities;
894    use crate::testing::{ManualClock, StaticProvider};
895
896    fn caps(structured: StructuredOutputCapability) -> ProviderCapabilities {
897        ProviderCapabilities::minimal().with_structured_output(structured)
898    }
899
900    fn provider(
901        provider_key: &str,
902        model: &str,
903        structured: StructuredOutputCapability,
904    ) -> Arc<dyn ModelProvider> {
905        Arc::new(StaticProvider::new(provider_key, model).with_capabilities(caps(structured)))
906    }
907
908    fn pool(entries: Vec<Arc<dyn ModelProvider>>) -> Arc<ProviderPool> {
909        let mut builder = ProviderPool::builder();
910        for entry in entries {
911            builder = builder.provider(entry);
912        }
913        Arc::new(builder.build().unwrap())
914    }
915
916    fn understand() -> CapabilityRequirements {
917        ModelPurpose::Extract.requirements()
918    }
919
920    #[test]
921    fn a_required_tag_admits_only_the_profiles_carrying_it() {
922        let pool = Arc::new(
923            ProviderPool::builder()
924                .provider_tagged(
925                    provider("mini", "m", StructuredOutputCapability::NativeJsonSchema),
926                    ["small"],
927                )
928                .provider_tagged(
929                    provider("big", "m", StructuredOutputCapability::NativeJsonSchema),
930                    ["large"],
931                )
932                .build()
933                .unwrap(),
934        );
935        let router = PolicyRouter::new(pool);
936        let large = router
937            .select(
938                ModelPurpose::Extract,
939                &ModelPurpose::Extract.requirements(),
940                &RoutingPolicy::new().with_required_tag("large"),
941            )
942            .unwrap();
943        assert_eq!(large.len(), 1);
944        assert_eq!(large[0].profile.provider.as_str(), "big");
945
946        let error = router
947            .select(
948                ModelPurpose::Extract,
949                &ModelPurpose::Extract.requirements(),
950                &RoutingPolicy::new().with_required_tag("vision"),
951            )
952            .expect_err("no profile carries the tag, and nothing untagged stands in");
953        assert!(error.to_string().contains("missing_tag(vision)"), "{error}");
954    }
955
956    #[test]
957    fn capability_fit_runs_before_policy_and_is_never_relaxed() {
958        let router = PolicyRouter::new(pool(vec![
959            provider("weak", "m", StructuredOutputCapability::PromptOnly),
960            provider("strong", "m", StructuredOutputCapability::NativeJsonSchema),
961        ]));
962        let candidates = router
963            .select(ModelPurpose::Extract, &understand(), &RoutingPolicy::new())
964            .unwrap();
965        assert_eq!(candidates.len(), 1);
966        assert_eq!(candidates[0].profile.provider.as_str(), "strong");
967    }
968
969    #[test]
970    fn no_silent_downgrade_names_what_was_missing() {
971        let router = PolicyRouter::new(pool(vec![
972            provider("weak", "m", StructuredOutputCapability::PromptOnly),
973            provider("weaker", "m", StructuredOutputCapability::None),
974        ]));
975        let error = router
976            .select(ModelPurpose::Extract, &understand(), &RoutingPolicy::new())
977            .unwrap_err();
978        assert!(error.structured_output_unmet());
979        assert_eq!(error.purpose(), ModelPurpose::Extract);
980        let RoutingError::NoCandidate {
981            required_structured_output,
982            rejections,
983            ..
984        } = &error
985        else {
986            panic!("{error:?}");
987        };
988        assert_eq!(
989            required_structured_output,
990            &crate::purpose::MUTATION_SAFE_STRUCTURED_OUTPUT.to_vec()
991        );
992        assert_eq!(rejections.len(), 2);
993        let text = error.to_string();
994        assert!(text.contains("native_json_schema"), "{text}");
995        assert!(text.contains("prompt_only"), "{text}");
996        assert!(text.contains("weaker/m"), "{text}");
997    }
998
999    #[test]
1000    fn an_empty_pool_is_its_own_error() {
1001        let router = PolicyRouter::new(pool(vec![]));
1002        let error = router
1003            .select(
1004                ModelPurpose::Acknowledge,
1005                &CapabilityRequirements::none(),
1006                &RoutingPolicy::new(),
1007            )
1008            .unwrap_err();
1009        assert!(matches!(error, RoutingError::NoProvidersConfigured { .. }));
1010        assert!(!error.structured_output_unmet());
1011    }
1012
1013    #[test]
1014    fn allowlist_denylist_and_preference_shape_the_order() {
1015        let router = PolicyRouter::new(pool(vec![
1016            provider("a", "m", StructuredOutputCapability::NativeJsonSchema),
1017            provider("b", "m", StructuredOutputCapability::NativeJsonSchema),
1018            provider("c", "m", StructuredOutputCapability::NativeJsonSchema),
1019        ]));
1020
1021        let denied = RoutingPolicy::new().with_denylist([ProviderKey::from("a")]);
1022        let candidates = router
1023            .select(ModelPurpose::Extract, &understand(), &denied)
1024            .unwrap();
1025        assert_eq!(candidates.len(), 2);
1026        assert!(
1027            candidates
1028                .iter()
1029                .all(|c| c.profile.provider.as_str() != "a")
1030        );
1031
1032        let allowed = RoutingPolicy::new().with_allowlist([ProviderKey::from("c")]);
1033        let candidates = router
1034            .select(ModelPurpose::Extract, &understand(), &allowed)
1035            .unwrap();
1036        assert_eq!(candidates.len(), 1);
1037        assert_eq!(candidates[0].profile.provider.as_str(), "c");
1038
1039        let preferred = RoutingPolicy::new().preferring([ModelRef::new("c", "m")]);
1040        let candidates = router
1041            .select(ModelPurpose::Extract, &understand(), &preferred)
1042            .unwrap();
1043        let order: Vec<&str> = candidates
1044            .iter()
1045            .map(|c| c.profile.provider.as_str())
1046            .collect();
1047        assert_eq!(
1048            order,
1049            vec!["c", "a", "b"],
1050            "preferred first, then pool order"
1051        );
1052
1053        // A denylist beats an allowlist that names the same provider.
1054        let both = RoutingPolicy::new()
1055            .with_allowlist([ProviderKey::from("a")])
1056            .with_denylist([ProviderKey::from("a")]);
1057        let error = router
1058            .select(ModelPurpose::Extract, &understand(), &both)
1059            .unwrap_err();
1060        assert!(error.to_string().contains("denylisted"), "{error}");
1061    }
1062
1063    #[test]
1064    fn residency_and_cost_fail_closed_on_undeclared_profiles() {
1065        let eu = Arc::new(
1066            StaticProvider::new("eu", "m")
1067                .with_capabilities(caps(StructuredOutputCapability::NativeJsonSchema))
1068                .with_profile_region("eu")
1069                .with_profile_cost(MicroCents::from_cents(10), MicroCents::from_cents(20)),
1070        );
1071        let unknown = provider("unknown", "m", StructuredOutputCapability::NativeJsonSchema);
1072        let router = PolicyRouter::new(pool(vec![eu, unknown]));
1073
1074        let residency = RoutingPolicy::new().with_regions(["eu".to_owned()]);
1075        let candidates = router
1076            .select(ModelPurpose::Extract, &understand(), &residency)
1077            .unwrap();
1078        assert_eq!(
1079            candidates.len(),
1080            1,
1081            "a profile with no region does not pass"
1082        );
1083        assert_eq!(candidates[0].profile.provider.as_str(), "eu");
1084
1085        let ceiling = RoutingPolicy::new().with_max_cost(MicroCents::from_cents(20));
1086        let candidates = router
1087            .select(ModelPurpose::Extract, &understand(), &ceiling)
1088            .unwrap();
1089        assert_eq!(
1090            candidates.len(),
1091            1,
1092            "an unknown price does not pass a ceiling"
1093        );
1094
1095        let too_low = RoutingPolicy::new().with_max_cost(MicroCents::from_cents(5));
1096        let error = router
1097            .select(ModelPurpose::Extract, &understand(), &too_low)
1098            .unwrap_err();
1099        assert!(error.to_string().contains("cost_ceiling"), "{error}");
1100    }
1101
1102    #[test]
1103    fn confidential_data_needs_an_explicit_provider_allowlist() {
1104        let router = PolicyRouter::new(pool(vec![provider(
1105            "openai",
1106            "m",
1107            StructuredOutputCapability::NativeJsonSchema,
1108        )]));
1109
1110        let undeclared = RoutingPolicy::new().with_sensitivity(DataSensitivity::Confidential);
1111        let error = router
1112            .select(ModelPurpose::Extract, &understand(), &undeclared)
1113            .unwrap_err();
1114        assert!(error.to_string().contains("sensitivity"), "{error}");
1115
1116        let declared = RoutingPolicy::new()
1117            .with_sensitivity(DataSensitivity::Confidential)
1118            .allowing(DataSensitivity::Confidential, [ProviderKey::from("openai")]);
1119        assert_eq!(
1120            router
1121                .select(ModelPurpose::Extract, &understand(), &declared)
1122                .unwrap()
1123                .len(),
1124            1
1125        );
1126
1127        // Internal data needs no declaration.
1128        let internal = RoutingPolicy::new().with_sensitivity(DataSensitivity::Internal);
1129        assert_eq!(
1130            router
1131                .select(ModelPurpose::Extract, &understand(), &internal)
1132                .unwrap()
1133                .len(),
1134            1
1135        );
1136    }
1137
1138    #[test]
1139    fn health_degrades_on_a_clock_we_control_and_recovers() {
1140        let clock = Arc::new(ManualClock::at_epoch());
1141        let pool = Arc::new(
1142            ProviderPool::builder()
1143                .provider(provider(
1144                    "a",
1145                    "m",
1146                    StructuredOutputCapability::NativeJsonSchema,
1147                ))
1148                .provider(provider(
1149                    "b",
1150                    "m",
1151                    StructuredOutputCapability::NativeJsonSchema,
1152                ))
1153                .clock(Arc::clone(&clock))
1154                .health_policy(HealthPolicy {
1155                    failure_threshold: 2,
1156                    cooldown: Duration::from_secs(60),
1157                })
1158                .build()
1159                .unwrap(),
1160        );
1161        let router = PolicyRouter::new(Arc::clone(&pool));
1162        let a = ModelRef::new("a", "m");
1163
1164        pool.record_failure(&a, RetryClass::Retry);
1165        assert!(pool.is_healthy(&a), "one failure is not a cooldown");
1166        pool.record_failure(&a, RetryClass::Retry);
1167        assert!(!pool.is_healthy(&a));
1168        assert_eq!(pool.health(&a).consecutive_failures, 2);
1169
1170        // `a` is dropped while `b` is healthy.
1171        let candidates = router
1172            .select(ModelPurpose::Extract, &understand(), &RoutingPolicy::new())
1173            .unwrap();
1174        assert_eq!(candidates.len(), 1);
1175        assert_eq!(candidates[0].profile.provider.as_str(), "b");
1176
1177        // The cooldown expires on our clock, not on wall time.
1178        clock.advance(Duration::from_secs(61));
1179        assert!(pool.is_healthy(&a));
1180        assert_eq!(
1181            router
1182                .select(ModelPurpose::Extract, &understand(), &RoutingPolicy::new())
1183                .unwrap()
1184                .len(),
1185            2
1186        );
1187
1188        // A success clears the counter outright.
1189        pool.record_failure(&a, RetryClass::Fallback);
1190        pool.record_failure(&a, RetryClass::Fallback);
1191        assert!(!pool.is_healthy(&a));
1192        pool.record_success(&a);
1193        assert!(pool.is_healthy(&a));
1194        assert_eq!(pool.health(&a).consecutive_failures, 0);
1195        assert_eq!(pool.health(&a).successes, 1);
1196    }
1197
1198    #[test]
1199    fn a_fatal_outcome_does_not_degrade_health() {
1200        let clock = Arc::new(ManualClock::at_epoch());
1201        let pool = ProviderPool::builder()
1202            .provider(provider(
1203                "a",
1204                "m",
1205                StructuredOutputCapability::NativeJsonSchema,
1206            ))
1207            .clock(clock)
1208            .health_policy(HealthPolicy {
1209                failure_threshold: 1,
1210                cooldown: Duration::from_secs(60),
1211            })
1212            .build()
1213            .unwrap();
1214        let a = ModelRef::new("a", "m");
1215        pool.record_failure(&a, RetryClass::Fatal);
1216        assert!(pool.is_healthy(&a));
1217        assert_eq!(pool.health(&a).failures, 0);
1218        pool.record_failure(&a, RetryClass::Retry);
1219        assert!(!pool.is_healthy(&a));
1220        pool.reset_health();
1221        assert!(pool.is_healthy(&a));
1222    }
1223
1224    #[test]
1225    fn a_degraded_profile_is_still_offered_when_it_is_the_only_fit() {
1226        let clock = Arc::new(ManualClock::at_epoch());
1227        let pool = Arc::new(
1228            ProviderPool::builder()
1229                .provider(provider(
1230                    "a",
1231                    "m",
1232                    StructuredOutputCapability::NativeJsonSchema,
1233                ))
1234                .clock(clock)
1235                .health_policy(HealthPolicy {
1236                    failure_threshold: 1,
1237                    cooldown: Duration::from_secs(60),
1238                })
1239                .build()
1240                .unwrap(),
1241        );
1242        let a = ModelRef::new("a", "m");
1243        pool.record_failure(&a, RetryClass::Retry);
1244        assert!(!pool.is_healthy(&a));
1245
1246        let router = PolicyRouter::new(pool);
1247        let candidates = router
1248            .select(ModelPurpose::Extract, &understand(), &RoutingPolicy::new())
1249            .unwrap();
1250        assert_eq!(candidates.len(), 1);
1251        assert!(!candidates[0].healthy, "offered, and honestly labelled");
1252    }
1253
1254    #[test]
1255    fn a_duplicate_profile_is_a_configuration_error() {
1256        let error = ProviderPool::builder()
1257            .provider(provider(
1258                "a",
1259                "m",
1260                StructuredOutputCapability::NativeJsonSchema,
1261            ))
1262            .provider(provider("a", "m", StructuredOutputCapability::JsonObject))
1263            .build()
1264            .unwrap_err();
1265        assert_eq!(
1266            error,
1267            PoolError::DuplicateProfile {
1268                model: ModelRef::new("a", "m")
1269            }
1270        );
1271    }
1272
1273    #[test]
1274    fn the_pool_answers_lookups_and_renders_safely() {
1275        let pool = pool(vec![provider(
1276            "a",
1277            "m",
1278            StructuredOutputCapability::NativeJsonSchema,
1279        )]);
1280        assert_eq!(pool.len(), 1);
1281        assert!(!pool.is_empty());
1282        assert_eq!(pool.profiles().len(), 1);
1283        assert!(pool.provider(&ModelRef::new("a", "m")).is_some());
1284        assert!(pool.provider(&ModelRef::new("a", "other")).is_none());
1285        let rendered = format!("{pool:?}");
1286        assert!(rendered.contains("a/m"), "{rendered}");
1287    }
1288
1289    #[test]
1290    fn select_for_request_folds_in_the_requests_own_needs() {
1291        use crate::request::{ContentPart, Message, ModelRequest};
1292        let vision = Arc::new(StaticProvider::new("vision", "m").with_capabilities(
1293            caps(StructuredOutputCapability::NativeJsonSchema).with_vision(true),
1294        ));
1295        let blind = provider("blind", "m", StructuredOutputCapability::NativeJsonSchema);
1296        let router = PolicyRouter::new(pool(vec![vision, blind]));
1297
1298        let request = ModelRequest::new(ModelPurpose::Extract).with_message(
1299            Message::user("guarda").with_part(ContentPart::image_url("https://x.test/a")),
1300        );
1301        let candidates = router
1302            .select_for_request(&request, &RoutingPolicy::new())
1303            .unwrap();
1304        assert_eq!(candidates.len(), 1);
1305        assert_eq!(candidates[0].profile.provider.as_str(), "vision");
1306        assert!(format!("{:?}", candidates[0]).contains("vision/m"));
1307        assert!(format!("{router:?}").contains("PolicyRouter"));
1308    }
1309}