mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-11 19:59:50 +08:00
feat: unify provider scheduling workspace
This commit is contained in:
@@ -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>,
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user