1use 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
53pub trait Clock: Send + Sync + fmt::Debug {
55 fn now(&self) -> DateTime<Utc>;
57}
58
59#[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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
71pub struct HealthPolicy {
72 pub failure_threshold: u32,
74 pub cooldown: Duration,
76}
77
78impl HealthPolicy {
79 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#[derive(Debug, Clone, PartialEq, Eq, Default)]
94pub struct ProviderHealth {
95 pub consecutive_failures: u32,
97 pub degraded_until: Option<DateTime<Utc>>,
99 pub successes: u64,
101 pub failures: u64,
103}
104
105impl ProviderHealth {
106 #[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#[derive(Clone)]
115pub struct ProviderCandidate {
116 pub provider: Arc<dyn ModelProvider>,
118 pub profile: ModelProfile,
120 pub healthy: bool,
122}
123
124impl ProviderCandidate {
125 #[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#[derive(Debug, Clone)]
149pub struct RoutingPolicy {
150 pub allowlist: Option<BTreeSet<ProviderKey>>,
152 pub denylist: BTreeSet<ProviderKey>,
154 pub max_cost_per_million: Option<MicroCents>,
157 pub allowed_regions: Option<BTreeSet<String>>,
160 pub sensitivity: DataSensitivity,
162 pub sensitivity_allowlist: BTreeMap<DataSensitivity, BTreeSet<ProviderKey>>,
170 pub preference: Vec<ModelRef>,
173 pub skip_degraded: bool,
176 pub required_tag: Option<String>,
179}
180
181impl Default for RoutingPolicy {
182 fn default() -> Self {
184 Self::new()
185 }
186}
187
188impl RoutingPolicy {
189 #[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 #[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 #[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 #[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 #[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 #[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 #[must_use]
244 pub fn with_sensitivity(mut self, sensitivity: DataSensitivity) -> Self {
245 self.sensitivity = sensitivity;
246 self
247 }
248
249 #[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 #[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 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 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#[derive(Debug, Clone, PartialEq, Eq)]
330#[non_exhaustive]
331pub enum RejectionReason {
332 Capability(CapabilityMismatch),
334 MissingTag {
336 tag: String,
338 },
339 NotAllowlisted,
341 Denylisted,
343 Region {
345 declared: Option<String>,
347 },
348 CostCeiling {
350 declared: Option<MicroCents>,
352 ceiling: MicroCents,
354 },
355 Sensitivity {
357 level: DataSensitivity,
359 },
360 Degraded,
362}
363
364impl RejectionReason {
365 #[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#[derive(Debug, Clone, PartialEq, Eq)]
402pub struct CandidateRejection {
403 pub model: ModelRef,
405 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#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
417#[non_exhaustive]
418pub enum RoutingError {
419 #[error("no provider is configured for {purpose}")]
421 NoProvidersConfigured {
422 purpose: ModelPurpose,
424 },
425 #[error("no provider can serve {purpose}: {}", DisplayRejections(.rejections))]
427 NoCandidate {
428 purpose: ModelPurpose,
430 required_structured_output: Vec<StructuredOutputCapability>,
434 rejections: Vec<CandidateRejection>,
436 },
437}
438
439impl RoutingError {
440 #[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 #[must_use]
455 pub const fn purpose(&self) -> ModelPurpose {
456 match self {
457 Self::NoProvidersConfigured { purpose } | Self::NoCandidate { purpose, .. } => *purpose,
458 }
459 }
460}
461
462struct 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
477pub trait ProviderRouter: Send + Sync {
479 fn select(
486 &self,
487 purpose: ModelPurpose,
488 requirements: &CapabilityRequirements,
489 policy: &RoutingPolicy,
490 ) -> Result<Vec<ProviderCandidate>, RoutingError>;
491}
492
493struct PoolEntry {
495 provider: Arc<dyn ModelProvider>,
496 profile: ModelProfile,
497}
498
499#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
501#[non_exhaustive]
502pub enum PoolError {
503 #[error("profile {model} is configured twice")]
505 DuplicateProfile {
506 model: ModelRef,
508 },
509}
510
511pub 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 #[must_use]
525 pub fn builder() -> ProviderPoolBuilder {
526 ProviderPoolBuilder::new()
527 }
528
529 #[must_use]
531 pub fn len(&self) -> usize {
532 self.entries.len()
533 }
534
535 #[must_use]
537 pub fn is_empty(&self) -> bool {
538 self.entries.is_empty()
539 }
540
541 #[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 #[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 #[must_use]
561 pub fn now(&self) -> DateTime<Utc> {
562 self.clock.now()
563 }
564
565 #[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 #[must_use]
578 pub fn is_healthy(&self, model: &ModelRef) -> bool {
579 self.health(model).is_healthy_at(self.clock.now())
580 }
581
582 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 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 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
638pub struct ProviderPoolBuilder {
640 entries: Vec<PoolEntry>,
641 clock: Arc<dyn Clock>,
642 policy: HealthPolicy,
643}
644
645impl ProviderPoolBuilder {
646 #[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 #[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 #[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 #[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 #[must_use]
700 pub fn clock<C: Clock + 'static>(mut self, clock: Arc<C>) -> Self {
701 self.clock = clock;
702 self
703 }
704
705 #[must_use]
707 pub fn health_policy(mut self, policy: HealthPolicy) -> Self {
708 self.policy = policy;
709 self
710 }
711
712 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
751pub struct PolicyRouter {
782 pool: Arc<ProviderPool>,
783}
784
785impl PolicyRouter {
786 #[must_use]
788 pub fn new(pool: Arc<ProviderPool>) -> Self {
789 Self { pool }
790 }
791
792 #[must_use]
794 pub fn pool(&self) -> &Arc<ProviderPool> {
795 &self.pool
796 }
797
798 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 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 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 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 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 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 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 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 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 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}