fix scheduler affinity candidate selection

This commit is contained in:
fawney19
2026-04-30 16:27:24 +08:00
parent 558abfcfa3
commit 33aa70c22b
24 changed files with 1715 additions and 104 deletions
@@ -115,6 +115,15 @@ mod tests {
"gemini:generate_content",
]
);
assert_eq!(
request_candidate_api_formats("claude:messages", false),
vec![
"claude:messages",
"openai:chat",
"openai:responses",
"gemini:generate_content",
]
);
assert_eq!(
request_candidate_api_formats("openai:cli", false),
Vec::<&'static str>::new()
@@ -609,8 +609,7 @@ mod tests {
}
#[tokio::test]
async fn fixed_order_local_execution_ranking_keeps_provider_priority_before_format_preference()
{
async fn fixed_order_local_execution_ranking_demotes_cross_format_before_provider_priority() {
let provider_catalog = InMemoryProviderCatalogReadRepository::seed(
vec![
sample_provider_with_options("provider-same", false, 10),
@@ -662,8 +661,8 @@ mod tests {
)
.await;
assert_eq!(ranked[0].endpoint_id, "endpoint-cross");
assert_eq!(ranked[1].endpoint_id, "endpoint-same");
assert_eq!(ranked[0].endpoint_id, "endpoint-same");
assert_eq!(ranked[1].endpoint_id, "endpoint-cross");
}
#[tokio::test]
@@ -1405,6 +1404,181 @@ mod tests {
);
}
#[tokio::test]
async fn first_request_same_key_exact_endpoint_beats_cross_format_without_affinity() {
let mut openai_endpoint =
sample_endpoint_for_provider("provider-shared", "endpoint-openai", "openai:chat");
openai_endpoint.format_acceptance_config = Some(json!({
"enabled": true,
"accept_formats": ["claude:messages"],
}));
let provider_catalog = InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider_with_options("provider-shared", false, 0)],
vec![
openai_endpoint,
sample_endpoint_for_provider(
"provider-shared",
"endpoint-claude",
"claude:messages",
),
],
vec![sample_key_for_provider_with_options(
"provider-shared",
"key-shared",
"",
true,
Some(json!(["openai:chat", "claude:messages"])),
None,
)],
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
);
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data_state);
let (ranked, skipped) = resolve_and_rank_local_execution_candidates(
PlannerAppState::new(&state),
vec![
sample_priority_candidate(
"provider-shared",
"endpoint-openai",
"key-shared",
"openai:chat",
Some(0),
0,
),
sample_priority_candidate(
"provider-shared",
"endpoint-claude",
"key-shared",
"claude:messages",
Some(0),
0,
),
],
"claude:messages",
"gpt-4.1",
None,
None,
None,
)
.await;
assert!(skipped.is_empty());
assert_eq!(ranked[0].candidate.endpoint_id, "endpoint-claude");
assert_eq!(ranked[1].candidate.endpoint_id, "endpoint-openai");
assert_eq!(
ranked[1]
.ranking
.as_ref()
.and_then(|ranking| ranking.promoted_by),
None
);
assert_eq!(
ranked[1]
.ranking
.as_ref()
.and_then(|ranking| ranking.demoted_by),
Some(aether_scheduler_core::RANKING_REASON_CROSS_FORMAT)
);
}
#[tokio::test]
async fn cached_affinity_promotes_cross_format_over_same_key_exact_endpoint() {
let mut openai_endpoint =
sample_endpoint_for_provider("provider-shared", "endpoint-openai", "openai:chat");
openai_endpoint.format_acceptance_config = Some(json!({
"enabled": true,
"accept_formats": ["claude:messages"],
}));
let provider_catalog = InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider_with_options("provider-shared", false, 0)],
vec![
openai_endpoint,
sample_endpoint_for_provider(
"provider-shared",
"endpoint-claude",
"claude:messages",
),
],
vec![sample_key_for_provider_with_options(
"provider-shared",
"key-shared",
"",
true,
Some(json!(["openai:chat", "claude:messages"])),
None,
)],
);
let data_state = GatewayDataState::with_provider_transport_reader_for_tests(
std::sync::Arc::new(provider_catalog),
"development-key",
);
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data_state);
let auth_snapshot = sample_auth_snapshot();
let cached_cross_format = sample_priority_candidate(
"provider-shared",
"endpoint-openai",
"key-shared",
"openai:chat",
Some(0),
0,
);
remember_scheduler_affinity_for_candidate(
PlannerAppState::new(&state),
Some(&auth_snapshot),
"claude:messages",
"gpt-4.1",
&cached_cross_format,
);
let (ranked, skipped) = resolve_and_rank_local_execution_candidates(
PlannerAppState::new(&state),
vec![
cached_cross_format,
sample_priority_candidate(
"provider-shared",
"endpoint-claude",
"key-shared",
"claude:messages",
Some(0),
0,
),
],
"claude:messages",
"gpt-4.1",
Some(&auth_snapshot),
None,
None,
)
.await;
assert!(skipped.is_empty());
assert_eq!(ranked[0].candidate.endpoint_id, "endpoint-openai");
assert_eq!(
ranked[0]
.ranking
.as_ref()
.and_then(|ranking| ranking.promoted_by),
Some(RANKING_REASON_CACHED_AFFINITY)
);
assert_eq!(
ranked[0]
.ranking
.as_ref()
.and_then(|ranking| ranking.demoted_by),
Some(aether_scheduler_core::RANKING_REASON_CROSS_FORMAT)
);
assert_eq!(ranked[1].candidate.endpoint_id, "endpoint-claude");
}
#[tokio::test]
async fn non_pool_key_affinity_does_not_promote_sibling_key_when_cached_key_is_inactive() {
let provider_catalog = InMemoryProviderCatalogReadRepository::seed(
@@ -5,6 +5,7 @@ mod provider;
pub(crate) use self::provider::{
build_local_stream_plan_and_reports as build_local_same_format_stream_plan_and_reports,
build_local_sync_plan_and_reports as build_local_same_format_sync_plan_and_reports,
maybe_build_local_same_format_provider_decision_payload_for_candidate,
maybe_build_stream_local_same_format_provider_decision_payload,
maybe_build_sync_local_same_format_provider_decision_payload,
resolve_same_format_provider_transport_unsupported_reason_for_trace,
@@ -6,6 +6,7 @@ use crate::ai_pipeline::planner::candidate_metadata::build_request_trace_proxy_v
use crate::ai_pipeline::planner::materialization_policy::{
build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind,
};
use crate::ai_pipeline::planner::passthrough::maybe_build_local_same_format_provider_decision_payload_for_candidate;
use crate::ai_pipeline::planner::payload_metadata::{
build_local_execution_decision_response, LocalExecutionDecisionResponseParts,
};
@@ -17,7 +18,10 @@ use crate::ai_pipeline::planner::CandidateFailureDiagnostic;
use crate::ai_pipeline::transport::{
resolve_transport_execution_timeouts, resolve_transport_tls_profile,
};
use crate::ai_pipeline::{ConversionMode, ExecutionStrategy};
use crate::ai_pipeline::{
api_format_alias_matches, resolve_local_same_format_stream_spec,
resolve_local_same_format_sync_spec, ConversionMode, ExecutionStrategy,
};
use crate::{
append_execution_contract_fields_to_value, append_local_failover_policy_to_value, AppState,
GatewayControlSyncDecisionResponse,
@@ -36,6 +40,29 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
spec: LocalStandardSpec,
) -> Option<GatewayControlSyncDecisionResponse> {
let spec_metadata = local_standard_spec_metadata(spec);
if api_format_alias_matches(
&attempt.eligible.provider_api_format,
spec_metadata.api_format,
) {
let same_format_spec = if spec_metadata.require_streaming {
resolve_local_same_format_stream_spec(spec_metadata.decision_kind)
} else {
resolve_local_same_format_sync_spec(spec_metadata.decision_kind)
};
if let Some(same_format_spec) = same_format_spec {
return maybe_build_local_same_format_provider_decision_payload_for_candidate(
state,
parts,
trace_id,
body_json,
input,
attempt,
same_format_spec,
)
.await;
}
}
let LocalStandardCandidateAttempt {
eligible,
candidate_index,
@@ -243,3 +270,281 @@ pub(super) async fn mark_skipped_local_standard_candidate_with_failure_diagnosti
)
.await;
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use aether_provider_transport::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportSnapshot,
};
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
use serde_json::json;
use super::maybe_build_local_standard_decision_payload_for_candidate;
use crate::ai_pipeline::planner::candidate_materialization::LocalExecutionCandidateAttempt;
use crate::ai_pipeline::planner::candidate_resolution::EligibleLocalExecutionCandidate;
use crate::ai_pipeline::planner::decision_input::LocalRequestedModelDecisionInput;
use crate::ai_pipeline::{
ExecutionRuntimeAuthContext, GatewayAuthApiKeySnapshot, LocalStandardSourceFamily,
LocalStandardSourceMode, LocalStandardSpec,
};
use crate::orchestration::LocalExecutionCandidateMetadata;
fn sample_auth_snapshot() -> GatewayAuthApiKeySnapshot {
GatewayAuthApiKeySnapshot {
user_id: "user-1".to_string(),
username: "alice".to_string(),
email: None,
user_role: "user".to_string(),
user_auth_source: "local".to_string(),
user_is_active: true,
user_is_deleted: false,
user_rate_limit: None,
user_allowed_providers: None,
user_allowed_api_formats: None,
user_allowed_models: None,
api_key_id: "api-key-1".to_string(),
api_key_name: Some("default".to_string()),
api_key_is_active: true,
api_key_is_locked: false,
api_key_is_standalone: false,
api_key_rate_limit: None,
api_key_concurrent_limit: None,
api_key_expires_at_unix_secs: None,
api_key_allowed_providers: None,
api_key_allowed_api_formats: None,
api_key_allowed_models: None,
currently_usable: true,
}
}
fn sample_input() -> LocalRequestedModelDecisionInput {
LocalRequestedModelDecisionInput {
auth_context: ExecutionRuntimeAuthContext {
user_id: "user-1".to_string(),
api_key_id: "api-key-1".to_string(),
username: Some("alice".to_string()),
api_key_name: Some("default".to_string()),
balance_remaining: Some(10.0),
access_allowed: true,
api_key_is_standalone: false,
},
requested_model: "claude-sonnet-4-5".to_string(),
auth_snapshot: sample_auth_snapshot(),
required_capabilities: None,
}
}
fn sample_transport(api_format: &str, endpoint_id: &str) -> GatewayProviderTransportSnapshot {
GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "provider".to_string(),
provider_type: "custom".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: true,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: endpoint_id.to_string(),
provider_id: "provider-1".to_string(),
api_format: api_format.to_string(),
api_family: Some(
api_format
.split_once(':')
.map(|(family, _)| family)
.unwrap_or(api_format)
.to_string(),
),
endpoint_kind: Some("chat".to_string()),
is_active: true,
base_url: "https://api.example.test".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: if api_format == "openai:chat" {
Some(json!({
"enabled": true,
"accept_formats": ["claude:messages"],
}))
} else {
None
},
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "key".to_string(),
auth_type: "api_key".to_string(),
is_active: true,
api_formats: Some(vec![
"claude:messages".to_string(),
"openai:chat".to_string(),
]),
auth_type_by_format: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: Some(json!({
"claude:messages": 1,
"openai:chat": 1,
})),
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
decrypted_api_key: "sk-upstream".to_string(),
decrypted_auth_config: None,
},
}
}
fn sample_candidate(
api_format: &str,
endpoint_id: &str,
) -> SchedulerMinimalCandidateSelectionCandidate {
SchedulerMinimalCandidateSelectionCandidate {
provider_id: "provider-1".to_string(),
provider_name: "provider".to_string(),
provider_type: "custom".to_string(),
provider_priority: 1,
endpoint_id: endpoint_id.to_string(),
endpoint_api_format: api_format.to_string(),
key_id: "key-1".to_string(),
key_name: "key".to_string(),
key_auth_type: "api_key".to_string(),
key_internal_priority: 1,
key_global_priority_for_format: Some(1),
key_capabilities: None,
model_id: format!("model-{endpoint_id}"),
global_model_id: "global-model-1".to_string(),
global_model_name: "claude-sonnet-4-5".to_string(),
selected_provider_model_name: if api_format == "claude:messages" {
"claude-sonnet-4-5-upstream".to_string()
} else {
"gpt-4o-upstream".to_string()
},
mapping_matched_model: None,
}
}
fn sample_attempt(
api_format: &str,
endpoint_id: &str,
candidate_index: u32,
) -> LocalExecutionCandidateAttempt {
LocalExecutionCandidateAttempt {
eligible: EligibleLocalExecutionCandidate {
candidate: sample_candidate(api_format, endpoint_id),
transport: Arc::new(sample_transport(api_format, endpoint_id)),
provider_api_format: api_format.to_string(),
orchestration: LocalExecutionCandidateMetadata::default(),
ranking: None,
},
candidate_index,
retry_index: 0,
candidate_id: format!("candidate-{candidate_index}"),
}
}
fn claude_stream_spec() -> LocalStandardSpec {
LocalStandardSpec {
api_format: "claude:messages",
decision_kind: "claude_chat_stream",
report_kind: "claude_chat_stream_success",
family: LocalStandardSourceFamily::Standard,
mode: LocalStandardSourceMode::Chat,
require_streaming: true,
}
}
#[tokio::test]
async fn standard_family_builds_same_format_candidate_before_cross_format_candidate() {
let state = crate::AppState::new().expect("state should build");
let request = http::Request::builder()
.method("POST")
.uri("/v1/messages?beta=true")
.header(http::header::CONTENT_TYPE, "application/json")
.body(())
.expect("request should build");
let (parts, _) = request.into_parts();
let body_json = json!({
"model": "claude-sonnet-4-5",
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 32,
"stream": true
});
let input = sample_input();
let payload = maybe_build_local_standard_decision_payload_for_candidate(
&state,
&parts,
"trace-standard-same-format-first",
&body_json,
&input,
sample_attempt("claude:messages", "endpoint-claude", 0),
claude_stream_spec(),
)
.await
.expect("same-format candidate should build a standard-family payload");
assert_eq!(payload.endpoint_id.as_deref(), Some("endpoint-claude"));
assert_eq!(
payload.execution_strategy.as_deref(),
Some("local_same_format")
);
assert_eq!(payload.conversion_mode.as_deref(), Some("none"));
assert_eq!(
payload.provider_api_format.as_deref(),
Some("claude:messages")
);
assert_eq!(
payload.client_api_format.as_deref(),
Some("claude:messages")
);
assert_eq!(
payload
.provider_request_body
.as_ref()
.and_then(|body| body.get("model"))
.and_then(serde_json::Value::as_str),
Some("claude-sonnet-4-5-upstream")
);
let cross_format_payload = maybe_build_local_standard_decision_payload_for_candidate(
&state,
&parts,
"trace-standard-same-format-first",
&body_json,
&input,
sample_attempt("openai:chat", "endpoint-openai-chat", 1),
claude_stream_spec(),
)
.await
.expect("cross-format candidate should still build after the same-format candidate");
assert_eq!(
cross_format_payload.endpoint_id.as_deref(),
Some("endpoint-openai-chat")
);
assert_eq!(
cross_format_payload.execution_strategy.as_deref(),
Some("local_cross_format")
);
assert_eq!(
cross_format_payload.conversion_mode.as_deref(),
Some("bidirectional")
);
}
}