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
@@ -1892,6 +1892,84 @@ mod tests {
.is_none());
}
#[tokio::test]
async fn routing_policy_excludes_group_disabled_providers_from_candidate_pages() {
let repository: Arc<dyn MinimalCandidateSelectionReadRepository> =
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed([
standard_candidate_row("provider-disabled", "openai:chat", 0),
standard_candidate_row("provider-enabled", "openai:chat", 1),
]));
let app = AppState::new()
.expect("gateway state should build")
.with_data_state_for_tests(
GatewayDataState::with_minimal_candidate_selection_reader_for_tests(repository),
);
let auth_snapshot = unrestricted_auth_snapshot();
let model_directive_policy =
crate::system_features::ModelDirectivePolicySnapshot::load(&app).await;
let config = serde_json::from_value(serde_json::json!({
"disabled_providers": ["provider-disabled"],
"model_policies": [{
"model": "*",
"allowed_providers": ["provider-disabled", "provider-enabled"]
}]
}))
.expect("routing config should parse");
let routing_policy = aether_routing_core::resolve_routing_policy(
&config,
aether_routing_core::RoutingPolicyInput {
group_id: Some("routing-group-1"),
group_version: Some(1),
selection_source: "test",
requested_model: "gpt-5",
resolved_model: "gpt-5",
api_format: "openai:chat",
user_id: None,
api_key_id: None,
headers: &serde_json::json!({}),
body: &serde_json::json!({}),
phase: aether_routing_core::RoutingRulePhase::ClientRequest,
},
)
.expect("routing policy should resolve");
let mut cursor = LocalCandidatePreselectionPageCursor::new(
PlannerAppState::new(&app),
&model_directive_policy,
"openai:chat",
"gpt-5",
None,
false,
None,
&auth_snapshot,
Some(&routing_policy),
None,
None,
true,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
false,
None,
)
.await;
let page = cursor
.next_page()
.await
.expect("routing candidate scan should succeed")
.expect("the enabled provider should remain");
assert_eq!(
page.candidates
.iter()
.map(|candidate| candidate.provider_id.as_str())
.collect::<Vec<_>>(),
vec!["provider-enabled"]
);
assert!(cursor
.next_page()
.await
.expect("routing scan should finish")
.is_none());
}
#[tokio::test]
async fn routing_policy_collects_candidate_pages_before_final_ranking() {
let rows = (0..300)
@@ -1735,7 +1735,7 @@ mod tests {
assert_eq!(policy.group_version, Some(4));
assert_eq!(
policy.priority_mode,
aether_routing_core::RoutingSetPriorityMode::GlobalKey
aether_routing_core::RoutingSetPriorityMode::Provider
);
assert_eq!(
policy.scheduling_mode,
@@ -562,6 +562,30 @@ impl GatewayDataState {
Ok(created)
}
pub(crate) async fn create_provider_catalog_provider_in_routing_group(
&self,
provider: &StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
routing_group_id: &str,
) -> Result<Option<StoredProviderCatalogProvider>, DataLayerError> {
let created = match &self.provider_catalog_writer {
Some(repository) => repository
.create_provider_in_routing_group(
provider,
shift_existing_priorities_from,
routing_group_id,
)
.await
.map(Some),
None => Ok(None),
}?;
if created.is_some() {
self.clear_provider_catalog_cache();
self.clear_routing_group_cache();
}
Ok(created)
}
pub(crate) async fn update_provider_catalog_provider(
&self,
provider: &StoredProviderCatalogProvider,
@@ -62,6 +62,16 @@ pub(crate) async fn maybe_build_local_admin_provider_writes_response(
)));
}
};
let routing_group_id = payload
.routing_group_id
.as_deref()
.map(str::trim)
.map(str::to_string);
if routing_group_id.as_deref() == Some("") {
return Ok(Some(build_admin_provider_bad_request_response(
"routing_group_id 不能为空",
)));
}
let (record, shift_existing_priorities_from) =
match state.build_admin_create_provider_record(payload).await {
Ok(record) => record,
@@ -69,10 +79,23 @@ pub(crate) async fn maybe_build_local_admin_provider_writes_response(
return Ok(Some(build_admin_provider_bad_request_response(message)));
}
};
let Some(created_provider) = state
.create_provider_catalog_provider(&record, shift_existing_priorities_from)
.await?
else {
let created = match routing_group_id.as_deref() {
Some(group_id) => {
state
.create_provider_catalog_provider_in_routing_group(
&record,
shift_existing_priorities_from,
group_id,
)
.await?
}
None => {
state
.create_provider_catalog_provider(&record, shift_existing_priorities_from)
.await?
}
};
let Some(created_provider) = created else {
return Ok(Some(build_admin_providers_data_unavailable_response()));
};
@@ -179,6 +179,8 @@ pub(crate) struct AdminCodexResetCreditConsumeRequest {
pub(crate) struct AdminProviderCreateRequest {
pub(crate) name: String,
#[serde(default)]
pub(crate) routing_group_id: Option<String>,
#[serde(default)]
pub(crate) provider_type: Option<String>,
#[serde(default)]
pub(crate) description: Option<String>,
@@ -420,6 +420,24 @@ impl<'a> AdminAppState<'a> {
.await
}
pub(crate) async fn create_provider_catalog_provider_in_routing_group(
&self,
provider: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
routing_group_id: &str,
) -> Result<
Option<aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider>,
GatewayError,
> {
self.app
.create_provider_catalog_provider_in_routing_group(
provider,
shift_existing_priorities_from,
routing_group_id,
)
.await
}
pub(crate) async fn update_provider_catalog_provider(
&self,
provider: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider,
@@ -350,6 +350,7 @@ async fn publish_routing_group(
.update_routing_group(
group_id,
UpdateRoutingGroupRecord {
expected_version: Some(group.version),
version: Some(next_version),
updated_at: now,
published_at: Some(Some(now)),
@@ -455,6 +456,9 @@ fn build_routing_group_update_patch(
updated_at: current_unix_secs() as i64,
..UpdateRoutingGroupRecord::default()
};
if let Some(value) = object.get("expected_version") {
patch.expected_version = Some(required_i64(value, "expected_version")?);
}
if let Some(value) = object.get("name") {
patch.name = Some(required_string(value, "name")?);
}
@@ -2909,6 +2909,10 @@ mod tests {
allowed_keys: vec!["key-other".to_string()],
..matching.clone()
};
let disabled_provider = aether_routing_core::RankingOverlay {
disabled_providers: vec!["provider-allowed".to_string()],
..matching.clone()
};
assert!(routing_overlay_allows_affinity_target(None, &target));
assert!(routing_overlay_allows_affinity_target(
@@ -2923,6 +2927,10 @@ mod tests {
Some(&wrong_key),
&target
));
assert!(!routing_overlay_allows_affinity_target(
Some(&disabled_provider),
&target
));
}
#[test]
+87 -6
View File
@@ -92,7 +92,8 @@ pub(crate) fn resolve_gateway_static_default_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: default_policy.priority_mode,
// Match the full resolver even when a stored default still says global_key.
priority_mode: RoutingSetPriorityMode::Provider,
scheduling_mode: default_policy.scheduling_mode,
keep_priority_on_conversion: default_policy.keep_priority_on_conversion,
sticky_key_attempts: default_policy.sticky_key_attempts,
@@ -110,10 +111,11 @@ fn static_default_policy_fields(
let Some(object) = config_json.as_object() else {
return Ok(None);
};
// A strategy's default policy applies to every model. Only model policies
// and rules require the request-context-aware resolver; unknown legacy
// fields (including the removed group allowlist) are intentionally ignored.
if !routing_array_field_is_missing_or_empty(object, "model_policies")
// A strategy's default policy applies to every model. Provider exclusions,
// model policies and rules require the full resolver to build the overlay;
// unknown legacy fields (including the removed group allowlist) are ignored.
if !routing_array_field_is_missing_or_empty(object, "disabled_providers")
|| !routing_array_field_is_missing_or_empty(object, "model_policies")
|| !routing_array_field_is_missing_or_empty(object, "rules")
{
return Ok(None);
@@ -275,7 +277,7 @@ mod tests {
);
assert_eq!(
static_policy.priority_mode,
RoutingSetPriorityMode::GlobalKey
RoutingSetPriorityMode::Provider
);
assert_eq!(
static_policy.scheduling_mode,
@@ -313,6 +315,85 @@ mod tests {
assert!(policy.is_none());
}
#[test]
fn provider_exclusions_without_model_policies_or_rules_use_full_resolver() {
for model in ["model-a", "future-model"] {
let config = json!({
"disabled_providers": ["provider-disabled"],
"model_policies": [],
"rules": []
});
let static_policy = resolve_gateway_static_default_routing_policy(
GatewayStaticRoutingPolicyInput {
group_id: Some("group-1"),
group_version: Some(1),
group_config_json: &config,
selection_source: "system_default",
requested_model: model,
resolved_model: model,
},
)
.expect("provider exclusions should defer to the full resolver");
assert!(static_policy.is_none());
let policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput {
group_id: Some("group-1"),
group_version: Some(1),
group_config_json: &config,
selection_source: "system_default",
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("group exclusions should resolve");
assert!(!policy.ranking_overlay.provider_allowed("provider-disabled"));
assert!(policy.ranking_overlay.provider_allowed("provider-enabled"));
}
}
#[test]
fn empty_provider_exclusions_preserve_static_default_fast_path() {
for config in [json!({}), json!({"disabled_providers": []})] {
assert!(static_default_policy_fields(&config)
.expect("empty provider exclusions should be valid")
.is_some());
}
}
#[test]
fn malformed_provider_exclusions_cannot_bypass_full_config_validation() {
for disabled in [json!("provider-disabled"), json!([42]), Value::Null] {
let config = json!({"disabled_providers": disabled});
let error = resolve_gateway_routing_policy(GatewayRoutingPolicyInput {
group_id: Some("group-1"),
group_version: Some(1),
group_config_json: &config,
selection_source: "system_default",
requested_model: "model-a",
resolved_model: "model-a",
api_format: "openai:chat",
user_id: None,
api_key_id: None,
headers: &json!({}),
body: &json!({}),
phase: RoutingRulePhase::ClientRequest,
})
.expect_err("malformed provider exclusions must not be silently ignored");
assert!(matches!(
error,
GatewayError::Client {
status: StatusCode::BAD_REQUEST,
..
}
));
}
}
#[test]
fn routing_config_errors_do_not_echo_config_values() {
let secret = "https://internal.example/?token=Bearer-secret";
+35 -4
View File
@@ -102,8 +102,12 @@ impl SchedulerOrderingConfig {
fn scheduler_priority_mode_from_routing(mode: RoutingSetPriorityMode) -> SchedulerPriorityMode {
match mode {
RoutingSetPriorityMode::Provider => SchedulerPriorityMode::Provider,
RoutingSetPriorityMode::GlobalKey => SchedulerPriorityMode::GlobalKey,
// Defaults and previously resolved snapshots may still carry the
// removed global_key mode. Keep that compatibility at this boundary,
// without changing the scheduler's independent GlobalKey capability.
RoutingSetPriorityMode::Provider | RoutingSetPriorityMode::GlobalKey => {
SchedulerPriorityMode::Provider
}
}
}
@@ -161,6 +165,33 @@ mod tests {
use super::*;
use crate::data::GatewayDataState;
#[test]
fn legacy_resolved_key_snapshot_uses_provider_ordering() {
let snapshot: ResolvedRoutingPolicy = serde_json::from_value(json!({
"group_id": "legacy-group",
"selection_source": "system_default",
"requested_model": "model-a",
"resolved_model": "model-a",
"priority_mode": "global_key",
"scheduling_mode": "fixed_order",
"keep_priority_on_conversion": true,
"sticky_key_attempts": 4,
"ranking_overlay": { "key_priority_overrides": { "key-a": 2 } },
"mutation_plan": { "body_patch": [], "header_patch": [] }
}))
.expect("legacy snapshots must remain readable");
let ordering = SchedulerOrderingConfig::from_routing_policy(&snapshot);
assert_eq!(ordering.priority_mode, SchedulerPriorityMode::Provider);
assert_eq!(
ordering.scheduling_mode,
SchedulerSchedulingMode::FixedOrder
);
assert!(ordering.keep_priority_on_conversion);
assert_eq!(ordering.sticky_key_attempts, 4);
assert_eq!(snapshot.ranking_overlay.key_priority_overrides["key-a"], 2);
}
async fn create_system_default(
repository: &InMemoryRoutingGroupRepository,
enabled: bool,
@@ -185,14 +216,14 @@ mod tests {
}
#[tokio::test]
async fn system_default_routing_group_exposes_strategy_ordering() {
async fn legacy_system_default_routing_group_uses_provider_ordering() {
let repository = Arc::new(InMemoryRoutingGroupRepository::default());
create_system_default(
&repository,
true,
json!({
"default_policy": {
"priority_mode": "provider",
"priority_mode": "global_key",
"scheduling_mode": "fixed_order",
"keep_priority_on_conversion": false
}
+38
View File
@@ -639,6 +639,44 @@ impl AppState {
}
}
pub(crate) async fn create_provider_catalog_provider_in_routing_group(
&self,
provider: &provider_catalog::StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
routing_group_id: &str,
) -> Result<Option<provider_catalog::StoredProviderCatalogProvider>, GatewayError> {
let protected = self.protect_provider_catalog_provider(provider)?;
let created = self
.data
.create_provider_catalog_provider_in_routing_group(
&protected,
shift_existing_priorities_from,
routing_group_id,
)
.await
.map_err(|err| match err {
aether_data_contracts::DataLayerError::InvalidInput(ref message)
if message == "routing_group_not_found" =>
{
GatewayError::Client {
status: axum::http::StatusCode::NOT_FOUND,
message: "策略分组不存在,请刷新后重试".to_string(),
}
}
_ => GatewayError::Internal(err.to_string()),
})?;
if created.is_some() {
self.invalidate_provider_routing_caches();
}
match created {
Some(provider) => self
.open_provider_catalog_provider(provider)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn update_provider_catalog_provider(
&self,
provider: &provider_catalog::StoredProviderCatalogProvider,
@@ -160,7 +160,17 @@ impl AppState {
.data
.update_routing_group(id, patch)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
.map_err(|err| match err {
aether_data_contracts::DataLayerError::InvalidInput(ref message)
if message == "routing_group_version_conflict" =>
{
GatewayError::Client {
status: axum::http::StatusCode::CONFLICT,
message: "策略分组已被修改,请刷新后重试".to_string(),
}
}
_ => GatewayError::Internal(err.to_string()),
})?;
if updated.is_some() {
self.invalidate_provider_routing_caches();
}
@@ -42,6 +42,90 @@ const ADMIN_PROVIDERS_DATA_UNAVAILABLE_DETAIL: &str = "Admin provider catalog da
mod health;
#[tokio::test]
async fn gateway_creates_provider_in_selected_routing_group_and_rejects_stale_saves() {
use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository;
use aether_data_contracts::repository::routing_profiles::{
CreateRoutingGroupRecord, RoutingGroupReadRepository, RoutingGroupWriteRepository,
};
let groups = Arc::new(InMemoryRoutingGroupRepository::default());
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!({}),
version: 1,
created_at: 1,
updated_at: 1,
published_at: None,
})
.await
.unwrap();
}
let providers = Arc::new(
InMemoryProviderCatalogReadRepository::default().with_routing_groups(groups.clone()),
);
let gateway = build_router_with_state(
AppState::new().unwrap().with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(providers.clone())
.with_routing_group_repository_for_tests(groups.clone()),
),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let client = reqwest::Client::new();
let request = |name: &str, group: &str| {
client
.post(format!("{gateway_url}/api/admin/providers/"))
.header(GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({"name": name, "routing_group_id": group}))
};
let missing = request("missing", "gone").send().await.unwrap();
assert_eq!(missing.status(), StatusCode::NOT_FOUND);
assert!(providers.list_providers(false).await.unwrap().is_empty());
let response = request("scoped", "selected").send().await.unwrap();
let status = response.status();
let payload: serde_json::Value = response.json().await.unwrap();
assert_eq!(status, StatusCode::OK, "{payload}");
let provider_id = payload["id"].as_str().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!(null)
} else {
json!([provider_id])
}
);
}
let response = client
.patch(format!("{gateway_url}/api/admin/routing/groups/other"))
.header(GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({"expected_version": 1, "config_json": {"disabled_providers": []}}))
.send()
.await
.unwrap();
assert_eq!(
response.status(),
StatusCode::CONFLICT,
"{}",
response.text().await.unwrap()
);
gateway_handle.abort();
}
async fn provider_health_summary(
endpoints: &[StoredProviderCatalogEndpoint],
keys: &[StoredProviderCatalogKey],