use std::collections::{BTreeMap, BTreeSet};
use std::fmt;
use std::sync::{Arc, Mutex, PoisonError};
use std::time::Duration;
use chrono::{DateTime, Utc};
use turnframe_core::read::DataSensitivity;
use crate::capabilities::{
CapabilityMismatch, CapabilityRequirements, MicroCents, ModelProfile,
StructuredOutputCapability,
};
use crate::error::RetryClass;
use crate::ids::{ModelRef, ProviderKey};
use crate::provider::ModelProvider;
use crate::purpose::ModelPurpose;
pub trait Clock: Send + Sync + fmt::Debug {
fn now(&self) -> DateTime<Utc>;
}
#[derive(Debug, Clone, Copy, Default)]
pub struct SystemClock;
impl Clock for SystemClock {
fn now(&self) -> DateTime<Utc> {
Utc::now()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct HealthPolicy {
pub failure_threshold: u32,
pub cooldown: Duration,
}
impl HealthPolicy {
pub const DEFAULT: Self = Self {
failure_threshold: 3,
cooldown: Duration::from_secs(30),
};
}
impl Default for HealthPolicy {
fn default() -> Self {
Self::DEFAULT
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct ProviderHealth {
pub consecutive_failures: u32,
pub degraded_until: Option<DateTime<Utc>>,
pub successes: u64,
pub failures: u64,
}
impl ProviderHealth {
#[must_use]
pub fn is_healthy_at(&self, now: DateTime<Utc>) -> bool {
self.degraded_until.is_none_or(|until| now >= until)
}
}
#[derive(Clone)]
pub struct ProviderCandidate {
pub provider: Arc<dyn ModelProvider>,
pub profile: ModelProfile,
pub healthy: bool,
}
impl ProviderCandidate {
#[must_use]
pub fn reference(&self) -> ModelRef {
self.profile.reference()
}
}
impl fmt::Debug for ProviderCandidate {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ProviderCandidate")
.field("model", &self.reference().to_string())
.field("healthy", &self.healthy)
.finish_non_exhaustive()
}
}
#[derive(Debug, Clone)]
pub struct RoutingPolicy {
pub allowlist: Option<BTreeSet<ProviderKey>>,
pub denylist: BTreeSet<ProviderKey>,
pub max_cost_per_million: Option<MicroCents>,
pub allowed_regions: Option<BTreeSet<String>>,
pub sensitivity: DataSensitivity,
pub sensitivity_allowlist: BTreeMap<DataSensitivity, BTreeSet<ProviderKey>>,
pub preference: Vec<ModelRef>,
pub skip_degraded: bool,
pub required_tag: Option<String>,
}
impl Default for RoutingPolicy {
fn default() -> Self {
Self::new()
}
}
impl RoutingPolicy {
#[must_use]
pub fn new() -> Self {
Self {
allowlist: None,
denylist: BTreeSet::new(),
max_cost_per_million: None,
allowed_regions: None,
sensitivity: DataSensitivity::Internal,
sensitivity_allowlist: BTreeMap::new(),
preference: Vec::new(),
skip_degraded: true,
required_tag: None,
}
}
#[must_use]
pub fn with_required_tag(mut self, tag: impl Into<String>) -> Self {
self.required_tag = Some(tag.into());
self
}
#[must_use]
pub fn with_allowlist<I: IntoIterator<Item = ProviderKey>>(mut self, providers: I) -> Self {
self.allowlist = Some(providers.into_iter().collect());
self
}
#[must_use]
pub fn with_denylist<I: IntoIterator<Item = ProviderKey>>(mut self, providers: I) -> Self {
self.denylist = providers.into_iter().collect();
self
}
#[must_use]
pub fn with_max_cost(mut self, ceiling: MicroCents) -> Self {
self.max_cost_per_million = Some(ceiling);
self
}
#[must_use]
pub fn with_regions<I: IntoIterator<Item = String>>(mut self, regions: I) -> Self {
self.allowed_regions = Some(regions.into_iter().collect());
self
}
#[must_use]
pub fn with_sensitivity(mut self, sensitivity: DataSensitivity) -> Self {
self.sensitivity = sensitivity;
self
}
#[must_use]
pub fn allowing<I: IntoIterator<Item = ProviderKey>>(
mut self,
level: DataSensitivity,
providers: I,
) -> Self {
self.sensitivity_allowlist
.insert(level, providers.into_iter().collect());
self
}
#[must_use]
pub fn preferring<I: IntoIterator<Item = ModelRef>>(mut self, order: I) -> Self {
self.preference = order.into_iter().collect();
self
}
fn admits(&self, profile: &ModelProfile) -> Result<(), RejectionReason> {
if let Some(tag) = &self.required_tag
&& !profile.tags.contains(tag)
{
return Err(RejectionReason::MissingTag { tag: tag.clone() });
}
if self.denylist.contains(&profile.provider) {
return Err(RejectionReason::Denylisted);
}
if let Some(allowlist) = &self.allowlist
&& !allowlist.contains(&profile.provider)
{
return Err(RejectionReason::NotAllowlisted);
}
match self.sensitivity_allowlist.get(&self.sensitivity) {
Some(allowed) if !allowed.contains(&profile.provider) => {
return Err(RejectionReason::Sensitivity {
level: self.sensitivity,
});
}
None if self.sensitivity >= DataSensitivity::Confidential => {
return Err(RejectionReason::Sensitivity {
level: self.sensitivity,
});
}
_ => {}
}
if let Some(regions) = &self.allowed_regions {
let admitted = profile
.region
.as_ref()
.is_some_and(|region| regions.contains(region));
if !admitted {
return Err(RejectionReason::Region {
declared: profile.region.clone(),
});
}
}
if let Some(ceiling) = self.max_cost_per_million {
let cost = profile.max_cost_per_million();
if !cost.is_some_and(|cost| cost <= ceiling) {
return Err(RejectionReason::CostCeiling {
declared: cost,
ceiling,
});
}
}
Ok(())
}
fn preference_rank(&self, model: &ModelRef) -> usize {
self.preference
.iter()
.position(|preferred| preferred == model)
.unwrap_or(usize::MAX)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum RejectionReason {
Capability(CapabilityMismatch),
MissingTag {
tag: String,
},
NotAllowlisted,
Denylisted,
Region {
declared: Option<String>,
},
CostCeiling {
declared: Option<MicroCents>,
ceiling: MicroCents,
},
Sensitivity {
level: DataSensitivity,
},
Degraded,
}
impl RejectionReason {
#[must_use]
pub const fn as_str(&self) -> &'static str {
match self {
Self::Capability(_) => "capability",
Self::MissingTag { .. } => "missing_tag",
Self::NotAllowlisted => "not_allowlisted",
Self::Denylisted => "denylisted",
Self::Region { .. } => "region",
Self::CostCeiling { .. } => "cost_ceiling",
Self::Sensitivity { .. } => "sensitivity",
Self::Degraded => "degraded",
}
}
}
impl fmt::Display for RejectionReason {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Capability(mismatch) => write!(f, "{mismatch}"),
Self::Region { declared } => match declared {
Some(region) => write!(f, "region({region})"),
None => f.write_str("region(undeclared)"),
},
Self::CostCeiling { declared, ceiling } => match declared {
Some(cost) => write!(f, "cost_ceiling({cost} > {ceiling})"),
None => write!(f, "cost_ceiling(unknown price, ceiling {ceiling})"),
},
Self::Sensitivity { level } => write!(f, "sensitivity({level:?})"),
Self::MissingTag { tag } => write!(f, "missing_tag({tag})"),
other => f.write_str(other.as_str()),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CandidateRejection {
pub model: ModelRef,
pub reason: RejectionReason,
}
impl fmt::Display for CandidateRejection {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}: {}", self.model, self.reason)
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum RoutingError {
#[error("no provider is configured for {purpose}")]
NoProvidersConfigured {
purpose: ModelPurpose,
},
#[error("no provider can serve {purpose}: {}", DisplayRejections(.rejections))]
NoCandidate {
purpose: ModelPurpose,
required_structured_output: Vec<StructuredOutputCapability>,
rejections: Vec<CandidateRejection>,
},
}
impl RoutingError {
#[must_use]
pub fn structured_output_unmet(&self) -> bool {
match self {
Self::NoProvidersConfigured { .. } => false,
Self::NoCandidate { rejections, .. } => rejections.iter().any(|rejection| {
matches!(&rejection.reason, RejectionReason::Capability(mismatch)
if mismatch.structured_output_unmet())
}),
}
}
#[must_use]
pub const fn purpose(&self) -> ModelPurpose {
match self {
Self::NoProvidersConfigured { purpose } | Self::NoCandidate { purpose, .. } => *purpose,
}
}
}
struct DisplayRejections<'a>(&'a [CandidateRejection]);
impl fmt::Display for DisplayRejections<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
for (index, rejection) in self.0.iter().enumerate() {
if index > 0 {
f.write_str("; ")?;
}
write!(f, "{rejection}")?;
}
Ok(())
}
}
pub trait ProviderRouter: Send + Sync {
fn select(
&self,
purpose: ModelPurpose,
requirements: &CapabilityRequirements,
policy: &RoutingPolicy,
) -> Result<Vec<ProviderCandidate>, RoutingError>;
}
struct PoolEntry {
provider: Arc<dyn ModelProvider>,
profile: ModelProfile,
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum PoolError {
#[error("profile {model} is configured twice")]
DuplicateProfile {
model: ModelRef,
},
}
pub struct ProviderPool {
entries: Vec<PoolEntry>,
health: Mutex<BTreeMap<ModelRef, ProviderHealth>>,
clock: Arc<dyn Clock>,
policy: HealthPolicy,
}
impl ProviderPool {
#[must_use]
pub fn builder() -> ProviderPoolBuilder {
ProviderPoolBuilder::new()
}
#[must_use]
pub fn len(&self) -> usize {
self.entries.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
#[must_use]
pub fn profiles(&self) -> Vec<ModelProfile> {
self.entries
.iter()
.map(|entry| entry.profile.clone())
.collect()
}
#[must_use]
pub fn provider(&self, model: &ModelRef) -> Option<Arc<dyn ModelProvider>> {
self.entries
.iter()
.find(|entry| entry.profile.reference() == *model)
.map(|entry| Arc::clone(&entry.provider))
}
#[must_use]
pub fn now(&self) -> DateTime<Utc> {
self.clock.now()
}
#[must_use]
pub fn health(&self, model: &ModelRef) -> ProviderHealth {
self.health
.lock()
.unwrap_or_else(PoisonError::into_inner)
.get(model)
.cloned()
.unwrap_or_default()
}
#[must_use]
pub fn is_healthy(&self, model: &ModelRef) -> bool {
self.health(model).is_healthy_at(self.clock.now())
}
pub fn record_success(&self, model: &ModelRef) {
let mut health = self.health.lock().unwrap_or_else(PoisonError::into_inner);
let entry = health.entry(model.clone()).or_default();
entry.consecutive_failures = 0;
entry.degraded_until = None;
entry.successes = entry.successes.saturating_add(1);
}
pub fn record_failure(&self, model: &ModelRef, class: RetryClass) {
if class == RetryClass::Fatal {
return;
}
let now = self.clock.now();
let mut health = self.health.lock().unwrap_or_else(PoisonError::into_inner);
let entry = health.entry(model.clone()).or_default();
entry.consecutive_failures = entry.consecutive_failures.saturating_add(1);
entry.failures = entry.failures.saturating_add(1);
if entry.consecutive_failures >= self.policy.failure_threshold {
entry.degraded_until = Some(
now + chrono::Duration::from_std(self.policy.cooldown)
.unwrap_or_else(|_| chrono::Duration::seconds(30)),
);
}
}
pub fn reset_health(&self) {
self.health
.lock()
.unwrap_or_else(PoisonError::into_inner)
.clear();
}
}
impl fmt::Debug for ProviderPool {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let models: Vec<String> = self
.entries
.iter()
.map(|entry| entry.profile.reference().to_string())
.collect();
f.debug_struct("ProviderPool")
.field("profiles", &models)
.field("health_policy", &self.policy)
.finish_non_exhaustive()
}
}
pub struct ProviderPoolBuilder {
entries: Vec<PoolEntry>,
clock: Arc<dyn Clock>,
policy: HealthPolicy,
}
impl ProviderPoolBuilder {
#[must_use]
pub fn new() -> Self {
Self {
entries: Vec::new(),
clock: Arc::new(SystemClock),
policy: HealthPolicy::DEFAULT,
}
}
#[must_use]
pub fn provider(self, provider: Arc<dyn ModelProvider>) -> Self {
let profile = provider.profile();
self.provider_with_profile(provider, profile)
}
#[must_use]
pub fn provider_tagged<I, T>(self, provider: Arc<dyn ModelProvider>, tags: I) -> Self
where
I: IntoIterator<Item = T>,
T: Into<String>,
{
let mut profile = provider.profile();
for tag in tags {
let tag = tag.into();
if !profile.tags.contains(&tag) {
profile.tags.push(tag);
}
}
self.provider_with_profile(provider, profile)
}
#[must_use]
pub fn provider_with_profile(
mut self,
provider: Arc<dyn ModelProvider>,
profile: ModelProfile,
) -> Self {
self.entries.push(PoolEntry { provider, profile });
self
}
#[must_use]
pub fn clock<C: Clock + 'static>(mut self, clock: Arc<C>) -> Self {
self.clock = clock;
self
}
#[must_use]
pub fn health_policy(mut self, policy: HealthPolicy) -> Self {
self.policy = policy;
self
}
pub fn build(self) -> Result<ProviderPool, PoolError> {
let mut seen = BTreeSet::new();
for entry in &self.entries {
let reference = entry.profile.reference();
if !seen.insert(reference.clone()) {
return Err(PoolError::DuplicateProfile { model: reference });
}
}
Ok(ProviderPool {
entries: self.entries,
health: Mutex::new(BTreeMap::new()),
clock: self.clock,
policy: self.policy,
})
}
}
impl Default for ProviderPoolBuilder {
fn default() -> Self {
Self::new()
}
}
impl fmt::Debug for ProviderPoolBuilder {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ProviderPoolBuilder")
.field("entries", &self.entries.len())
.field("health_policy", &self.policy)
.finish_non_exhaustive()
}
}
pub struct PolicyRouter {
pool: Arc<ProviderPool>,
}
impl PolicyRouter {
#[must_use]
pub fn new(pool: Arc<ProviderPool>) -> Self {
Self { pool }
}
#[must_use]
pub fn pool(&self) -> &Arc<ProviderPool> {
&self.pool
}
pub fn select_for_request(
&self,
request: &crate::request::ModelRequest,
policy: &RoutingPolicy,
) -> Result<Vec<ProviderCandidate>, RoutingError> {
let requirements = request.requirements();
self.select(request.purpose, &requirements, policy)
}
}
impl ProviderRouter for PolicyRouter {
fn select(
&self,
purpose: ModelPurpose,
requirements: &CapabilityRequirements,
policy: &RoutingPolicy,
) -> Result<Vec<ProviderCandidate>, RoutingError> {
if self.pool.is_empty() {
return Err(RoutingError::NoProvidersConfigured { purpose });
}
let now = self.pool.now();
let mut rejections = Vec::new();
let mut admitted: Vec<ProviderCandidate> = Vec::new();
for entry in &self.pool.entries {
let model = entry.profile.reference();
if let Err(mismatch) = requirements.satisfied_by(&entry.profile.capabilities) {
rejections.push(CandidateRejection {
model,
reason: RejectionReason::Capability(mismatch),
});
continue;
}
if let Err(reason) = policy.admits(&entry.profile) {
rejections.push(CandidateRejection { model, reason });
continue;
}
let healthy = self.pool.health(&model).is_healthy_at(now);
admitted.push(ProviderCandidate {
provider: Arc::clone(&entry.provider),
profile: entry.profile.clone(),
healthy,
});
}
if policy.skip_degraded && admitted.iter().any(|candidate| candidate.healthy) {
admitted.retain(|candidate| {
if candidate.healthy {
return true;
}
rejections.push(CandidateRejection {
model: candidate.reference(),
reason: RejectionReason::Degraded,
});
false
});
}
if admitted.is_empty() {
return Err(RoutingError::NoCandidate {
purpose,
required_structured_output: requirements.structured_output.clone(),
rejections,
});
}
admitted.sort_by_key(|candidate| {
(
usize::from(!candidate.healthy),
policy.preference_rank(&candidate.reference()),
)
});
Ok(admitted)
}
}
impl fmt::Debug for PolicyRouter {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PolicyRouter")
.field("pool", &self.pool)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::capabilities::ProviderCapabilities;
use crate::testing::{ManualClock, StaticProvider};
fn caps(structured: StructuredOutputCapability) -> ProviderCapabilities {
ProviderCapabilities::minimal().with_structured_output(structured)
}
fn provider(
provider_key: &str,
model: &str,
structured: StructuredOutputCapability,
) -> Arc<dyn ModelProvider> {
Arc::new(StaticProvider::new(provider_key, model).with_capabilities(caps(structured)))
}
fn pool(entries: Vec<Arc<dyn ModelProvider>>) -> Arc<ProviderPool> {
let mut builder = ProviderPool::builder();
for entry in entries {
builder = builder.provider(entry);
}
Arc::new(builder.build().unwrap())
}
fn understand() -> CapabilityRequirements {
ModelPurpose::Extract.requirements()
}
#[test]
fn a_required_tag_admits_only_the_profiles_carrying_it() {
let pool = Arc::new(
ProviderPool::builder()
.provider_tagged(
provider("mini", "m", StructuredOutputCapability::NativeJsonSchema),
["small"],
)
.provider_tagged(
provider("big", "m", StructuredOutputCapability::NativeJsonSchema),
["large"],
)
.build()
.unwrap(),
);
let router = PolicyRouter::new(pool);
let large = router
.select(
ModelPurpose::Extract,
&ModelPurpose::Extract.requirements(),
&RoutingPolicy::new().with_required_tag("large"),
)
.unwrap();
assert_eq!(large.len(), 1);
assert_eq!(large[0].profile.provider.as_str(), "big");
let error = router
.select(
ModelPurpose::Extract,
&ModelPurpose::Extract.requirements(),
&RoutingPolicy::new().with_required_tag("vision"),
)
.expect_err("no profile carries the tag, and nothing untagged stands in");
assert!(error.to_string().contains("missing_tag(vision)"), "{error}");
}
#[test]
fn capability_fit_runs_before_policy_and_is_never_relaxed() {
let router = PolicyRouter::new(pool(vec![
provider("weak", "m", StructuredOutputCapability::PromptOnly),
provider("strong", "m", StructuredOutputCapability::NativeJsonSchema),
]));
let candidates = router
.select(ModelPurpose::Extract, &understand(), &RoutingPolicy::new())
.unwrap();
assert_eq!(candidates.len(), 1);
assert_eq!(candidates[0].profile.provider.as_str(), "strong");
}
#[test]
fn no_silent_downgrade_names_what_was_missing() {
let router = PolicyRouter::new(pool(vec![
provider("weak", "m", StructuredOutputCapability::PromptOnly),
provider("weaker", "m", StructuredOutputCapability::None),
]));
let error = router
.select(ModelPurpose::Extract, &understand(), &RoutingPolicy::new())
.unwrap_err();
assert!(error.structured_output_unmet());
assert_eq!(error.purpose(), ModelPurpose::Extract);
let RoutingError::NoCandidate {
required_structured_output,
rejections,
..
} = &error
else {
panic!("{error:?}");
};
assert_eq!(
required_structured_output,
&crate::purpose::MUTATION_SAFE_STRUCTURED_OUTPUT.to_vec()
);
assert_eq!(rejections.len(), 2);
let text = error.to_string();
assert!(text.contains("native_json_schema"), "{text}");
assert!(text.contains("prompt_only"), "{text}");
assert!(text.contains("weaker/m"), "{text}");
}
#[test]
fn an_empty_pool_is_its_own_error() {
let router = PolicyRouter::new(pool(vec![]));
let error = router
.select(
ModelPurpose::Acknowledge,
&CapabilityRequirements::none(),
&RoutingPolicy::new(),
)
.unwrap_err();
assert!(matches!(error, RoutingError::NoProvidersConfigured { .. }));
assert!(!error.structured_output_unmet());
}
#[test]
fn allowlist_denylist_and_preference_shape_the_order() {
let router = PolicyRouter::new(pool(vec![
provider("a", "m", StructuredOutputCapability::NativeJsonSchema),
provider("b", "m", StructuredOutputCapability::NativeJsonSchema),
provider("c", "m", StructuredOutputCapability::NativeJsonSchema),
]));
let denied = RoutingPolicy::new().with_denylist([ProviderKey::from("a")]);
let candidates = router
.select(ModelPurpose::Extract, &understand(), &denied)
.unwrap();
assert_eq!(candidates.len(), 2);
assert!(
candidates
.iter()
.all(|c| c.profile.provider.as_str() != "a")
);
let allowed = RoutingPolicy::new().with_allowlist([ProviderKey::from("c")]);
let candidates = router
.select(ModelPurpose::Extract, &understand(), &allowed)
.unwrap();
assert_eq!(candidates.len(), 1);
assert_eq!(candidates[0].profile.provider.as_str(), "c");
let preferred = RoutingPolicy::new().preferring([ModelRef::new("c", "m")]);
let candidates = router
.select(ModelPurpose::Extract, &understand(), &preferred)
.unwrap();
let order: Vec<&str> = candidates
.iter()
.map(|c| c.profile.provider.as_str())
.collect();
assert_eq!(
order,
vec!["c", "a", "b"],
"preferred first, then pool order"
);
let both = RoutingPolicy::new()
.with_allowlist([ProviderKey::from("a")])
.with_denylist([ProviderKey::from("a")]);
let error = router
.select(ModelPurpose::Extract, &understand(), &both)
.unwrap_err();
assert!(error.to_string().contains("denylisted"), "{error}");
}
#[test]
fn residency_and_cost_fail_closed_on_undeclared_profiles() {
let eu = Arc::new(
StaticProvider::new("eu", "m")
.with_capabilities(caps(StructuredOutputCapability::NativeJsonSchema))
.with_profile_region("eu")
.with_profile_cost(MicroCents::from_cents(10), MicroCents::from_cents(20)),
);
let unknown = provider("unknown", "m", StructuredOutputCapability::NativeJsonSchema);
let router = PolicyRouter::new(pool(vec![eu, unknown]));
let residency = RoutingPolicy::new().with_regions(["eu".to_owned()]);
let candidates = router
.select(ModelPurpose::Extract, &understand(), &residency)
.unwrap();
assert_eq!(
candidates.len(),
1,
"a profile with no region does not pass"
);
assert_eq!(candidates[0].profile.provider.as_str(), "eu");
let ceiling = RoutingPolicy::new().with_max_cost(MicroCents::from_cents(20));
let candidates = router
.select(ModelPurpose::Extract, &understand(), &ceiling)
.unwrap();
assert_eq!(
candidates.len(),
1,
"an unknown price does not pass a ceiling"
);
let too_low = RoutingPolicy::new().with_max_cost(MicroCents::from_cents(5));
let error = router
.select(ModelPurpose::Extract, &understand(), &too_low)
.unwrap_err();
assert!(error.to_string().contains("cost_ceiling"), "{error}");
}
#[test]
fn confidential_data_needs_an_explicit_provider_allowlist() {
let router = PolicyRouter::new(pool(vec![provider(
"openai",
"m",
StructuredOutputCapability::NativeJsonSchema,
)]));
let undeclared = RoutingPolicy::new().with_sensitivity(DataSensitivity::Confidential);
let error = router
.select(ModelPurpose::Extract, &understand(), &undeclared)
.unwrap_err();
assert!(error.to_string().contains("sensitivity"), "{error}");
let declared = RoutingPolicy::new()
.with_sensitivity(DataSensitivity::Confidential)
.allowing(DataSensitivity::Confidential, [ProviderKey::from("openai")]);
assert_eq!(
router
.select(ModelPurpose::Extract, &understand(), &declared)
.unwrap()
.len(),
1
);
let internal = RoutingPolicy::new().with_sensitivity(DataSensitivity::Internal);
assert_eq!(
router
.select(ModelPurpose::Extract, &understand(), &internal)
.unwrap()
.len(),
1
);
}
#[test]
fn health_degrades_on_a_clock_we_control_and_recovers() {
let clock = Arc::new(ManualClock::at_epoch());
let pool = Arc::new(
ProviderPool::builder()
.provider(provider(
"a",
"m",
StructuredOutputCapability::NativeJsonSchema,
))
.provider(provider(
"b",
"m",
StructuredOutputCapability::NativeJsonSchema,
))
.clock(Arc::clone(&clock))
.health_policy(HealthPolicy {
failure_threshold: 2,
cooldown: Duration::from_secs(60),
})
.build()
.unwrap(),
);
let router = PolicyRouter::new(Arc::clone(&pool));
let a = ModelRef::new("a", "m");
pool.record_failure(&a, RetryClass::Retry);
assert!(pool.is_healthy(&a), "one failure is not a cooldown");
pool.record_failure(&a, RetryClass::Retry);
assert!(!pool.is_healthy(&a));
assert_eq!(pool.health(&a).consecutive_failures, 2);
let candidates = router
.select(ModelPurpose::Extract, &understand(), &RoutingPolicy::new())
.unwrap();
assert_eq!(candidates.len(), 1);
assert_eq!(candidates[0].profile.provider.as_str(), "b");
clock.advance(Duration::from_secs(61));
assert!(pool.is_healthy(&a));
assert_eq!(
router
.select(ModelPurpose::Extract, &understand(), &RoutingPolicy::new())
.unwrap()
.len(),
2
);
pool.record_failure(&a, RetryClass::Fallback);
pool.record_failure(&a, RetryClass::Fallback);
assert!(!pool.is_healthy(&a));
pool.record_success(&a);
assert!(pool.is_healthy(&a));
assert_eq!(pool.health(&a).consecutive_failures, 0);
assert_eq!(pool.health(&a).successes, 1);
}
#[test]
fn a_fatal_outcome_does_not_degrade_health() {
let clock = Arc::new(ManualClock::at_epoch());
let pool = ProviderPool::builder()
.provider(provider(
"a",
"m",
StructuredOutputCapability::NativeJsonSchema,
))
.clock(clock)
.health_policy(HealthPolicy {
failure_threshold: 1,
cooldown: Duration::from_secs(60),
})
.build()
.unwrap();
let a = ModelRef::new("a", "m");
pool.record_failure(&a, RetryClass::Fatal);
assert!(pool.is_healthy(&a));
assert_eq!(pool.health(&a).failures, 0);
pool.record_failure(&a, RetryClass::Retry);
assert!(!pool.is_healthy(&a));
pool.reset_health();
assert!(pool.is_healthy(&a));
}
#[test]
fn a_degraded_profile_is_still_offered_when_it_is_the_only_fit() {
let clock = Arc::new(ManualClock::at_epoch());
let pool = Arc::new(
ProviderPool::builder()
.provider(provider(
"a",
"m",
StructuredOutputCapability::NativeJsonSchema,
))
.clock(clock)
.health_policy(HealthPolicy {
failure_threshold: 1,
cooldown: Duration::from_secs(60),
})
.build()
.unwrap(),
);
let a = ModelRef::new("a", "m");
pool.record_failure(&a, RetryClass::Retry);
assert!(!pool.is_healthy(&a));
let router = PolicyRouter::new(pool);
let candidates = router
.select(ModelPurpose::Extract, &understand(), &RoutingPolicy::new())
.unwrap();
assert_eq!(candidates.len(), 1);
assert!(!candidates[0].healthy, "offered, and honestly labelled");
}
#[test]
fn a_duplicate_profile_is_a_configuration_error() {
let error = ProviderPool::builder()
.provider(provider(
"a",
"m",
StructuredOutputCapability::NativeJsonSchema,
))
.provider(provider("a", "m", StructuredOutputCapability::JsonObject))
.build()
.unwrap_err();
assert_eq!(
error,
PoolError::DuplicateProfile {
model: ModelRef::new("a", "m")
}
);
}
#[test]
fn the_pool_answers_lookups_and_renders_safely() {
let pool = pool(vec![provider(
"a",
"m",
StructuredOutputCapability::NativeJsonSchema,
)]);
assert_eq!(pool.len(), 1);
assert!(!pool.is_empty());
assert_eq!(pool.profiles().len(), 1);
assert!(pool.provider(&ModelRef::new("a", "m")).is_some());
assert!(pool.provider(&ModelRef::new("a", "other")).is_none());
let rendered = format!("{pool:?}");
assert!(rendered.contains("a/m"), "{rendered}");
}
#[test]
fn select_for_request_folds_in_the_requests_own_needs() {
use crate::request::{ContentPart, Message, ModelRequest};
let vision = Arc::new(StaticProvider::new("vision", "m").with_capabilities(
caps(StructuredOutputCapability::NativeJsonSchema).with_vision(true),
));
let blind = provider("blind", "m", StructuredOutputCapability::NativeJsonSchema);
let router = PolicyRouter::new(pool(vec![vision, blind]));
let request = ModelRequest::new(ModelPurpose::Extract).with_message(
Message::user("guarda").with_part(ContentPart::image_url("https://x.test/a")),
);
let candidates = router
.select_for_request(&request, &RoutingPolicy::new())
.unwrap();
assert_eq!(candidates.len(), 1);
assert_eq!(candidates[0].profile.provider.as_str(), "vision");
assert!(format!("{:?}", candidates[0]).contains("vision/m"));
assert!(format!("{router:?}").contains("PolicyRouter"));
}
}