fix(routing): harden routed pool scheduling

This commit is contained in:
elky
2026-07-27 22:06:28 +08:00
parent 550cc36760
commit 4148ab1931
62 changed files with 5510 additions and 850 deletions
+2 -1
View File
@@ -55,7 +55,8 @@ pub(crate) use self::planner::{
maybe_build_sync_decision_payload, maybe_build_sync_plan_payload,
planner_is_matching_stream_request, provider_key_pool_score_id, provider_key_pool_score_scope,
read_candidate_transport_snapshot, record_local_runtime_candidate_skip_reason,
resolve_upstream_is_stream_for_provider, set_local_openai_chat_execution_exhausted_diagnostic,
resolve_tunnel_scheduler_affinity_context, resolve_upstream_is_stream_for_provider,
set_local_openai_chat_execution_exhausted_diagnostic,
set_local_openai_image_execution_exhausted_diagnostic, validate_final_openai_provider_request,
CandidateFailureDiagnostic, CandidateFailureDiagnosticKind, EligibleLocalExecutionCandidate,
GatewayAuthApiKeySnapshot, GatewayProviderTransportSnapshot, LocalExecutionAttemptSource,
@@ -1,6 +1,8 @@
use aether_routing_core::ResolvedRoutingPolicy;
use aether_scheduler_core::{
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session, ClientSessionAffinity,
SchedulerAffinityTarget, SchedulerMinimalCandidateSelectionCandidate,
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope,
ClientSessionAffinity, SchedulerAffinityScope, SchedulerAffinityTarget,
SchedulerMinimalCandidateSelectionCandidate,
};
use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState};
@@ -20,6 +22,7 @@ pub(crate) fn read_cached_scheduler_affinity_target(
client_session_affinity: Option<&ClientSessionAffinity>,
client_api_format: &str,
requested_model: Option<&str>,
routing_policy: Option<&ResolvedRoutingPolicy>,
) -> Option<SchedulerAffinityTarget> {
if !has_explicit_session_affinity(client_session_affinity) {
return None;
@@ -30,12 +33,15 @@ pub(crate) fn read_cached_scheduler_affinity_target(
let api_key_id = auth_snapshot
.map(|snapshot| snapshot.api_key_id.trim())
.filter(|value| !value.is_empty())?;
let cache_key = build_scheduler_affinity_cache_key_for_api_key_id_with_client_session(
api_key_id,
client_api_format,
requested_model,
client_session_affinity,
)?;
let affinity_scope = scheduler_affinity_scope_for_routing_policy(routing_policy);
let cache_key =
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope(
api_key_id,
client_api_format,
requested_model,
client_session_affinity,
affinity_scope.as_ref(),
)?;
state
.app()
@@ -72,6 +78,53 @@ pub(crate) fn remember_scheduler_affinity_for_candidate_at_epoch(
requested_model: &str,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
expected_epoch: Option<u64>,
) {
remember_scheduler_affinity_for_candidate_with_scope_at_epoch(
state,
auth_snapshot,
client_session_affinity,
client_api_format,
requested_model,
candidate,
None,
expected_epoch,
);
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn remember_scheduler_affinity_for_candidate_with_routing_policy_at_epoch(
state: PlannerAppState<'_>,
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>,
client_api_format: &str,
requested_model: &str,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
routing_policy: Option<&ResolvedRoutingPolicy>,
expected_epoch: Option<u64>,
) {
let affinity_scope = scheduler_affinity_scope_for_routing_policy(routing_policy);
remember_scheduler_affinity_for_candidate_with_scope_at_epoch(
state,
auth_snapshot,
client_session_affinity,
client_api_format,
requested_model,
candidate,
affinity_scope.as_ref(),
expected_epoch,
);
}
#[allow(clippy::too_many_arguments)]
fn remember_scheduler_affinity_for_candidate_with_scope_at_epoch(
state: PlannerAppState<'_>,
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>,
client_api_format: &str,
requested_model: &str,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
affinity_scope: Option<&SchedulerAffinityScope>,
expected_epoch: Option<u64>,
) {
if !has_explicit_session_affinity(client_session_affinity) {
return;
@@ -82,12 +135,15 @@ pub(crate) fn remember_scheduler_affinity_for_candidate_at_epoch(
else {
return;
};
let Some(cache_key) = build_scheduler_affinity_cache_key_for_api_key_id_with_client_session(
api_key_id,
client_api_format,
requested_model,
client_session_affinity,
) else {
let Some(cache_key) =
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope(
api_key_id,
client_api_format,
requested_model,
client_session_affinity,
affinity_scope,
)
else {
return;
};
@@ -103,3 +159,15 @@ pub(crate) fn remember_scheduler_affinity_for_candidate_at_epoch(
expected_epoch,
);
}
fn scheduler_affinity_scope_for_routing_policy(
routing_policy: Option<&ResolvedRoutingPolicy>,
) -> Option<SchedulerAffinityScope> {
let policy = routing_policy?;
let group_id = policy
.group_id
.as_deref()
.map(str::trim)
.filter(|group_id| !group_id.is_empty())?;
Some(SchedulerAffinityScope::new(group_id, policy.group_version))
}
@@ -24,7 +24,7 @@ use tokio::time::Instant;
use tracing::warn;
use uuid::Uuid;
use crate::ai_serving::planner::candidate_affinity_cache::remember_scheduler_affinity_for_candidate_at_epoch;
use crate::ai_serving::planner::candidate_affinity_cache::remember_scheduler_affinity_for_candidate_with_routing_policy_at_epoch;
use crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy;
use crate::ai_serving::planner::candidate_resolution::{
resolve_and_rank_logical_local_execution_candidates, EligibleLocalExecutionCandidate,
@@ -421,6 +421,7 @@ where
self.client_session_affinity,
self.client_api_format,
self.requested_model,
self.routing_policy,
candidates,
);
}
@@ -691,6 +692,7 @@ where
client_session_affinity,
client_api_format,
requested_model,
routing_policy,
&candidates,
);
}
@@ -983,15 +985,7 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
"candidate_page_load",
page_started_at.elapsed().as_millis() as u64,
);
if matches!(error, GatewayError::AdmissionTimeout { .. }) {
return Err(error);
}
warn!(
trace_id = %self.trace_id,
error = ?error,
"gateway lazy requested-model candidate page read failed"
);
return Ok(false);
return Err(error);
}
};
observe_gateway_stage_ms(
@@ -1035,6 +1029,7 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
self.client_session_affinity.as_ref(),
&self.client_api_format,
Some(&self.requested_model),
self.routing_policy.as_ref(),
&candidates,
);
self.remembered_affinity = true;
@@ -1259,6 +1254,7 @@ pub(crate) fn remember_first_local_candidate_affinity(
client_session_affinity: Option<&ClientSessionAffinity>,
client_api_format: &str,
requested_model: Option<&str>,
routing_policy: Option<&ResolvedRoutingPolicy>,
candidates: &[EligibleLocalExecutionCandidate],
) {
let Some(first_candidate) = candidates.first() else {
@@ -1268,13 +1264,14 @@ pub(crate) fn remember_first_local_candidate_affinity(
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(first_candidate.candidate.global_model_name.as_str());
remember_scheduler_affinity_for_candidate_at_epoch(
remember_scheduler_affinity_for_candidate_with_routing_policy_at_epoch(
state,
auth_snapshot,
client_session_affinity,
client_api_format,
affinity_requested_model,
&first_candidate.candidate,
routing_policy,
first_candidate.orchestration.scheduler_affinity_epoch,
);
}
@@ -66,6 +66,7 @@ impl AiCandidateRankingPort for GatewayLocalCandidateRankingPort<'_> {
self.client_session_affinity,
normalized_client_api_format,
affinity_requested_model,
self.routing_policy,
))
}
@@ -221,17 +222,21 @@ fn routing_overlaid_candidate(
let mut overlaid = candidate.clone();
overlaid.provider_priority = policy
.ranking_overlay
.provider_priority_or_unspecified(candidate.provider_id.as_str());
.provider_priority(candidate.provider_id.as_str(), candidate.provider_priority);
let overlaid_key_priority = match kind {
LocalExecutionCandidateKind::SingleKey => policy
.ranking_overlay
.key_priority_or_unspecified(candidate.key_id.as_str()),
.key_priority_overrides
.get(candidate.key_id.as_str()),
LocalExecutionCandidateKind::PoolGroup => policy
.ranking_overlay
.pool_priority_or_unspecified(candidate.provider_id.as_str()),
.pool_priority_overrides
.get(candidate.provider_id.as_str()),
};
overlaid.key_internal_priority = overlaid_key_priority;
overlaid.key_global_priority_for_format = Some(overlaid_key_priority);
if let Some(overlaid_key_priority) = overlaid_key_priority.copied() {
overlaid.key_internal_priority = overlaid_key_priority;
overlaid.key_global_priority_for_format = Some(overlaid_key_priority);
}
overlaid
}
@@ -354,7 +359,7 @@ mod tests {
}
#[test]
fn routing_policy_priorities_do_not_fall_back_to_candidate_priorities() {
fn routing_policy_priorities_fall_back_to_candidate_priorities() {
let mut candidate = sample_candidate("endpoint-1", "key-1");
candidate.provider_priority = 7;
candidate.key_internal_priority = 3;
@@ -380,18 +385,9 @@ mod tests {
&candidate,
);
assert_eq!(
overlaid.provider_priority,
aether_routing_core::ROUTING_PRIORITY_UNSPECIFIED
);
assert_eq!(
overlaid.key_internal_priority,
aether_routing_core::ROUTING_PRIORITY_UNSPECIFIED
);
assert_eq!(
overlaid.key_global_priority_for_format,
Some(aether_routing_core::ROUTING_PRIORITY_UNSPECIFIED)
);
assert_eq!(overlaid.provider_priority, 7);
assert_eq!(overlaid.key_internal_priority, 3);
assert_eq!(overlaid.key_global_priority_for_format, Some(2));
}
#[test]
@@ -29,9 +29,10 @@ use crate::cache::{
};
use crate::clock::request_distribution_seed;
use crate::data::candidate_selection::{
read_requested_model_rows_fast_path_page, requested_model_candidate_names,
MinimalCandidateSelectionRowSource, RequestedModelCandidateRowsPage,
REQUESTED_MODEL_CANDIDATE_PAGE_SIZE, REQUESTED_MODEL_MAX_SCANNED_ROWS,
read_api_format_rows_fallback_page, read_requested_model_rows_fast_path_page,
requested_model_candidate_names, MinimalCandidateSelectionRowSource,
RequestedModelCandidateRowsPage, REQUESTED_MODEL_CANDIDATE_PAGE_SIZE,
REQUESTED_MODEL_MAX_SCANNED_ROWS,
};
use crate::scheduler::candidate::SchedulerSkippedCandidate;
use crate::scheduler::config::{SchedulerOrderingConfig, SchedulerSchedulingMode};
@@ -120,7 +121,10 @@ impl AiCandidatePreselectionPort for GatewayLocalCandidatePreselectionPort<'_> {
self.require_streaming,
self.required_capabilities,
auth_snapshot,
self.client_session_affinity,
self.routing_policy
.is_none()
.then_some(self.client_session_affinity)
.flatten(),
self.ranking_seed,
false,
self.request_operation,
@@ -320,7 +324,8 @@ pub(crate) struct LocalCandidatePreselectionPageCursor<'a> {
requested_name_offsets: BTreeMap<String, u32>,
scanned_rows_by_format: BTreeMap<String, u32>,
resolved_global_model_names: BTreeMap<String, String>,
fallback_scanned_api_formats: BTreeSet<String>,
fallback_offsets: BTreeMap<String, u32>,
fallback_scan_epoch: u32,
exhausted_api_formats: BTreeSet<String>,
seen_candidate_keys: BTreeSet<String>,
}
@@ -402,7 +407,8 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
requested_name_offsets: BTreeMap::new(),
scanned_rows_by_format: BTreeMap::new(),
resolved_global_model_names: BTreeMap::new(),
fallback_scanned_api_formats: BTreeSet::new(),
fallback_offsets: BTreeMap::new(),
fallback_scan_epoch: 0,
exhausted_api_formats: BTreeSet::new(),
seen_candidate_keys: BTreeSet::new(),
}
@@ -421,13 +427,35 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
> {
if !self.priority_page_emitted {
self.priority_page_emitted = true;
let priority_page = self.cached_next_priority_page().await?;
let mut priority_page = self.cached_next_priority_page().await?;
if self.routing_policy.is_some() {
while let Some(mut page) = self.next_page_after_priority().await? {
priority_page.candidates.append(&mut page.candidates);
priority_page
.skipped_candidates
.append(&mut page.skipped_candidates);
}
}
if !priority_page.candidates.is_empty() || !priority_page.skipped_candidates.is_empty()
{
return Ok(Some(priority_page));
}
}
self.next_page_after_priority().await
}
async fn next_page_after_priority(
&mut self,
) -> Result<
Option<
AiCandidatePreselectionOutcome<
SchedulerMinimalCandidateSelectionCandidate,
SkippedLocalExecutionCandidate,
>,
>,
GatewayError,
> {
// Deferred pages and formats already proven exhausted require no planning
// permit. This is the common second-target path for a single-candidate
// model, so keep it entirely in memory before joining the shared gate.
@@ -477,7 +505,8 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
self.requested_name_offsets.clear();
self.scanned_rows_by_format.clear();
self.resolved_global_model_names.clear();
self.fallback_scanned_api_formats.clear();
self.fallback_offsets.clear();
self.fallback_scan_epoch = self.fallback_scan_epoch.wrapping_add(1);
self.exhausted_api_formats.clear();
self.seen_candidate_keys.clear();
self.priority_page_emitted = false;
@@ -967,6 +996,63 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
}
}
async fn read_api_format_rows_fallback_page_cached(
&self,
normalized_api_format: &str,
offset: u32,
limit: u32,
) -> Result<RequestedModelCandidateRowsPage, GatewayError> {
let key = CandidateRowPageCacheKey::for_api_format_fallback(
normalized_api_format,
offset,
limit,
self.fallback_scan_epoch,
);
let cache = self.state.app().candidate_row_page_cache.clone();
let ttl = candidate_page_cache_ttl_from_env();
let stale_ttl = candidate_page_cache_stale_ttl(ttl);
let cached = cache
.get_or_load_once_stale_while_refreshing(
key,
ttl,
stale_ttl,
|| async {
let page = read_api_format_rows_fallback_page(
self.state.app().data.as_ref(),
normalized_api_format,
offset,
limit,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
Ok::<_, GatewayError>(Some(Arc::new(page)))
},
CacheLoadObserver::new()
.on_hit(record_candidate_row_page_cache_hit)
.on_miss(record_candidate_row_page_cache_miss)
.on_load(record_candidate_row_page_cache_load)
.on_follower_wait(record_candidate_row_page_cache_follower_wait),
)
.await?;
match cached {
Some(page) => {
if page.rows.is_empty() {
record_candidate_row_page_cache_none();
}
Ok(page.as_ref().clone())
}
None => {
record_candidate_row_page_cache_none();
Ok(RequestedModelCandidateRowsPage {
rows: Vec::new(),
scanned_rows: 0,
end_of_requested_name: true,
})
}
}
}
async fn next_fallback_page_for_api_format(
&mut self,
candidate_api_format: &str,
@@ -980,43 +1066,67 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
>,
GatewayError,
> {
if self
.fallback_scanned_api_formats
.contains(normalized_api_format)
{
self.exhausted_api_formats
.insert(normalized_api_format.to_string());
return Ok(None);
}
let routing_model = self.routing_model(candidate_api_format).to_string();
let rows = self
.state
.app()
.data
.read_minimal_candidate_selection_rows_for_api_format(normalized_api_format)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
.into_iter()
.filter(|row| {
row_supports_requested_model_with_model_directives_and_request_operation(
row,
&routing_model,
normalized_api_format,
false,
self.request_operation.as_deref(),
)
})
.collect::<Vec<_>>();
loop {
let scanned = *self
.scanned_rows_by_format
.get(normalized_api_format)
.unwrap_or(&0);
let remaining = REQUESTED_MODEL_MAX_SCANNED_ROWS.saturating_sub(scanned);
if remaining == 0 {
self.exhausted_api_formats
.insert(normalized_api_format.to_string());
return Ok(None);
}
let limit = REQUESTED_MODEL_CANDIDATE_PAGE_SIZE.min(remaining);
let offset = *self
.fallback_offsets
.get(normalized_api_format)
.unwrap_or(&0);
let page = self
.read_api_format_rows_fallback_page_cached(normalized_api_format, offset, limit)
.await?;
let page_scanned = page.scanned_rows.min(limit);
let end_of_format = page.end_of_requested_name || page_scanned < limit;
self.fallback_offsets.insert(
normalized_api_format.to_string(),
offset.saturating_add(page_scanned),
);
let total_scanned = scanned.saturating_add(page_scanned);
self.scanned_rows_by_format
.insert(normalized_api_format.to_string(), total_scanned);
if end_of_format || total_scanned >= REQUESTED_MODEL_MAX_SCANNED_ROWS {
self.exhausted_api_formats
.insert(normalized_api_format.to_string());
}
if page_scanned == 0 {
return Ok(None);
}
let outcome = self
.build_page_outcome_from_rows(candidate_api_format, normalized_api_format, rows)
.await?;
self.fallback_scanned_api_formats
.insert(normalized_api_format.to_string());
self.exhausted_api_formats
.insert(normalized_api_format.to_string());
Ok(outcome)
let rows = page
.rows
.into_iter()
.take(page_scanned as usize)
.filter(|row| {
row_supports_requested_model_with_model_directives_and_request_operation(
row,
&routing_model,
normalized_api_format,
false,
self.request_operation.as_deref(),
)
})
.collect::<Vec<_>>();
if let Some(outcome) = self
.build_page_outcome_from_rows(candidate_api_format, normalized_api_format, rows)
.await?
{
return Ok(Some(outcome));
}
if self.exhausted_api_formats.contains(normalized_api_format) {
return Ok(None);
}
}
}
fn api_format_is_exhausted(&self, candidate_api_format: &str) -> bool {
@@ -1125,7 +1235,10 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
candidates,
self.required_capabilities.as_ref(),
auth_snapshot,
self.client_session_affinity.as_ref(),
self.routing_policy
.is_none()
.then_some(self.client_session_affinity.as_ref())
.flatten(),
self.ranking_seed,
)
.await?;
@@ -1313,22 +1426,116 @@ mod tests {
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::DataLayerError;
use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use async_trait::async_trait;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::sync::{Arc, Mutex};
#[derive(Default)]
struct EmptyFallbackCountingRepository {
fallback_reads: AtomicUsize,
}
struct PagedFallbackRepository {
total_rows: u32,
page_queries: Mutex<Vec<StoredApiFormatCandidateRowsQuery>>,
}
impl PagedFallbackRepository {
fn new(total_rows: u32) -> Self {
Self {
total_rows,
page_queries: Mutex::new(Vec::new()),
}
}
fn page_queries(&self) -> Vec<StoredApiFormatCandidateRowsQuery> {
self.page_queries
.lock()
.expect("fallback query lock")
.clone()
}
}
#[async_trait]
impl MinimalCandidateSelectionReadRepository for PagedFallbackRepository {
async fn list_for_exact_api_format(
&self,
_api_format: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
panic!("routing fallback must not use the unbounded API-format query")
}
async fn list_for_exact_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.page_queries
.lock()
.expect("fallback query lock")
.push(query.clone());
if normalize_api_format(&query.api_format) != "openai:chat" {
return Ok(Vec::new());
}
let end = query
.offset
.saturating_add(query.limit)
.min(self.total_rows);
Ok((query.offset..end)
.map(|index| {
standard_candidate_row(
format!("fallback-provider-{index:04}").as_str(),
"openai:chat",
i32::try_from(index).expect("test provider priority should fit"),
)
})
.collect())
}
async fn list_for_exact_api_format_and_global_model(
&self,
_api_format: &str,
_global_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Ok(Vec::new())
}
async fn list_for_exact_api_format_and_requested_model(
&self,
_api_format: &str,
_requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Ok(Vec::new())
}
async fn list_for_exact_api_format_and_requested_model_page(
&self,
_query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Ok(Vec::new())
}
async fn list_pool_key_rows_for_group(
&self,
_query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Ok(Vec::new())
}
async fn list_pool_key_rows_for_group_key_ids(
&self,
_query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Ok(Vec::new())
}
}
impl EmptyFallbackCountingRepository {
fn fallback_reads(&self) -> usize {
self.fallback_reads.load(Ordering::Acquire)
@@ -1571,6 +1778,155 @@ mod tests {
.is_none());
}
#[tokio::test]
async fn routing_policy_collects_candidate_pages_before_final_ranking() {
let rows = (0..300)
.map(|index| {
standard_candidate_row(
format!("provider-{index:03}").as_str(),
"openai:chat",
index,
)
})
.collect::<Vec<_>>();
let repository: Arc<dyn MinimalCandidateSelectionReadRepository> =
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(rows));
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 routing_policy = ResolvedRoutingPolicy {
group_id: Some("routing-group-1".to_string()),
group_version: Some(1),
selection_source: "test".to_string(),
requested_model: "gpt-5".to_string(),
resolved_model: "gpt-5".to_string(),
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder,
keep_priority_on_conversion: false,
ranking_overlay: Default::default(),
mutation_plan: Default::default(),
pool_policy_overrides: Default::default(),
matched_rules: Vec::new(),
};
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 candidates = cursor
.next_page()
.await
.expect("routing candidate scan should succeed")
.expect("routing candidates should be present")
.candidates;
assert_eq!(candidates.len(), 300);
assert!(cursor
.next_page()
.await
.expect("routing scan should be exhausted")
.is_none());
}
#[tokio::test]
async fn routing_fallback_uses_bounded_api_format_pages() {
let repository = Arc::new(PagedFallbackRepository::new(
REQUESTED_MODEL_MAX_SCANNED_ROWS + REQUESTED_MODEL_CANDIDATE_PAGE_SIZE,
));
let app = AppState::new()
.expect("gateway state should build")
.with_data_state_for_tests(
GatewayDataState::with_minimal_candidate_selection_reader_for_tests(
repository.clone(),
),
);
let auth_snapshot = unrestricted_auth_snapshot();
let model_directive_policy =
crate::system_features::ModelDirectivePolicySnapshot::load(&app).await;
let routing_policy = ResolvedRoutingPolicy {
group_id: Some("routing-group-fallback".to_string()),
group_version: Some(1),
selection_source: "test".to_string(),
requested_model: "gpt-5".to_string(),
resolved_model: "gpt-5".to_string(),
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
scheduling_mode: aether_routing_core::RoutingSchedulingMode::FixedOrder,
keep_priority_on_conversion: false,
ranking_overlay: Default::default(),
mutation_plan: Default::default(),
pool_policy_overrides: Default::default(),
matched_rules: Vec::new(),
};
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 candidates = cursor
.next_page()
.await
.expect("routing fallback scan should succeed")
.expect("routing fallback candidates should be present")
.candidates;
assert_eq!(candidates.len(), REQUESTED_MODEL_MAX_SCANNED_ROWS as usize);
assert!(cursor
.next_page()
.await
.expect("bounded routing fallback should be exhausted")
.is_none());
let page_queries = repository
.page_queries()
.into_iter()
.filter(|query| normalize_api_format(&query.api_format) == "openai:chat")
.collect::<Vec<_>>();
assert_eq!(
page_queries.len(),
(REQUESTED_MODEL_MAX_SCANNED_ROWS / REQUESTED_MODEL_CANDIDATE_PAGE_SIZE) as usize
);
for (index, query) in page_queries.iter().enumerate() {
assert_eq!(query.limit, REQUESTED_MODEL_CANDIDATE_PAGE_SIZE);
assert_eq!(
query.offset,
u32::try_from(index).expect("page index should fit")
* REQUESTED_MODEL_CANDIDATE_PAGE_SIZE
);
}
}
#[tokio::test]
async fn priority_page_cache_requires_fixed_order_or_explicit_affinity() {
let repository: Arc<dyn MinimalCandidateSelectionReadRepository> =
@@ -11,7 +11,6 @@ use async_trait::async_trait;
use http::StatusCode;
use http::{HeaderMap, HeaderName, HeaderValue};
use serde_json::{json, Value};
use tracing::warn;
use crate::ai_serving::planner::common::extract_standard_requested_model;
use crate::ai_serving::{
@@ -408,29 +407,28 @@ pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
let principal_context_required = if explicit_group.is_some() {
true
} else {
!matches!(repository.has_any_routing_group_binding().await, Ok(false))
repository
.has_any_routing_group_binding()
.await
.map_err(|error| {
routing_selection_error(GatewayRoutingSelectionError::Repository(
error.to_string(),
))
})?
};
let user_group_ids = if principal_context_required {
let user_groups_lookup_started_at = std::time::Instant::now();
let user_group_ids = match state
let user_groups = state
.list_user_groups_for_user(&input.auth_context.user_id)
.await
{
Ok(groups) => groups.into_iter().map(|group| group.id).collect::<Vec<_>>(),
Err(error) => {
warn!(
user_id = %input.auth_context.user_id,
error = ?error,
"gateway routing profile user group lookup failed"
);
Vec::new()
}
};
.await;
observe_gateway_stage_ms(
"routing_user_groups_lookup",
user_groups_lookup_started_at.elapsed().as_millis() as u64,
);
user_group_ids
user_groups?
.into_iter()
.map(|group| group.id)
.collect::<Vec<_>>()
} else {
Vec::new()
};
@@ -735,9 +733,14 @@ pub(crate) async fn resolve_local_authenticated_decision_input(
}
fn routing_selection_error(error: GatewayRoutingSelectionError) -> GatewayError {
GatewayError::Client {
status: StatusCode::FORBIDDEN,
message: error.to_string(),
match error {
GatewayRoutingSelectionError::Repository(message) => {
GatewayError::Internal(format!("routing group repository lookup failed: {message}"))
}
error => GatewayError::Client {
status: StatusCode::FORBIDDEN,
message: error.to_string(),
},
}
}
@@ -1006,6 +1009,21 @@ mod tests {
assert!(first.contains("groups=team-1"));
}
#[test]
fn routing_repository_failure_maps_to_internal_gateway_error() {
let error = routing_selection_error(GatewayRoutingSelectionError::Repository(
"sql error: database unavailable".to_string(),
));
match error {
GatewayError::Internal(message) => {
assert!(message.contains("routing group repository lookup failed"));
assert!(message.contains("database unavailable"));
}
other => panic!("unexpected routing repository error mapping: {other:?}"),
}
}
#[tokio::test]
async fn explicit_routing_attachment_authorizes_and_caches_per_principal() {
let repository = Arc::new(InMemoryRoutingGroupRepository::default());
@@ -93,6 +93,71 @@ pub(crate) use aether_ai_serving::{
CandidateFailureDiagnostic, CandidateFailureDiagnosticKind,
};
pub(crate) struct ResolvedTunnelSchedulerAffinityContext {
pub(crate) requested_model: String,
pub(crate) client_session_affinity: Option<aether_scheduler_core::ClientSessionAffinity>,
pub(crate) policy_context: Option<crate::scheduler::affinity::SchedulerAffinityPolicyContext>,
pub(crate) routing_overlay: Option<aether_routing_core::RankingOverlay>,
}
pub(crate) async fn resolve_tunnel_scheduler_affinity_context(
state: &AppState,
parts: &http::request::Parts,
decision: &GatewayControlDecision,
requested_model: String,
body_json: &serde_json::Value,
client_api_format: &str,
) -> Result<Option<ResolvedTunnelSchedulerAffinityContext>, GatewayError> {
let Some(auth_context) = decision.auth_context.as_ref() else {
return Ok(None);
};
let execution_auth_context =
crate::ai_serving::build_execution_runtime_auth_context(auth_context);
let Some(auth_snapshot) = state
.read_cached_auth_api_key_snapshot(
&execution_auth_context.user_id,
&execution_auth_context.api_key_id,
crate::clock::current_unix_secs(),
)
.await?
else {
return Ok(None);
};
let resolved_auth_input = decision_input::ResolvedLocalDecisionAuthInput {
auth_context: execution_auth_context,
auth_snapshot,
required_capabilities: None,
model_directive_policy: decision.model_directive_policy.clone(),
};
let mut input = decision_input::build_local_requested_model_decision_input(
resolved_auth_input,
requested_model,
);
decision_input::attach_routing_policy_to_local_requested_model_input(
state,
parts,
&mut input,
body_json,
client_api_format,
)
.await?;
let policy_context = input
.routing_policy
.as_ref()
.map(crate::scheduler::affinity::SchedulerAffinityPolicyContext::from_routing_policy);
let routing_overlay = input
.routing_policy
.as_ref()
.map(|policy| policy.ranking_overlay.clone());
Ok(Some(ResolvedTunnelSchedulerAffinityContext {
requested_model: input.requested_model,
client_session_affinity: input.client_session_affinity,
policy_context,
routing_overlay,
}))
}
pub(crate) async fn maybe_build_sync_decision_payload(
state: &AppState,
parts: &http::request::Parts,
@@ -199,6 +199,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
original_request_body_json,
original_request_body_base64: None,
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
client_requested_stream: body_json
.get("stream")
@@ -6,6 +6,7 @@ use aether_ai_serving::{
provider_stream_event_api_format_for_provider_type as ai_provider_stream_event_api_format_for_provider_type,
AiExecutionReportContextParts, AiRequestOrigin,
};
use aether_routing_core::ResolvedRoutingPolicy;
use aether_runtime_state::RuntimeLockLease;
use aether_scheduler_core::{ClientSessionAffinity, SchedulerRankingOutcome};
use serde_json::{Map, Value};
@@ -20,8 +21,9 @@ use crate::client_session_affinity::{
};
use crate::orchestration::{
insert_pool_key_lease_report_context_fields, ExecutionAttemptIdentity,
SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD,
ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD, SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD,
};
use crate::scheduler::affinity::insert_scheduler_affinity_policy_report_context_field;
pub(crate) struct LocalExecutionReportContextParts<'a> {
pub(crate) auth_context: &'a ExecutionRuntimeAuthContext,
@@ -55,6 +57,7 @@ pub(crate) struct LocalExecutionReportContextParts<'a> {
pub(crate) original_request_body_json: Option<&'a Value>,
pub(crate) original_request_body_base64: Option<&'a str>,
pub(crate) client_session_affinity: Option<&'a ClientSessionAffinity>,
pub(crate) routing_policy: Option<&'a ResolvedRoutingPolicy>,
pub(crate) scheduler_affinity_epoch: Option<u64>,
pub(crate) client_requested_stream: bool,
pub(crate) upstream_is_stream: bool,
@@ -105,6 +108,16 @@ pub(crate) fn build_local_execution_report_context(
merge_incoming_tls_fingerprint(&mut extra_fields, incoming_tls);
}
insert_pool_key_lease_report_context_fields(&mut extra_fields, parts.pool_key_lease);
insert_scheduler_affinity_policy_report_context_field(&mut extra_fields, parts.routing_policy);
if let Some(override_policy) = parts
.routing_policy
.and_then(|policy| policy.pool_policy_overrides.get(parts.provider_id))
.filter(|override_policy| !override_policy.scheduling_presets.is_empty())
{
if let Ok(value) = serde_json::to_value(override_policy) {
extra_fields.insert(ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD.to_string(), value);
}
}
if let Some(epoch) = parts.scheduler_affinity_epoch {
extra_fields.insert(
SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD.to_string(),
@@ -315,6 +328,7 @@ mod tests {
original_request_body_json: Some(&json!({"model": "gpt-5"})),
original_request_body_base64: None,
client_session_affinity: Some(&client_session_affinity),
routing_policy: None,
scheduler_affinity_epoch: None,
client_requested_stream: false,
upstream_is_stream: false,
@@ -397,6 +411,7 @@ mod tests {
})),
original_request_body_base64: None,
client_session_affinity: None,
routing_policy: None,
scheduler_affinity_epoch: None,
client_requested_stream: false,
upstream_is_stream: true,
@@ -463,6 +478,7 @@ mod tests {
original_request_body_json: Some(&json!({"model": "gpt-5"})),
original_request_body_base64: None,
client_session_affinity: None,
routing_policy: None,
scheduler_affinity_epoch: None,
client_requested_stream: false,
upstream_is_stream: false,
@@ -107,6 +107,7 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
original_request_body_json: Some(body_json),
original_request_body_base64: resolved.provider_request_body_base64.as_deref(),
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
client_requested_stream: spec_metadata.require_streaming,
upstream_is_stream: spec_metadata.require_streaming,
@@ -120,6 +120,7 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
original_request_body_json: Some(body_json),
original_request_body_base64: body_base64,
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
client_requested_stream: spec_metadata.require_streaming,
upstream_is_stream,
@@ -88,6 +88,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
original_request_body_json: Some(body_json),
original_request_body_base64: None,
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
client_requested_stream: false,
upstream_is_stream: false,
@@ -140,6 +140,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
original_request_body_json,
original_request_body_base64: None,
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
client_requested_stream: body_json
.get("stream")
@@ -193,6 +193,7 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
original_request_body_json,
original_request_body_base64: None,
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
client_requested_stream: body_json
.get("stream")
@@ -142,6 +142,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
original_request_body_json,
original_request_body_base64: None,
client_session_affinity: input.client_session_affinity.as_ref(),
routing_policy: input.routing_policy.as_ref(),
scheduler_affinity_epoch: eligible.orchestration.scheduler_affinity_epoch,
client_requested_stream: body_json
.get("stream")
+32 -6
View File
@@ -72,11 +72,21 @@ struct CandidatePageCacheMetrics {
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub(crate) struct CandidateRowPageCacheKey {
api_format: String,
requested_model_name: String,
requested_name: String,
kind: CandidateRowPageCacheKind,
offset: u32,
limit: u32,
enable_model_directives: bool,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
enum CandidateRowPageCacheKind {
RequestedModel {
requested_model_name: String,
requested_name: String,
enable_model_directives: bool,
},
ApiFormatFallback {
scan_epoch: u32,
},
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
@@ -119,11 +129,27 @@ impl CandidateRowPageCacheKey {
) -> Self {
Self {
api_format: normalize_api_format(api_format),
requested_model_name: normalize_text_key(requested_model_name),
requested_name: normalize_text_key(requested_name),
kind: CandidateRowPageCacheKind::RequestedModel {
requested_model_name: normalize_text_key(requested_model_name),
requested_name: normalize_text_key(requested_name),
enable_model_directives,
},
offset,
limit,
}
}
pub(crate) fn for_api_format_fallback(
api_format: &str,
offset: u32,
limit: u32,
scan_epoch: u32,
) -> Self {
Self {
api_format: normalize_api_format(api_format),
kind: CandidateRowPageCacheKind::ApiFormatFallback { scan_epoch },
offset,
limit,
enable_model_directives,
}
}
}
@@ -1,7 +1,7 @@
use aether_data::DataLayerError;
use aether_data_contracts::repository::candidate_selection::{
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
StoredRequestedModelCandidateRowsQuery,
StoredApiFormatCandidateRowsQuery, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
};
use aether_scheduler_core::{
auth_constraints_allow_api_format, collect_global_model_names_for_required_capability,
@@ -39,6 +39,19 @@ pub(crate) trait MinimalCandidateSelectionRowSource {
api_format: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>;
async fn read_minimal_candidate_selection_rows_for_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Ok(self
.read_minimal_candidate_selection_rows_for_api_format(&query.api_format)
.await?
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
}
async fn read_pool_key_candidate_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
@@ -231,6 +244,30 @@ pub(crate) async fn read_requested_model_rows_fast_path_page(
})
}
pub(crate) async fn read_api_format_rows_fallback_page(
state: &(impl MinimalCandidateSelectionRowSource + Sync),
api_format: &str,
offset: u32,
limit: u32,
) -> Result<RequestedModelCandidateRowsPage, DataLayerError> {
let limit = limit.max(1);
let rows = state
.read_minimal_candidate_selection_rows_for_api_format_page(
&StoredApiFormatCandidateRowsQuery {
api_format: api_format.to_string(),
offset,
limit,
},
)
.await?;
let scanned_rows = rows.len() as u32;
Ok(RequestedModelCandidateRowsPage {
rows,
scanned_rows,
end_of_requested_name: scanned_rows < limit,
})
}
pub(crate) async fn enumerate_minimal_candidate_selection_with_required_capabilities(
state: &(impl MinimalCandidateSelectionRowSource + Sync),
api_format: &str,
@@ -7,9 +7,10 @@ use std::time::Duration;
use aether_cache::ExpiringMap;
use aether_data::DataLayerError;
use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredRequestedModelCandidateRowsQuery,
};
use async_trait::async_trait;
use tokio::sync::{Notify, OwnedSemaphorePermit, Semaphore};
@@ -389,6 +390,19 @@ impl MinimalCandidateSelectionReadRepository for CachedMinimalCandidateSelection
.await
}
async fn list_for_exact_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let key = CandidateSelectionCacheKey::ApiFormatPage {
api_format: normalize_api_format_key(&query.api_format),
offset: query.offset,
limit: query.limit,
};
self.get_or_load(key, || self.inner.list_for_exact_api_format_page(query))
.await
}
async fn list_for_exact_api_format_and_global_model(
&self,
api_format: &str,
@@ -480,6 +494,11 @@ enum CandidateSelectionCacheKey {
ApiFormat {
api_format: String,
},
ApiFormatPage {
api_format: String,
offset: u32,
limit: u32,
},
ApiFormatAndGlobalModel {
api_format: String,
global_model_name: String,
@@ -1044,6 +1063,36 @@ mod tests {
assert!(cache.inflight.lock().unwrap().is_empty());
}
#[tokio::test]
async fn candidate_selection_cache_keys_api_format_pages_by_offset() {
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
let cache = CachedMinimalCandidateSelectionReadRepository::new(inner.clone());
let first = StoredApiFormatCandidateRowsQuery {
api_format: "openai:chat".to_string(),
offset: 0,
limit: 256,
};
let second = StoredApiFormatCandidateRowsQuery {
offset: 256,
..first.clone()
};
cache
.list_for_exact_api_format_page(&first)
.await
.expect("first page should load");
cache
.list_for_exact_api_format_page(&first)
.await
.expect("first page should be cached");
cache
.list_for_exact_api_format_page(&second)
.await
.expect("second page should load independently");
assert_eq!(inner.calls(), 2);
}
#[tokio::test]
async fn candidate_selection_cache_releases_inflight_when_leader_is_cancelled() {
let inner = Arc::new(FirstLoadPendingThenFastRepository::new());
@@ -176,6 +176,14 @@ impl MinimalCandidateSelectionRowSource for GatewayDataState {
.await
}
async fn read_minimal_candidate_selection_rows_for_api_format_page(
&self,
query: &aether_data_contracts::repository::candidate_selection::StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.list_minimal_candidate_selection_rows_for_api_format_page(query)
.await
}
async fn read_pool_key_candidate_rows_for_group(
&self,
query: &aether_data_contracts::repository::candidate_selection::StoredPoolKeyCandidateRowsQuery,
+3 -3
View File
@@ -97,9 +97,9 @@ use aether_data_contracts::repository::billing::{
UserPlanEntitlementRecord,
};
use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredRequestedModelCandidateRowsQuery,
MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
};
use aether_data_contracts::repository::candidates::{
PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateReadRepository,
+23 -5
View File
@@ -2,11 +2,12 @@ use super::{
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
DataLayerError, GatewayDataState, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery,
PublicGlobalModelQuery, StoredAdminGlobalModel, StoredAdminGlobalModelPage,
StoredAdminProviderModel, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredProviderActiveGlobalModel, StoredProviderModelStats, StoredPublicCatalogModel,
StoredPublicGlobalModel, StoredPublicGlobalModelPage, StoredRequestedModelCandidateRowsQuery,
UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
StoredAdminProviderModel, StoredApiFormatCandidateRowsQuery,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery, StoredProviderActiveGlobalModel, StoredProviderModelStats,
StoredPublicCatalogModel, StoredPublicGlobalModel, StoredPublicGlobalModelPage,
StoredRequestedModelCandidateRowsQuery, UpdateAdminGlobalModelRecord,
UpsertAdminProviderModelRecord,
};
impl GatewayDataState {
@@ -98,6 +99,23 @@ impl GatewayDataState {
.await
}
pub(crate) async fn list_minimal_candidate_selection_rows_for_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
crate::request_diagnostics::observe_db_operation(
"candidate_selection",
self.database_pool_summary(),
async {
match &self.minimal_candidate_selection_reader {
Some(repository) => repository.list_for_exact_api_format_page(query).await,
None => Ok(Vec::new()),
}
},
)
.await
}
pub(crate) async fn list_pool_key_candidate_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
File diff suppressed because it is too large Load Diff
@@ -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(
+105 -27
View File
@@ -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()
@@ -35,6 +35,7 @@ pub(crate) struct LocalExecutionCandidateMetadata {
}
pub(crate) const SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD: &str = "scheduler_affinity_epoch";
pub(crate) const ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD: &str = "routing_pool_policy_override";
pub(crate) const POOL_KEY_LEASE_KEY_REPORT_FIELD: &str = "pool_key_lease_key";
pub(crate) const POOL_KEY_LEASE_OWNER_REPORT_FIELD: &str = "pool_key_lease_owner";
pub(crate) const POOL_KEY_LEASE_TOKEN_REPORT_FIELD: &str = "pool_key_lease_token";
+257 -12
View File
@@ -13,8 +13,9 @@ use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
ProviderCatalogKeyHealthStateUpdate,
};
use aether_routing_core::RoutingPoolPolicyOverride;
use aether_scheduler_core::{
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session,
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope,
count_recent_rpm_requests_for_provider_key, ClientSessionAffinity, SchedulerAffinityTarget,
};
use aether_usage_runtime::{
@@ -41,9 +42,16 @@ use crate::handlers::shared::provider_pool::{
admin_provider_pool_key_terminal_error_reason, record_admin_provider_pool_error,
record_admin_provider_pool_stream_timeout, record_admin_provider_pool_success,
release_admin_provider_pool_key_lease, AdminProviderPoolConfig,
AdminProviderPoolSchedulingPreset,
};
use crate::orchestration::{
local_execution_candidate_metadata_from_report_context,
ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD,
};
use crate::scheduler::affinity::{
scheduler_affinity_policy_context_from_report_context, SCHEDULER_AFFINITY_POLICY_REPORT_FIELD,
SCHEDULER_AFFINITY_TTL,
};
use crate::orchestration::local_execution_candidate_metadata_from_report_context;
use crate::scheduler::affinity::SCHEDULER_AFFINITY_TTL;
use crate::scheduler::config::{read_scheduler_ordering_config, SchedulerSchedulingMode};
use crate::AppState;
@@ -343,11 +351,22 @@ fn report_context_string_field<'a>(
fn local_scheduler_affinity_cache_key(report_context: Option<&Value>) -> Option<String> {
let client_session_affinity = local_client_session_affinity(report_context);
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session(
let policy_context = scheduler_affinity_policy_context_from_report_context(report_context);
if report_context
.and_then(|context| context.get(SCHEDULER_AFFINITY_POLICY_REPORT_FIELD))
.is_some()
&& policy_context.is_none()
{
return None;
}
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope(
report_context_string_field(report_context, "api_key_id")?,
report_context_string_field(report_context, "client_api_format")?,
report_context_string_field(report_context, "model")?,
client_session_affinity.as_ref(),
policy_context
.as_ref()
.and_then(|context| context.scope.as_ref()),
)
}
@@ -504,7 +523,17 @@ async fn local_scheduler_affinity_matches_failed_target(
local_execution_plan_uses_pool(state, plan).await
}
async fn scheduler_cache_affinity_enabled(state: &AppState) -> bool {
async fn scheduler_cache_affinity_enabled(
state: &AppState,
report_context: Option<&Value>,
) -> bool {
if report_context
.and_then(|context| context.get(SCHEDULER_AFFINITY_POLICY_REPORT_FIELD))
.is_some()
{
return scheduler_affinity_policy_context_from_report_context(report_context)
.is_some_and(|context| context.cache_affinity_enabled());
}
match read_scheduler_ordering_config(state).await {
Ok(config) => config.scheduling_mode == SchedulerSchedulingMode::CacheAffinity,
Err(error) => {
@@ -523,7 +552,7 @@ async fn remember_successful_local_scheduler_affinity(
state: &AppState,
context: LocalExecutionEffectContext<'_>,
) {
if !scheduler_cache_affinity_enabled(state).await {
if !scheduler_cache_affinity_enabled(state, context.report_context).await {
return;
}
let Some(cache_key) = local_scheduler_affinity_cache_key(context.report_context) else {
@@ -577,12 +606,33 @@ async fn resolve_pool_feedback_context(
}
};
let Some(pool_config) =
let Some(mut pool_config) =
admin_provider_pool_config_from_config_value(transport.provider.config.as_ref())
else {
return None;
};
if let Some(override_policy) = context
.report_context
.and_then(|report_context| report_context.get(ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD))
.and_then(|value| serde_json::from_value::<RoutingPoolPolicyOverride>(value.clone()).ok())
.filter(|override_policy| !override_policy.scheduling_presets.is_empty())
{
let scheduling_presets = override_policy
.scheduling_presets
.into_iter()
.map(|preset| AdminProviderPoolSchedulingPreset {
preset: preset.preset,
enabled: preset.enabled,
mode: preset.mode,
})
.collect::<Vec<_>>();
pool_config.lru_enabled = scheduling_presets
.iter()
.any(|preset| preset.enabled && preset.preset.eq_ignore_ascii_case("lru"));
pool_config.scheduling_presets = scheduling_presets;
}
let sticky_session_token = pool_feedback_request_body(plan, context.report_context)
.and_then(extract_pool_sticky_session_token);
@@ -1758,10 +1808,11 @@ mod tests {
apply_local_execution_effect, execution_plan_bearer_matches_transport,
local_candidate_failure_should_apply_key_effects,
local_candidate_failure_should_record_pool_error, pool_score_feedback_gate_allows,
pool_score_hard_state_for_status, LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect,
LocalAttemptFailureEffect, LocalExecutionEffect, LocalExecutionEffectContext,
LocalHealthFailureEffect, LocalHealthSuccessEffect, LocalOAuthInvalidationEffect,
LocalPoolErrorEffect, ProviderKeyEffectLockPool,
pool_score_hard_state_for_status, resolve_pool_feedback_context,
LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect,
LocalExecutionEffect, LocalExecutionEffectContext, LocalHealthFailureEffect,
LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalPoolErrorEffect,
ProviderKeyEffectLockPool,
};
use crate::data::{GatewayDataConfig, GatewayDataState};
use crate::orchestration::LocalFailoverClassification;
@@ -1770,7 +1821,8 @@ mod tests {
use aether_scheduler_core::{
build_scheduler_affinity_cache_key_for_api_key_id,
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session,
ClientSessionAffinity, SchedulerAffinityTarget,
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope,
ClientSessionAffinity, SchedulerAffinityScope, SchedulerAffinityTarget,
};
async fn start_managed_redis_or_skip() -> Option<ManagedRedisServer> {
@@ -2221,6 +2273,61 @@ mod tests {
)
}
#[tokio::test]
async fn pool_feedback_uses_routing_profile_scheduling_override() {
let mut provider = sample_pool_health_provider();
provider.config = Some(json!({
"pool_advanced": {
"scheduling_presets": [{
"preset": "lru",
"enabled": true
}]
}
}));
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![sample_health_endpoint()],
vec![sample_health_key()],
));
let state = AppState::new()
.expect("gateway state should build")
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(repository)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
);
let plan = sample_plan();
let report_context = json!({
"routing_pool_policy_override": {
"scheduling_presets": [{
"preset": "cache_affinity",
"enabled": true
}]
}
});
let feedback = resolve_pool_feedback_context(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: Some(&report_context),
},
)
.await
.expect("pool feedback context should resolve");
assert!(
crate::handlers::shared::provider_pool::admin_provider_pool_cache_affinity_enabled(
&feedback.pool_config
)
);
assert!(!feedback.pool_config.lru_enabled);
assert_eq!(feedback.pool_config.scheduling_presets.len(), 1);
assert_eq!(
feedback.pool_config.scheduling_presets[0].preset,
"cache_affinity"
);
}
fn health_state_with_key(key: StoredProviderCatalogKey) -> AppState {
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_health_provider()],
@@ -2632,6 +2739,144 @@ mod tests {
);
}
#[tokio::test]
async fn routing_profile_cache_affinity_overrides_legacy_fixed_mode_on_success() {
let state = AppState::new()
.expect("gateway state should build")
.with_data_state_for_tests(
GatewayDataState::disabled().with_system_config_values_for_tests(vec![(
"scheduling_mode".to_string(),
json!("fixed_order"),
)]),
);
let plan = sample_plan();
let affinity = session_affinity();
let scope = SchedulerAffinityScope::new("routing-group-1", Some(7));
let report_context = json!({
"api_key_id": "api-key-1",
"client_api_format": "openai:chat",
"model": "gpt-5",
"client_session_affinity": {
"client_family": "generic",
"session_key": "session=session-1;agent=coder"
},
"scheduler_affinity_policy": {
"scheduling_mode": "cache_affinity",
"scope": {
"routing_group_id": "routing-group-1",
"routing_group_version": 7
}
}
});
let scoped_cache_key =
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope(
"api-key-1",
"openai:chat",
"gpt-5",
Some(&affinity),
Some(&scope),
)
.expect("scoped scheduler affinity cache key should build");
apply_local_execution_effect(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: Some(&report_context),
},
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
)
.await;
assert_eq!(
state.read_scheduler_affinity_target(scoped_cache_key.as_str(), SCHEDULER_AFFINITY_TTL),
Some(SchedulerAffinityTarget {
provider_id: "prov-1".to_string(),
endpoint_id: "ep-1".to_string(),
key_id: "key-1".to_string(),
})
);
assert!(state
.read_scheduler_affinity_target(
session_scheduler_affinity_cache_key().as_str(),
SCHEDULER_AFFINITY_TTL
)
.is_none());
}
#[tokio::test]
async fn routing_profile_fixed_mode_overrides_legacy_cache_affinity_on_success() {
let state = AppState::new().expect("gateway state should build");
let plan = sample_plan();
let report_context = json!({
"api_key_id": "api-key-1",
"client_api_format": "openai:chat",
"model": "gpt-5",
"scheduler_affinity_policy": {
"scheduling_mode": "fixed_order",
"scope": {
"routing_group_id": "routing-group-1",
"routing_group_version": 7
}
}
});
let scope = SchedulerAffinityScope::new("routing-group-1", Some(7));
let scoped_cache_key =
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope(
"api-key-1",
"openai:chat",
"gpt-5",
None,
Some(&scope),
)
.expect("scoped scheduler affinity cache key should build");
apply_local_execution_effect(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: Some(&report_context),
},
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
)
.await;
assert!(state
.read_scheduler_affinity_target(scoped_cache_key.as_str(), SCHEDULER_AFFINITY_TTL)
.is_none());
}
#[tokio::test]
async fn malformed_routing_affinity_context_does_not_fall_back_to_legacy_mode() {
let state = AppState::new().expect("gateway state should build");
let plan = sample_plan();
let report_context = json!({
"api_key_id": "api-key-1",
"client_api_format": "openai:chat",
"model": "gpt-5",
"scheduler_affinity_policy": {
"scheduling_mode": "unknown"
}
});
let legacy_cache_key =
build_scheduler_affinity_cache_key_for_api_key_id("api-key-1", "openai:chat", "gpt-5")
.expect("legacy scheduler affinity cache key should build");
apply_local_execution_effect(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: Some(&report_context),
},
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
)
.await;
assert!(state
.read_scheduler_affinity_target(legacy_cache_key.as_str(), SCHEDULER_AFFINITY_TTL)
.is_none());
}
#[tokio::test]
async fn health_success_keeps_scheduler_affinity_after_health_state_update() {
let state = health_state();
+2 -1
View File
@@ -22,7 +22,8 @@ pub(crate) use self::attempt::{
attempt_identity_from_report_context, build_local_attempt_identities,
insert_pool_key_lease_report_context_fields, local_attempt_slot_count,
local_execution_candidate_metadata_from_report_context, ExecutionAttemptIdentity,
LocalExecutionCandidateMetadata, SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD,
LocalExecutionCandidateMetadata, ROUTING_POOL_POLICY_OVERRIDE_REPORT_FIELD,
SCHEDULER_AFFINITY_EPOCH_REPORT_FIELD,
};
pub(crate) use self::classifier::{
classify_anthropic_failure_disposition, classify_failure_disposition, classify_local_failover,
+130 -32
View File
@@ -14,6 +14,8 @@ pub(crate) enum GatewayRoutingSelectionError {
Disabled(String),
#[error("routing group was explicitly requested but is not allowed for this principal: {0}")]
Forbidden(String),
#[error("routing group repository lookup failed: {0}")]
Repository(String),
}
#[derive(Debug, Clone, Default)]
@@ -42,22 +44,21 @@ pub(crate) async fn select_gateway_routing_group(
let group = repository
.find_routing_group(RoutingGroupLookupKey::Id(explicit))
.await
.ok()
.flatten()
.or({
let group: Option<StoredRoutingGroup> = repository
.find_routing_group(RoutingGroupLookupKey::Name(explicit))
.await
.unwrap_or_default();
group
});
.map_err(repository_selection_error)?;
let group = match group {
Some(group) => Some(group),
None => repository
.find_routing_group(RoutingGroupLookupKey::Name(explicit))
.await
.map_err(repository_selection_error)?,
};
let Some(group) = group else {
return Err(GatewayRoutingSelectionError::NotFound(explicit.to_string()));
};
if !group.enabled {
return Err(GatewayRoutingSelectionError::Disabled(group.id));
}
if !explicit_group_allowed(repository, &group.id, &input).await {
if !explicit_group_allowed(repository, &group, &input).await? {
return Err(GatewayRoutingSelectionError::Forbidden(group.id));
}
return Ok(GatewayRoutingGroupSelection {
@@ -70,13 +71,15 @@ pub(crate) async fn select_gateway_routing_group(
// produce a group. The data-state repository answers this with a cached
// existence query, so the common "routing configured but unused" case
// does not materialize the binding table per API key/user.
let has_bindings = repository.has_any_routing_group_binding().await;
if matches!(has_bindings, Ok(false)) {
let has_bindings = repository
.has_any_routing_group_binding()
.await
.map_err(repository_selection_error)?;
if !has_bindings {
let system_default = repository
.find_routing_group(RoutingGroupLookupKey::SystemDefault)
.await
.ok()
.flatten()
.map_err(repository_selection_error)?
.filter(|group| group.enabled);
return Ok(GatewayRoutingGroupSelection {
group: system_default,
@@ -92,13 +95,12 @@ pub(crate) async fn select_gateway_routing_group(
subject_id: Some(subject_id.to_string()),
})
.await
.unwrap_or_default();
.map_err(repository_selection_error)?;
for binding in bindings.into_iter().filter(|binding| binding.is_default) {
let group = repository
.find_routing_group(RoutingGroupLookupKey::Id(&binding.group_id))
.await
.ok()
.flatten();
.map_err(repository_selection_error)?;
if let Some(group) = group.filter(|group| group.enabled) {
return Ok(GatewayRoutingGroupSelection {
group: Some(group),
@@ -111,8 +113,7 @@ pub(crate) async fn select_gateway_routing_group(
let system_default = repository
.find_routing_group(RoutingGroupLookupKey::SystemDefault)
.await
.ok()
.flatten()
.map_err(repository_selection_error)?
.filter(|group| group.enabled);
Ok(GatewayRoutingGroupSelection {
group: system_default,
@@ -122,31 +123,30 @@ pub(crate) async fn select_gateway_routing_group(
async fn explicit_group_allowed(
repository: &(impl RoutingGroupReadRepository + ?Sized),
group_id: &str,
group: &StoredRoutingGroup,
input: &GatewayRoutingSelectionInput<'_>,
) -> bool {
if let Ok(Some(group)) = repository
.find_routing_group(RoutingGroupLookupKey::Id(group_id))
.await
{
if group.is_system_default {
return true;
}
) -> Result<bool, GatewayRoutingSelectionError> {
if group.is_system_default {
return Ok(true);
}
for (subject_type, subject_id, _) in default_binding_candidates(input) {
let bindings = repository
.list_routing_group_bindings(&RoutingGroupBindingQuery {
group_id: Some(group_id.to_string()),
group_id: Some(group.id.clone()),
subject_type: Some(subject_type),
subject_id: Some(subject_id.to_string()),
})
.await
.unwrap_or_default();
.map_err(repository_selection_error)?;
if bindings.iter().any(|binding| binding.allow_explicit_select) {
return true;
return Ok(true);
}
}
false
Ok(false)
}
fn repository_selection_error(error: impl std::fmt::Display) -> GatewayRoutingSelectionError {
GatewayRoutingSelectionError::Repository(error.to_string())
}
fn default_binding_candidates<'a>(
@@ -178,11 +178,57 @@ mod tests {
use aether_data::repository::routing_profiles::InMemoryRoutingGroupRepository;
use aether_data_contracts::repository::routing_profiles::{
CreateRoutingGroupBindingRecord, CreateRoutingGroupRecord, RoutingGroupWriteRepository,
StoredRoutingGroupBinding, StoredRoutingGroupVersion,
};
use aether_data_contracts::DataLayerError;
use async_trait::async_trait;
use serde_json::json;
use super::*;
struct FailingRoutingGroupRepository {
id_lookup_is_missing: bool,
}
impl FailingRoutingGroupRepository {
fn failure<T>() -> Result<T, DataLayerError> {
Err(DataLayerError::Sql(
"routing repository unavailable".to_string(),
))
}
}
#[async_trait]
impl RoutingGroupReadRepository for FailingRoutingGroupRepository {
async fn list_routing_groups(&self) -> Result<Vec<StoredRoutingGroup>, DataLayerError> {
Self::failure()
}
async fn find_routing_group(
&self,
lookup: RoutingGroupLookupKey<'_>,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
if self.id_lookup_is_missing && matches!(lookup, RoutingGroupLookupKey::Id(_)) {
return Ok(None);
}
Self::failure()
}
async fn list_routing_group_bindings(
&self,
_query: &RoutingGroupBindingQuery,
) -> Result<Vec<StoredRoutingGroupBinding>, DataLayerError> {
Self::failure()
}
async fn list_routing_group_versions(
&self,
_group_id: &str,
) -> Result<Vec<StoredRoutingGroupVersion>, DataLayerError> {
Self::failure()
}
}
#[tokio::test]
async fn selects_api_key_default_binding() {
let repository = InMemoryRoutingGroupRepository::default();
@@ -336,6 +382,58 @@ mod tests {
);
}
#[tokio::test]
async fn propagates_explicit_name_lookup_failure_after_missing_id() {
let repository = FailingRoutingGroupRepository {
id_lookup_is_missing: true,
};
let error = select_gateway_routing_group(
&repository,
GatewayRoutingSelectionInput {
explicit_group: Some("group-name"),
user_id: Some("user-1"),
api_key_id: Some("api-key-1"),
user_group_ids: &[],
},
)
.await
.unwrap_err();
assert_eq!(
error,
GatewayRoutingSelectionError::Repository(
"sql error: routing repository unavailable".to_string()
)
);
}
#[tokio::test]
async fn propagates_implicit_binding_lookup_failure() {
let repository = FailingRoutingGroupRepository {
id_lookup_is_missing: false,
};
let error = select_gateway_routing_group(
&repository,
GatewayRoutingSelectionInput {
explicit_group: None,
user_id: Some("user-1"),
api_key_id: Some("api-key-1"),
user_group_ids: &[],
},
)
.await
.unwrap_err();
assert_eq!(
error,
GatewayRoutingSelectionError::Repository(
"sql error: routing repository unavailable".to_string()
)
);
}
#[tokio::test]
async fn rejects_explicit_disabled_group() {
let repository = InMemoryRoutingGroupRepository::default();
+76 -2
View File
@@ -1,13 +1,65 @@
use std::time::Duration;
use aether_routing_core::{ResolvedRoutingPolicy, RoutingSchedulingMode};
use aether_scheduler_core::{
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session, ClientSessionAffinity,
SchedulerAffinityTarget,
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session,
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope,
ClientSessionAffinity, SchedulerAffinityScope, SchedulerAffinityTarget,
};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use super::state::SchedulerRuntimeState;
pub(crate) const SCHEDULER_AFFINITY_TTL: Duration = Duration::from_secs(300);
pub(crate) const SCHEDULER_AFFINITY_POLICY_REPORT_FIELD: &str = "scheduler_affinity_policy";
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub(crate) struct SchedulerAffinityPolicyContext {
pub(crate) scheduling_mode: RoutingSchedulingMode,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub(crate) scope: Option<SchedulerAffinityScope>,
}
impl SchedulerAffinityPolicyContext {
pub(crate) fn from_routing_policy(policy: &ResolvedRoutingPolicy) -> Self {
let scope = policy
.group_id
.as_deref()
.map(str::trim)
.filter(|group_id| !group_id.is_empty())
.map(|group_id| SchedulerAffinityScope::new(group_id, policy.group_version));
Self {
scheduling_mode: policy.scheduling_mode,
scope,
}
}
pub(crate) fn cache_affinity_enabled(&self) -> bool {
self.scheduling_mode == RoutingSchedulingMode::CacheAffinity
}
}
pub(crate) fn scheduler_affinity_policy_context_from_report_context(
report_context: Option<&Value>,
) -> Option<SchedulerAffinityPolicyContext> {
report_context
.and_then(|context| context.get(SCHEDULER_AFFINITY_POLICY_REPORT_FIELD))
.and_then(|value| serde_json::from_value(value.clone()).ok())
}
pub(crate) fn insert_scheduler_affinity_policy_report_context_field(
extra_fields: &mut serde_json::Map<String, Value>,
routing_policy: Option<&ResolvedRoutingPolicy>,
) {
let Some(routing_policy) = routing_policy else {
return;
};
let context = SchedulerAffinityPolicyContext::from_routing_policy(routing_policy);
if let Ok(value) = serde_json::to_value(context) {
extra_fields.insert(SCHEDULER_AFFINITY_POLICY_REPORT_FIELD.to_string(), value);
}
}
pub(crate) fn read_cached_scheduler_affinity_target(
state: &(impl SchedulerRuntimeState + ?Sized),
@@ -24,3 +76,25 @@ pub(crate) fn read_cached_scheduler_affinity_target(
)?;
state.read_cached_scheduler_affinity_target(&cache_key, SCHEDULER_AFFINITY_TTL)
}
pub(crate) fn read_cached_scheduler_affinity_target_with_policy_context(
state: &(impl SchedulerRuntimeState + ?Sized),
api_key_id: &str,
client_session_affinity: Option<&ClientSessionAffinity>,
api_format: &str,
global_model_name: &str,
policy_context: &SchedulerAffinityPolicyContext,
) -> Option<SchedulerAffinityTarget> {
if !policy_context.cache_affinity_enabled() {
return None;
}
let cache_key =
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope(
api_key_id,
api_format,
global_model_name,
client_session_affinity,
policy_context.scope.as_ref(),
)?;
state.read_cached_scheduler_affinity_target(&cache_key, SCHEDULER_AFFINITY_TTL)
}
@@ -279,11 +279,24 @@ async fn gateway_handles_admin_pool_scheduling_presets_locally_with_trusted_admi
assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse");
let items = payload.as_array().expect("payload should be an array");
assert_eq!(items.len(), 14);
assert_eq!(items[0]["name"], "lru");
assert_eq!(items[1]["name"], "cache_affinity");
assert_eq!(items[8]["name"], "pro_first");
assert_eq!(items[13]["name"], "team_first");
assert_eq!(items.len(), 15);
let preset_names = items
.iter()
.filter_map(|item| item["name"].as_str())
.collect::<Vec<_>>();
assert_eq!(preset_names.first(), Some(&"lru"));
assert_eq!(preset_names.get(1), Some(&"cache_affinity"));
assert_eq!(preset_names.last(), Some(&"team_first"));
let preset_index = |name| {
preset_names
.iter()
.position(|preset_name| *preset_name == name)
.unwrap_or_else(|| panic!("missing scheduling preset {name}"))
};
assert!(preset_index("cost_first") < preset_index("free_team_first"));
assert!(preset_index("free_team_first") < preset_index("free_first"));
assert!(preset_index("pro_first") < preset_index("team_first"));
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort();