mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-11 03:39:49 +08:00
fix scheduler affinity candidate selection
This commit is contained in:
@@ -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")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user