feat: unify provider scheduling workspace

This commit is contained in:
elky
2026-10-07 00:34:18 +08:00
parent e7de935e61
commit 466c7918a1
69 changed files with 5810 additions and 2666 deletions
@@ -1179,6 +1179,16 @@ WHERE id = $1
&self,
provider: &StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
) -> Result<StoredProviderCatalogProvider, DataLayerError> {
self.create_provider_with_routing_group(provider, shift_existing_priorities_from, None)
.await
}
async fn create_provider_with_routing_group(
&self,
provider: &StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
routing_group_id: Option<&str>,
) -> Result<StoredProviderCatalogProvider, DataLayerError> {
if provider.id.trim().is_empty() {
return Err(DataLayerError::InvalidInput(
@@ -1208,6 +1218,40 @@ WHERE id = $1
let mut tx = self.pool.begin().await.map_postgres_err()?;
if let Some(group_id) = routing_group_id {
// Group edits use the same lock: validation, exclusions, and provider
// creation are committed together, including concurrent deletions.
sqlx::query("LOCK TABLE routing_groups IN SHARE ROW EXCLUSIVE MODE")
.execute(&mut *tx)
.await
.map_postgres_err()?;
let exists: bool =
sqlx::query_scalar("SELECT EXISTS(SELECT 1 FROM routing_groups WHERE id = $1)")
.bind(group_id)
.fetch_one(&mut *tx)
.await
.map_postgres_err()?;
if !exists {
return Err(DataLayerError::InvalidInput(
"routing_group_not_found".to_string(),
));
}
sqlx::query(r#"
UPDATE routing_groups
SET config_json = jsonb_set(config_json::jsonb, '{disabled_providers}',
CASE WHEN id = $1
THEN COALESCE(config_json::jsonb -> 'disabled_providers', '[]'::jsonb) - $2::text
ELSE (COALESCE(config_json::jsonb -> 'disabled_providers', '[]'::jsonb) - $2::text) || jsonb_build_array($2::text)
END),
version = version + 1,
updated_at = EXTRACT(EPOCH FROM NOW())::bigint
WHERE (id <> $1 AND NOT (COALESCE(config_json::jsonb -> 'disabled_providers', '[]'::jsonb) ? $2::text))
OR (id = $1 AND (COALESCE(config_json::jsonb -> 'disabled_providers', '[]'::jsonb) ? $2::text))
"#)
.bind(group_id).bind(&provider.id)
.execute(&mut *tx).await.map_postgres_err()?;
}
if let Some(target_priority) = shift_existing_priorities_from {
sqlx::query(
r#"
@@ -3015,6 +3059,20 @@ impl ProviderCatalogReadRepository for SqlxProviderCatalogReadRepository {
#[async_trait]
impl ProviderCatalogWriteRepository for SqlxProviderCatalogReadRepository {
async fn create_provider_in_routing_group(
&self,
provider: &StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
routing_group_id: &str,
) -> Result<StoredProviderCatalogProvider, DataLayerError> {
self.create_provider_with_routing_group(
provider,
shift_existing_priorities_from,
Some(routing_group_id),
)
.await
}
async fn create_provider(
&self,
provider: &StoredProviderCatalogProvider,
@@ -1463,6 +1463,19 @@ pub trait ProviderCatalogReadRepository: Send + Sync {
#[async_trait]
pub trait ProviderCatalogWriteRepository: Send + Sync {
/// Create a provider and exclude it from every other existing routing group
/// in one transaction. Implementations must fail closed if unsupported.
async fn create_provider_in_routing_group(
&self,
_provider: &StoredProviderCatalogProvider,
_shift_existing_priorities_from: Option<i32>,
_routing_group_id: &str,
) -> Result<StoredProviderCatalogProvider, crate::DataLayerError> {
Err(crate::DataLayerError::InvalidConfiguration(
"atomic provider creation in a routing group is not supported".to_string(),
))
}
async fn create_provider(
&self,
provider: &StoredProviderCatalogProvider,
@@ -58,6 +58,7 @@ pub struct CreateRoutingGroupRecord {
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct UpdateRoutingGroupRecord {
pub expected_version: Option<i64>,
pub name: Option<String>,
pub description: Option<Option<String>>,
pub enabled: Option<bool>,
@@ -244,6 +245,16 @@ pub fn apply_group_patch(
group: &mut StoredRoutingGroup,
patch: UpdateRoutingGroupRecord,
) -> Result<(), crate::DataLayerError> {
if patch
.expected_version
.is_some_and(|version| version != group.version)
{
return Err(crate::DataLayerError::InvalidInput(
"routing_group_version_conflict".to_string(),
));
}
let previous_version = group.version;
let config_changed = patch.config_json.is_some();
if let Some(name) = patch.name {
if name.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
@@ -273,7 +284,13 @@ pub fn apply_group_patch(
group.config_json = config_json;
}
if let Some(version) = patch.version {
group.version = version.max(1);
group.version = if config_changed {
version.max(previous_version.saturating_add(1))
} else {
version.max(previous_version)
};
} else if config_changed {
group.version = previous_version.saturating_add(1);
}
if let Some(published_at) = patch.published_at {
group.published_at = published_at;
@@ -35,6 +35,7 @@ mod overview_fact_metadata;
mod overview_migration_safety;
mod policy_nulls;
mod provider_expenses;
mod scoped_provider_creation;
/// A clean PostgreSQL database is bootstrapped from the schema snapshot first;
/// migrations after the privacy/security frontier are intentionally left
@@ -0,0 +1,131 @@
use super::*;
use aether_data_contracts::repository::{
provider_catalog::{ProviderCatalogWriteRepository, StoredProviderCatalogProvider},
routing_profiles::{
CreateRoutingGroupRecord, RoutingGroupReadRepository, RoutingGroupWriteRepository,
UpdateRoutingGroupRecord,
},
};
use serde_json::json;
#[tokio::test]
async fn postgres_scoped_provider_creation_rolls_back_and_serializes_group_saves() {
let Some(server) = ManagedPostgresServer::try_start()
.await
.expect("local postgres should start or skip")
else {
return;
};
let pool = PgPool::connect(server.database_url()).await.unwrap();
prepare_and_apply_clean_postgres_database(&pool).await;
let groups =
crate::repository::routing_profiles::PostgresRoutingGroupRepository::new(pool.clone());
let providers =
crate::repository::provider_catalog::SqlxProviderCatalogReadRepository::new(pool.clone());
for id in ["selected", "other"] {
groups
.create_routing_group(CreateRoutingGroupRecord {
id: id.into(),
name: id.into(),
description: None,
enabled: true,
is_system_default: id == "selected",
sort_order: 0,
config_json: json!({"disabled_providers": ["existing-disabled"]}),
version: 1,
created_at: 1,
updated_at: 1,
published_at: None,
})
.await
.unwrap();
}
let provider =
StoredProviderCatalogProvider::new("new".into(), "new".into(), None, "custom".into())
.unwrap();
assert!(providers
.create_provider_in_routing_group(&provider, None, "missing")
.await
.is_err());
assert_eq!(
sqlx::query_scalar::<_, i64>("SELECT COUNT(*) FROM providers")
.fetch_one(&pool)
.await
.unwrap(),
0
);
assert!(groups
.list_routing_groups()
.await
.unwrap()
.iter()
.all(|group| group.version == 1));
providers
.create_provider_in_routing_group(&provider, None, "selected")
.await
.unwrap();
let before = groups.list_routing_groups().await.unwrap();
for group in &before {
assert_eq!(group.version, if group.id == "selected" { 1 } else { 2 });
assert_eq!(
group.config_json["disabled_providers"],
if group.id == "selected" {
json!(["existing-disabled"])
} else {
json!(["existing-disabled", "new"])
}
);
}
// The INSERT fails after group updates execute. Its transaction must undo
// every exclusion and version change along with any priority shifts.
assert!(providers
.create_provider_in_routing_group(&provider, Some(0), "other")
.await
.is_err());
assert_eq!(groups.list_routing_groups().await.unwrap(), before);
assert!(groups
.update_routing_group(
"other",
UpdateRoutingGroupRecord {
expected_version: Some(1),
config_json: Some(json!({"disabled_providers": []})),
..Default::default()
}
)
.await
.is_err());
assert_eq!(groups.list_routing_groups().await.unwrap(), before);
let concurrent_provider = StoredProviderCatalogProvider::new(
"concurrent".into(),
"concurrent".into(),
None,
"custom".into(),
)
.unwrap();
let (created, edited) = tokio::join!(
providers.create_provider_in_routing_group(&concurrent_provider, None, "selected"),
groups.update_routing_group(
"other",
UpdateRoutingGroupRecord {
expected_version: Some(2),
config_json: Some(json!({"disabled_providers": ["new"]})),
..Default::default()
}
)
);
created.unwrap();
let other = groups
.list_routing_groups()
.await
.unwrap()
.into_iter()
.find(|group| group.id == "other")
.unwrap();
assert!(other.config_json["disabled_providers"]
.as_array()
.unwrap()
.contains(&json!("concurrent")));
assert_eq!(other.version, if edited.is_ok() { 4 } else { 3 });
}
@@ -1,5 +1,5 @@
use std::collections::BTreeMap;
use std::sync::RwLock;
use std::sync::{Arc, RwLock};
use std::time::{SystemTime, UNIX_EPOCH};
use async_trait::async_trait;
@@ -31,6 +31,8 @@ struct MemoryProviderCatalogIndex {
#[derive(Debug, Default)]
pub struct InMemoryProviderCatalogReadRepository {
index: RwLock<MemoryProviderCatalogIndex>,
routing_groups:
Option<Arc<crate::repository::routing_profiles::InMemoryRoutingGroupRepository>>,
}
impl InMemoryProviderCatalogReadRepository {
@@ -40,6 +42,7 @@ impl InMemoryProviderCatalogReadRepository {
keys: Vec<StoredProviderCatalogKey>,
) -> Self {
Self {
routing_groups: None,
index: RwLock::new(MemoryProviderCatalogIndex {
providers: providers
.into_iter()
@@ -54,6 +57,14 @@ impl InMemoryProviderCatalogReadRepository {
}
}
pub fn with_routing_groups(
mut self,
repository: Arc<crate::repository::routing_profiles::InMemoryRoutingGroupRepository>,
) -> Self {
self.routing_groups = Some(repository);
self
}
fn snapshot(&self) -> ProviderCatalogSnapshot {
let index = self.index.read().expect("provider catalog repository lock");
ProviderCatalogSnapshot::new(
@@ -416,6 +427,46 @@ impl ProviderCatalogReadRepository for InMemoryProviderCatalogReadRepository {
#[async_trait]
impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
async fn create_provider_in_routing_group(
&self,
provider: &StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
routing_group_id: &str,
) -> Result<StoredProviderCatalogProvider, DataLayerError> {
let groups = self.routing_groups.as_ref().ok_or_else(|| {
DataLayerError::InvalidConfiguration(
"atomic provider creation requires a shared routing group repository".to_string(),
)
})?;
groups.create_scoped_provider(routing_group_id, &provider.id, || {
let mut index = self
.index
.write()
.expect("provider catalog repository lock");
if index.providers.contains_key(&provider.id)
|| index
.providers
.values()
.any(|existing| existing.name == provider.name)
{
return Err(DataLayerError::InvalidInput(
"provider already exists".to_string(),
));
}
if let Some(target_priority) = shift_existing_priorities_from {
for existing in index.providers.values_mut() {
if existing.provider_priority >= target_priority {
existing.provider_priority += 1;
}
}
}
index
.providers
.insert(provider.id.clone(), provider.clone());
Ok(provider.clone())
})
}
async fn create_provider(
&self,
provider: &StoredProviderCatalogProvider,
@@ -1587,6 +1638,102 @@ mod tests {
.expect("key should build")
}
#[tokio::test]
async fn scoped_provider_creation_is_atomic_and_rejects_stale_group_updates() {
use crate::repository::routing_profiles::InMemoryRoutingGroupRepository;
use aether_data_contracts::repository::routing_profiles::{
CreateRoutingGroupRecord, RoutingGroupReadRepository, RoutingGroupWriteRepository,
UpdateRoutingGroupRecord,
};
let groups = Arc::new(InMemoryRoutingGroupRepository::default());
for id in ["selected", "other", "disabled"] {
groups
.create_routing_group(CreateRoutingGroupRecord {
id: id.into(),
name: id.into(),
description: None,
enabled: id != "disabled",
is_system_default: id == "selected",
sort_order: 0,
config_json: json!({"disabled_providers": ["already-disabled"], "rules": []}),
version: 1,
created_at: 1,
updated_at: 1,
published_at: None,
})
.await
.unwrap();
}
let repository =
InMemoryProviderCatalogReadRepository::default().with_routing_groups(groups.clone());
let provider = sample_provider("new");
assert!(repository
.create_provider_in_routing_group(&provider, Some(0), "missing")
.await
.is_err());
assert!(repository.list_providers(false).await.unwrap().is_empty());
assert!(groups
.list_routing_groups()
.await
.unwrap()
.iter()
.all(|group| group.version == 1));
repository
.create_provider_in_routing_group(&provider, None, "selected")
.await
.unwrap();
for group in groups.list_routing_groups().await.unwrap() {
assert_eq!(group.version, if group.id == "selected" { 1 } else { 2 });
assert_eq!(
group.config_json["disabled_providers"],
if group.id == "selected" {
json!(["already-disabled"])
} else {
json!(["already-disabled", "new"])
}
);
}
let before = groups.list_routing_groups().await.unwrap();
assert!(repository
.create_provider_in_routing_group(&provider, Some(0), "other")
.await
.is_err());
assert_eq!(groups.list_routing_groups().await.unwrap(), before);
let stale = groups
.update_routing_group(
"other",
UpdateRoutingGroupRecord {
expected_version: Some(1),
config_json: Some(json!({"disabled_providers": []})),
..Default::default()
},
)
.await;
assert!(
matches!(stale, Err(DataLayerError::InvalidInput(message)) if message == "routing_group_version_conflict")
);
assert_eq!(groups.list_routing_groups().await.unwrap(), before);
let updated = groups
.update_routing_group(
"other",
UpdateRoutingGroupRecord {
expected_version: Some(2),
config_json: Some(json!({"disabled_providers": ["already-disabled"]})),
version: Some(2),
..Default::default()
},
)
.await
.unwrap()
.unwrap();
assert_eq!(updated.version, 3);
assert_eq!(
updated.config_json["disabled_providers"],
json!(["already-disabled"])
);
}
#[tokio::test]
async fn reads_provider_catalog_items_by_id() {
let repository = InMemoryProviderCatalogReadRepository::seed(
@@ -20,6 +20,58 @@ pub struct InMemoryRoutingGroupRepository {
}
impl InMemoryRoutingGroupRepository {
pub(crate) fn create_scoped_provider<T>(
&self,
selected_group_id: &str,
provider_id: &str,
create: impl FnOnce() -> Result<T, DataLayerError>,
) -> Result<T, DataLayerError> {
let mut groups = self.groups.write().expect("routing group repository lock");
if !groups.contains_key(selected_group_id) {
return Err(DataLayerError::InvalidInput(
"routing_group_not_found".to_string(),
));
}
let mut updated = groups.clone();
for group in updated.values_mut() {
let had_provider = group
.config_json
.get("disabled_providers")
.and_then(serde_json::Value::as_array)
.is_some_and(|disabled| {
disabled
.iter()
.any(|value| value.as_str() == Some(provider_id))
});
if (group.id == selected_group_id && !had_provider)
|| (group.id != selected_group_id && had_provider)
{
continue;
}
let object = group.config_json.as_object_mut().ok_or_else(|| {
DataLayerError::InvalidInput("routing group config must be an object".to_string())
})?;
let disabled = object
.entry("disabled_providers")
.or_insert_with(|| serde_json::json!([]));
let disabled = disabled.as_array_mut().ok_or_else(|| {
DataLayerError::InvalidInput("disabled_providers must be an array".to_string())
})?;
disabled.retain(|value| value.as_str() != Some(provider_id));
if group.id != selected_group_id {
disabled.push(serde_json::json!(provider_id));
}
group.version = group.version.saturating_add(1);
group.updated_at = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs() as i64;
}
let created = create()?;
*groups = updated;
Ok(created)
}
pub fn seed<I, B, V>(groups: I, bindings: B, versions: V) -> Self
where
I: IntoIterator<Item = StoredRoutingGroup>,
+4
View File
@@ -224,6 +224,10 @@ pub struct RoutingRule {
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
pub struct RoutingGroupConfig {
/// Providers excluded from every model in this group, including providers
/// otherwise selected by model policies or routing rules.
#[serde(default)]
pub disabled_providers: Vec<String>,
/// The default policy is global for the selected strategy group. Model
/// differences are expressed through `model_policies` and `rules`.
#[serde(default)]
+175 -10
View File
@@ -85,12 +85,17 @@ pub fn resolve_routing_policy(
selection_source: input.selection_source.to_string(),
requested_model: input.requested_model.to_string(),
resolved_model: input.resolved_model.to_string(),
priority_mode: config.default_policy.priority_mode,
// Legacy global_key values remain readable, but routing groups now
// always rank providers before their keys.
priority_mode: RoutingSetPriorityMode::Provider,
scheduling_mode: config.default_policy.scheduling_mode,
keep_priority_on_conversion: config.default_policy.keep_priority_on_conversion,
sticky_key_attempts: config.default_policy.sticky_key_attempts,
execution_policy: config.default_policy.execution_policy.clone(),
ranking_overlay: RankingOverlay::default(),
ranking_overlay: RankingOverlay {
disabled_providers: config.disabled_providers.clone(),
..RankingOverlay::default()
},
mutation_plan: MutationPlan::default(),
pool_policy_overrides: BTreeMap::new(),
matched_rules: Vec::new(),
@@ -203,14 +208,13 @@ fn apply_action(
policy.ranking_overlay.allowed_keys = key_ids.clone();
}
RoutingAction::SetScheduling {
priority_mode,
// Keep accepting the legacy field without re-enabling key-first
// scheduling through a model rule.
priority_mode: _,
scheduling_mode,
keep_priority_on_conversion,
sticky_key_attempts,
} => {
if let Some(priority_mode) = priority_mode {
policy.priority_mode = *priority_mode;
}
if let Some(scheduling_mode) = scheduling_mode {
policy.scheduling_mode = *scheduling_mode;
}
@@ -316,6 +320,164 @@ mod tests {
use super::*;
#[test]
fn legacy_key_scheduling_keeps_overrides_but_resolves_to_provider_ordering() {
let config: RoutingGroupConfig = serde_json::from_value(json!({
"default_policy": { "priority_mode": "global_key" },
"model_policies": [{
"model": "*",
"provider_priority_overrides": { "provider-a": 7 },
"key_priority_overrides": { "key-a": 2 },
"key_priority_overrides_by_format": { "openai:chat": { "key-a": 3 } },
"pool_priority_overrides": { "provider-pool": 4 }
}],
"rules": [{
"id": "legacy-key-client", "phase": "client_request",
"actions": [{ "type": "set_scheduling", "priority_mode": "global_key", "scheduling_mode": "fixed_order" }]
}]
}))
.expect("legacy key scheduling must stay readable");
let stored = serde_json::to_value(&config).unwrap();
assert_eq!(stored["default_policy"]["priority_mode"], "global_key");
assert_eq!(
stored["rules"][0]["actions"][0]["priority_mode"],
"global_key"
);
assert_eq!(
serde_json::from_value::<RoutingGroupConfig>(stored).unwrap(),
config
);
for phase in [
RoutingRulePhase::ClientRequest,
RoutingRulePhase::ProviderRequest,
] {
let policy = resolve_routing_policy(
&config,
RoutingPolicyInput {
group_id: Some("legacy-group"),
group_version: Some(1),
selection_source: "explicit",
requested_model: "model-a",
resolved_model: "model-a",
api_format: "openai:chat",
user_id: None,
api_key_id: None,
headers: &json!({}),
body: &json!({}),
phase,
},
)
.unwrap();
assert_eq!(policy.priority_mode, RoutingSetPriorityMode::Provider);
assert_eq!(
policy.scheduling_mode,
if phase == RoutingRulePhase::ClientRequest {
RoutingSchedulingMode::FixedOrder
} else {
RoutingSchedulingMode::CacheAffinity
}
);
assert_eq!(
policy.matched_rules.len(),
usize::from(phase == RoutingRulePhase::ClientRequest)
);
assert_eq!(
policy.ranking_overlay.provider_priority_overrides["provider-a"],
7
);
assert_eq!(policy.ranking_overlay.key_priority_overrides["key-a"], 2);
assert_eq!(
policy.ranking_overlay.pool_priority_overrides["provider-pool"],
4
);
assert_eq!(
policy
.ranking_overlay
.key_priority_for_format("key-a", "openai:chat", 99),
3
);
}
assert_eq!(
config.default_policy.priority_mode,
RoutingSetPriorityMode::GlobalKey
);
}
#[test]
fn group_disabled_providers_apply_to_every_model_and_cannot_be_reenabled() {
let config: RoutingGroupConfig = serde_json::from_value(json!({
"disabled_providers": ["provider-disabled"],
"model_policies": [{
"model": "model-allowlist",
"allowed_providers": ["provider-disabled", "provider-enabled"]
}],
"rules": [{
"id": "replace-provider-allowlist",
"conditions": { "field": "model", "op": "eq", "value": "rule-allowlist" },
"actions": [{
"type": "restrict_providers",
"provider_ids": ["provider-disabled", "provider-enabled"]
}, {
"type": "set_provider_priority",
"provider_id": "provider-disabled",
"priority": 0
}]
}, {
"id": "clear-provider-allowlist",
"conditions": { "field": "model", "op": "eq", "value": "rule-unrestricted" },
"actions": [{ "type": "restrict_providers", "provider_ids": [] }]
}]
}))
.expect("group provider exclusions should deserialize");
// The field survives the same round trip used when persisting or
// publishing strategy configuration.
let stored_config = serde_json::to_value(&config).unwrap();
assert_eq!(
stored_config["disabled_providers"],
json!(["provider-disabled"])
);
let config: RoutingGroupConfig = serde_json::from_value(stored_config).unwrap();
for model in [
"future-model",
"model-allowlist",
"rule-allowlist",
"rule-unrestricted",
] {
let policy = resolve_routing_policy(
&config,
RoutingPolicyInput {
group_id: Some("group-1"),
group_version: Some(1),
selection_source: "explicit",
requested_model: model,
resolved_model: model,
api_format: "openai:chat",
user_id: None,
api_key_id: None,
headers: &json!({}),
body: &json!({}),
phase: RoutingRulePhase::ClientRequest,
},
)
.expect("policy with group provider exclusions should resolve");
assert!(
!policy.ranking_overlay.provider_allowed("provider-disabled"),
"{model} must retain the group exclusion"
);
assert!(policy.ranking_overlay.provider_allowed("provider-enabled"));
let has_allowlist = matches!(model, "model-allowlist" | "rule-allowlist");
assert_eq!(
policy.ranking_overlay.provider_allowed("provider-unlisted"),
!has_allowlist,
"{model} should preserve its normal allowlist behavior"
);
}
}
#[test]
fn all_model_scheduling_and_rankings_apply_to_future_models() {
let config: RoutingGroupConfig = serde_json::from_value(json!({
@@ -330,6 +492,7 @@ mod tests {
"rules": []
}))
.expect("all-model scheduling config should deserialize");
assert!(config.disabled_providers.is_empty());
for model in ["existing-model", "future-model"] {
let policy = resolve_routing_policy(
@@ -349,7 +512,7 @@ mod tests {
},
)
.expect("all-model scheduling policy should resolve");
assert_eq!(policy.priority_mode, RoutingSetPriorityMode::GlobalKey);
assert_eq!(policy.priority_mode, RoutingSetPriorityMode::Provider);
assert_eq!(policy.scheduling_mode, RoutingSchedulingMode::LoadBalance);
assert_eq!(
policy
@@ -419,7 +582,7 @@ mod tests {
.is_empty());
assert!(policy.matched_rules.is_empty());
} else {
assert_eq!(policy.priority_mode, RoutingSetPriorityMode::GlobalKey);
assert_eq!(policy.priority_mode, RoutingSetPriorityMode::Provider);
assert_eq!(policy.scheduling_mode, RoutingSchedulingMode::FixedOrder);
assert_eq!(
policy
@@ -436,6 +599,7 @@ mod tests {
#[test]
fn resolves_model_policy_and_matching_rule() {
let config = RoutingGroupConfig {
disabled_providers: vec![],
default_policy: RoutingDefaultPolicy::default(),
model_policies: vec![RoutingModelPolicy {
model: "gpt-5".to_string(),
@@ -504,6 +668,7 @@ mod tests {
#[test]
fn default_policy_applies_to_models_without_an_override() {
let config = RoutingGroupConfig {
disabled_providers: vec![],
default_policy: RoutingDefaultPolicy {
priority_mode: RoutingSetPriorityMode::GlobalKey,
scheduling_mode: RoutingSchedulingMode::LoadBalance,
@@ -538,7 +703,7 @@ mod tests {
)
.expect("the specially configured model should resolve");
assert_eq!(special.priority_mode, RoutingSetPriorityMode::GlobalKey);
assert_eq!(special.priority_mode, RoutingSetPriorityMode::Provider);
assert_eq!(special.scheduling_mode, RoutingSchedulingMode::LoadBalance);
assert!(special.keep_priority_on_conversion);
assert_eq!(special.sticky_key_attempts, 3);
@@ -572,7 +737,7 @@ mod tests {
)
.expect("an unconfigured model should keep using the default policy");
assert_eq!(ordinary.priority_mode, RoutingSetPriorityMode::GlobalKey);
assert_eq!(ordinary.priority_mode, RoutingSetPriorityMode::Provider);
assert_eq!(ordinary.scheduling_mode, RoutingSchedulingMode::LoadBalance);
assert!(ordinary.keep_priority_on_conversion);
assert_eq!(ordinary.sticky_key_attempts, 3);
+47 -5
View File
@@ -13,6 +13,9 @@ pub enum CandidateKind {
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct RankingOverlay {
/// Group-wide exclusions take precedence over every provider allowlist.
#[serde(default)]
pub disabled_providers: Vec<String>,
#[serde(default)]
pub allowed_providers: Vec<String>,
#[serde(default)]
@@ -112,11 +115,15 @@ impl RankingOverlay {
}
pub fn provider_allowed(&self, provider_id: &str) -> bool {
self.allowed_providers.is_empty()
|| self
.allowed_providers
.iter()
.any(|item| item == provider_id)
!self
.disabled_providers
.iter()
.any(|item| item == provider_id)
&& (self.allowed_providers.is_empty()
|| self
.allowed_providers
.iter()
.any(|item| item == provider_id))
}
pub fn key_allowed(&self, key_id: &str) -> bool {
@@ -180,6 +187,41 @@ mod tests {
use super::*;
#[test]
fn disabled_providers_take_precedence_over_allowlists() {
let mut overlay = RankingOverlay {
disabled_providers: vec!["provider-disabled".to_string()],
..RankingOverlay::default()
};
assert!(!overlay.provider_allowed("provider-disabled"));
assert!(overlay.provider_allowed("provider-enabled"));
overlay.allowed_providers = vec![
"provider-disabled".to_string(),
"provider-enabled".to_string(),
];
assert!(!overlay.provider_allowed("provider-disabled"));
assert!(overlay.provider_allowed("provider-enabled"));
assert!(!overlay.provider_allowed("provider-unlisted"));
// An allowlist containing only disabled providers must not become an
// empty allowlist, which would otherwise allow unrelated providers.
overlay.allowed_providers = vec!["provider-disabled".to_string()];
assert!(!overlay.provider_allowed("provider-disabled"));
assert!(!overlay.provider_allowed("provider-enabled"));
}
#[test]
fn legacy_overlay_without_disabled_providers_preserves_provider_selection() {
let overlay: RankingOverlay = serde_json::from_value(serde_json::json!({
"allowed_providers": ["provider-enabled"]
}))
.expect("legacy overlays should remain readable");
assert!(overlay.disabled_providers.is_empty());
assert!(overlay.provider_allowed("provider-enabled"));
assert!(!overlay.provider_allowed("provider-unlisted"));
}
#[test]
fn overlay_applies_provider_and_key_priority() {
let overlay = RankingOverlay {