mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-11 03:39:49 +08:00
fix(routing): harden routed pool scheduling
This commit is contained in:
@@ -10,6 +10,7 @@ const POOL_ALLOWED_SCHEDULING_PRESETS: &[&str] = &[
|
||||
"load_balance",
|
||||
"single_account",
|
||||
"priority_first",
|
||||
"free_team_first",
|
||||
"free_first",
|
||||
"team_first",
|
||||
"plus_first",
|
||||
@@ -166,8 +167,9 @@ fn parse_pool_score_rules(pool_advanced: &Map<String, Value>) -> PoolMemberScore
|
||||
|
||||
fn normalize_pool_preset_mode(preset: &str, raw_mode: Option<&Value>) -> Option<String> {
|
||||
match preset {
|
||||
"free_first" | "team_first" | "plus_first" | "pro_first" => {
|
||||
"free_team_first" | "free_first" | "team_first" | "plus_first" | "pro_first" => {
|
||||
let default_mode = match preset {
|
||||
"free_team_first" => "both",
|
||||
"free_first" => "free_only",
|
||||
"team_first" => "team_only",
|
||||
"plus_first" => "plus_only",
|
||||
@@ -180,6 +182,9 @@ fn normalize_pool_preset_mode(preset: &str, raw_mode: Option<&Value>) -> Option<
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|value| value.to_ascii_lowercase())
|
||||
.filter(|value| match preset {
|
||||
"free_team_first" => {
|
||||
matches!(value.as_str(), "free_only" | "team_only" | "both")
|
||||
}
|
||||
"free_first" => value == "free_only",
|
||||
"team_first" => value == "team_only",
|
||||
"plus_first" => value == "plus_only",
|
||||
@@ -786,7 +791,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn retired_free_team_first_preset_is_rejected() {
|
||||
fn legacy_free_team_first_preset_is_preserved() {
|
||||
let config = admin_provider_pool_config_from_config_value(Some(&json!({
|
||||
"pool_advanced": {
|
||||
"scheduling_presets": [
|
||||
@@ -797,8 +802,11 @@ mod tests {
|
||||
.expect("pool config should parse");
|
||||
|
||||
assert_eq!(config.scheduling_presets.len(), 1);
|
||||
assert_eq!(config.scheduling_presets[0].preset, "lru");
|
||||
assert_eq!(config.scheduling_presets[0].mode, None);
|
||||
assert_eq!(config.scheduling_presets[0].preset, "free_team_first");
|
||||
assert_eq!(
|
||||
config.scheduling_presets[0].mode.as_deref(),
|
||||
Some("team_only")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -65,6 +65,8 @@ use aether_model_fetch::{
|
||||
aggregate_models_for_cache, fetch_models_from_transports, json_string_list,
|
||||
preset_models_for_provider, selected_models_fetch_endpoints,
|
||||
};
|
||||
use aether_pool_core::PoolSchedulingPreset;
|
||||
use aether_provider_pool::ProviderPoolService;
|
||||
use aether_scheduler_core::provider_key_circuit_payload_is_active_open_at;
|
||||
use axum::{
|
||||
body::{to_bytes, Body},
|
||||
@@ -1106,15 +1108,26 @@ fn provider_query_pool_sort_seed() -> String {
|
||||
|
||||
fn provider_query_ai_pool_scheduling_config(
|
||||
config: &AdminProviderPoolConfig,
|
||||
provider_type: &str,
|
||||
) -> AiPoolSchedulingConfig {
|
||||
let presets = config
|
||||
.scheduling_presets
|
||||
.iter()
|
||||
.map(|preset| PoolSchedulingPreset {
|
||||
preset: preset.preset.clone(),
|
||||
enabled: preset.enabled,
|
||||
mode: preset.mode.clone(),
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let normalized_presets = ProviderPoolService::with_builtin_adapters()
|
||||
.normalize_scheduling_presets(provider_type, &presets);
|
||||
AiPoolSchedulingConfig {
|
||||
scheduling_presets: config
|
||||
.scheduling_presets
|
||||
.iter()
|
||||
scheduling_presets: normalized_presets
|
||||
.into_iter()
|
||||
.map(|preset| AiPoolSchedulingPreset {
|
||||
preset: preset.preset.clone(),
|
||||
preset: preset.preset,
|
||||
enabled: preset.enabled,
|
||||
mode: preset.mode.clone(),
|
||||
mode: preset.mode,
|
||||
})
|
||||
.collect(),
|
||||
lru_enabled: config.lru_enabled,
|
||||
@@ -1324,7 +1337,8 @@ async fn provider_query_apply_pool_scheduler_to_test_candidates(
|
||||
provider.id.clone(),
|
||||
provider_query_ai_pool_runtime_state(&runtime),
|
||||
);
|
||||
let pool_config = provider_query_ai_pool_scheduling_config(pool_config);
|
||||
let pool_config =
|
||||
provider_query_ai_pool_scheduling_config(pool_config, provider.provider_type.as_str());
|
||||
let inputs = keys
|
||||
.into_iter()
|
||||
.map(|key| {
|
||||
|
||||
@@ -320,6 +320,31 @@ fn provider_query_model_test_empty_selected_key_ids_keep_default_selection() {
|
||||
assert!(provider_query_extract_api_key_ids(&json!({ "api_key_ids": [] })).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_codex_pool_config_injects_recent_refresh() {
|
||||
let raw_config = json!({
|
||||
"pool_advanced": {
|
||||
"scheduling_presets": [{
|
||||
"preset": "cache_affinity",
|
||||
"enabled": true
|
||||
}]
|
||||
}
|
||||
});
|
||||
let config = admin_provider_pool_config_from_config_value(Some(&raw_config))
|
||||
.expect("pool config should parse");
|
||||
|
||||
let normalized = provider_query_ai_pool_scheduling_config(&config, "codex");
|
||||
|
||||
assert_eq!(
|
||||
normalized
|
||||
.scheduling_presets
|
||||
.iter()
|
||||
.map(|preset| preset.preset.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
["cache_affinity", "recent_refresh"]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_query_standard_test_resolves_codex_responses_upstream_streaming() {
|
||||
assert!(provider_query_resolve_standard_test_upstream_is_stream(
|
||||
|
||||
@@ -346,20 +346,6 @@ async fn maybe_forward_public_request_to_tunnel_owner(
|
||||
}) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let cache_affinity_enabled = match read_scheduler_ordering_config(state).await {
|
||||
Ok(config) => config.scheduling_mode == SchedulerSchedulingMode::CacheAffinity,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %request_context.trace_id,
|
||||
error = ?err,
|
||||
"gateway failed to load scheduler config while checking tunnel affinity forwarding mode"
|
||||
);
|
||||
SchedulerSchedulingMode::default() == SchedulerSchedulingMode::CacheAffinity
|
||||
}
|
||||
};
|
||||
if !cache_affinity_enabled {
|
||||
return Ok(None);
|
||||
}
|
||||
let Some(api_format) = decision
|
||||
.auth_endpoint_signature
|
||||
.as_deref()
|
||||
@@ -382,21 +368,66 @@ async fn maybe_forward_public_request_to_tunnel_owner(
|
||||
crate::headers::decoded_request_body_bytes(&parts.headers, body.as_ref()).ok()?;
|
||||
serde_json::from_slice::<serde_json::Value>(body.as_ref()).ok()
|
||||
});
|
||||
let client_session_affinity =
|
||||
crate::client_session_affinity::client_session_affinity_from_api_request(
|
||||
api_format,
|
||||
&parts.headers,
|
||||
body_json.as_ref(),
|
||||
);
|
||||
let Some(target) = crate::scheduler::affinity::read_cached_scheduler_affinity_target(
|
||||
let empty_body_json = serde_json::Value::Null;
|
||||
let affinity_context = match crate::ai_serving::resolve_tunnel_scheduler_affinity_context(
|
||||
state,
|
||||
&auth_context.api_key_id,
|
||||
client_session_affinity.as_ref(),
|
||||
parts,
|
||||
decision,
|
||||
requested_model,
|
||||
body_json.as_ref().unwrap_or(&empty_body_json),
|
||||
api_format,
|
||||
&requested_model,
|
||||
) else {
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(context)) => context,
|
||||
Ok(None) => return Ok(None),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %request_context.trace_id,
|
||||
error = ?err,
|
||||
"gateway failed to resolve routing policy while checking tunnel affinity forwarding"
|
||||
);
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
let target = if let Some(policy_context) = affinity_context.policy_context.as_ref() {
|
||||
crate::scheduler::affinity::read_cached_scheduler_affinity_target_with_policy_context(
|
||||
state,
|
||||
&auth_context.api_key_id,
|
||||
affinity_context.client_session_affinity.as_ref(),
|
||||
api_format,
|
||||
&affinity_context.requested_model,
|
||||
policy_context,
|
||||
)
|
||||
} else {
|
||||
let cache_affinity_enabled = match read_scheduler_ordering_config(state).await {
|
||||
Ok(config) => config.scheduling_mode == SchedulerSchedulingMode::CacheAffinity,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %request_context.trace_id,
|
||||
error = ?err,
|
||||
"gateway failed to load scheduler config while checking tunnel affinity forwarding mode"
|
||||
);
|
||||
SchedulerSchedulingMode::default() == SchedulerSchedulingMode::CacheAffinity
|
||||
}
|
||||
};
|
||||
if !cache_affinity_enabled {
|
||||
return Ok(None);
|
||||
}
|
||||
crate::scheduler::affinity::read_cached_scheduler_affinity_target(
|
||||
state,
|
||||
&auth_context.api_key_id,
|
||||
affinity_context.client_session_affinity.as_ref(),
|
||||
api_format,
|
||||
&affinity_context.requested_model,
|
||||
)
|
||||
};
|
||||
let Some(target) = target else {
|
||||
return Ok(None);
|
||||
};
|
||||
if !routing_overlay_allows_affinity_target(affinity_context.routing_overlay.as_ref(), &target) {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let transport = match state
|
||||
.read_provider_transport_snapshot(&target.provider_id, &target.endpoint_id, &target.key_id)
|
||||
@@ -547,6 +578,16 @@ async fn maybe_forward_public_request_to_tunnel_owner(
|
||||
Ok(Some(response))
|
||||
}
|
||||
|
||||
fn routing_overlay_allows_affinity_target(
|
||||
routing_overlay: Option<&aether_routing_core::RankingOverlay>,
|
||||
target: &aether_scheduler_core::SchedulerAffinityTarget,
|
||||
) -> bool {
|
||||
routing_overlay.is_none_or(|overlay| {
|
||||
overlay.provider_allowed(target.provider_id.as_str())
|
||||
&& overlay.key_allowed(target.key_id.as_str())
|
||||
})
|
||||
}
|
||||
|
||||
fn owner_forward_request_is_stream(
|
||||
parts: &http::request::Parts,
|
||||
decision: &GatewayControlDecision,
|
||||
@@ -2327,14 +2368,51 @@ mod tests {
|
||||
api_key_remote_ip_allowed, buffer_and_normalize_request_body,
|
||||
diagnostic_is_auth_api_key_concurrency_limited, local_execution_runtime_miss_detail,
|
||||
owner_forward_request_is_stream, restore_redacted_stream_execution_response,
|
||||
restore_redacted_sync_execution_response, GatewayControlDecision,
|
||||
LocalExecutionRuntimeMissDiagnostic, RequestBodyBufferError, RequestBodyBufferPolicy,
|
||||
restore_redacted_sync_execution_response, routing_overlay_allows_affinity_target,
|
||||
GatewayControlDecision, LocalExecutionRuntimeMissDiagnostic, RequestBodyBufferError,
|
||||
RequestBodyBufferPolicy,
|
||||
};
|
||||
use axum::body::{to_bytes, Body, Bytes};
|
||||
use axum::http::{header, HeaderMap, HeaderValue, Method, Response};
|
||||
use serde_json::json;
|
||||
use tokio::sync::Semaphore;
|
||||
|
||||
#[test]
|
||||
fn routing_overlay_blocks_disallowed_tunnel_affinity_target() {
|
||||
let target = aether_scheduler_core::SchedulerAffinityTarget {
|
||||
provider_id: "provider-allowed".to_string(),
|
||||
endpoint_id: "endpoint-1".to_string(),
|
||||
key_id: "key-allowed".to_string(),
|
||||
};
|
||||
let matching = aether_routing_core::RankingOverlay {
|
||||
allowed_providers: vec!["provider-allowed".to_string()],
|
||||
allowed_keys: vec!["key-allowed".to_string()],
|
||||
..aether_routing_core::RankingOverlay::default()
|
||||
};
|
||||
let wrong_provider = aether_routing_core::RankingOverlay {
|
||||
allowed_providers: vec!["provider-other".to_string()],
|
||||
..matching.clone()
|
||||
};
|
||||
let wrong_key = aether_routing_core::RankingOverlay {
|
||||
allowed_keys: vec!["key-other".to_string()],
|
||||
..matching.clone()
|
||||
};
|
||||
|
||||
assert!(routing_overlay_allows_affinity_target(None, &target));
|
||||
assert!(routing_overlay_allows_affinity_target(
|
||||
Some(&matching),
|
||||
&target
|
||||
));
|
||||
assert!(!routing_overlay_allows_affinity_target(
|
||||
Some(&wrong_provider),
|
||||
&target
|
||||
));
|
||||
assert!(!routing_overlay_allows_affinity_target(
|
||||
Some(&wrong_key),
|
||||
&target
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn owner_forward_uses_search_protocol_timeout_semantics() {
|
||||
let request = http::Request::builder()
|
||||
|
||||
Reference in New Issue
Block a user