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();
@@ -0,0 +1,115 @@
UPDATE routing_groups
SET is_system_default = 0
WHERE id IN (
SELECT id
FROM (
SELECT
id,
ROW_NUMBER() OVER (ORDER BY enabled DESC, updated_at DESC, id ASC) AS default_rank
FROM routing_groups
WHERE is_system_default = 1
) AS ranked_defaults
WHERE default_rank > 1
);
UPDATE routing_group_bindings
SET is_default = 0
WHERE id IN (
SELECT id
FROM (
SELECT
id,
ROW_NUMBER() OVER (
PARTITION BY subject_type, subject_id
ORDER BY created_at ASC, id ASC
) AS default_rank
FROM routing_group_bindings
WHERE is_default = 1
) AS ranked_defaults
WHERE default_rank > 1
);
SET @aether_routing_groups_default_guard_column_sql := IF(
(
SELECT COUNT(*)
FROM information_schema.columns
WHERE table_schema = DATABASE()
AND table_name = 'routing_groups'
AND column_name = 'system_default_unique_guard'
) = 0,
'ALTER TABLE routing_groups ADD COLUMN system_default_unique_guard TINYINT GENERATED ALWAYS AS (CASE WHEN is_system_default = 1 THEN 1 ELSE NULL END) VIRTUAL INVISIBLE',
'DO 0'
);
PREPARE aether_routing_groups_default_guard_column_stmt
FROM @aether_routing_groups_default_guard_column_sql;
EXECUTE aether_routing_groups_default_guard_column_stmt;
DEALLOCATE PREPARE aether_routing_groups_default_guard_column_stmt;
SET @aether_routing_bindings_default_type_guard_column_sql := IF(
(
SELECT COUNT(*)
FROM information_schema.columns
WHERE table_schema = DATABASE()
AND table_name = 'routing_group_bindings'
AND column_name = 'default_subject_type_guard'
) = 0,
'ALTER TABLE routing_group_bindings ADD COLUMN default_subject_type_guard VARCHAR(32) GENERATED ALWAYS AS (CASE WHEN is_default = 1 THEN subject_type ELSE NULL END) VIRTUAL INVISIBLE',
'DO 0'
);
PREPARE aether_routing_bindings_default_type_guard_column_stmt
FROM @aether_routing_bindings_default_type_guard_column_sql;
EXECUTE aether_routing_bindings_default_type_guard_column_stmt;
DEALLOCATE PREPARE aether_routing_bindings_default_type_guard_column_stmt;
SET @aether_routing_bindings_default_id_guard_column_sql := IF(
(
SELECT COUNT(*)
FROM information_schema.columns
WHERE table_schema = DATABASE()
AND table_name = 'routing_group_bindings'
AND column_name = 'default_subject_id_guard'
) = 0,
'ALTER TABLE routing_group_bindings ADD COLUMN default_subject_id_guard VARCHAR(64) GENERATED ALWAYS AS (CASE WHEN is_default = 1 THEN subject_id ELSE NULL END) VIRTUAL INVISIBLE',
'DO 0'
);
PREPARE aether_routing_bindings_default_id_guard_column_stmt
FROM @aether_routing_bindings_default_id_guard_column_sql;
EXECUTE aether_routing_bindings_default_id_guard_column_stmt;
DEALLOCATE PREPARE aether_routing_bindings_default_id_guard_column_stmt;
SET @aether_routing_groups_default_unique_index_sql := IF(
(
SELECT COUNT(*)
FROM information_schema.statistics
WHERE table_schema = DATABASE()
AND table_name = 'routing_groups'
AND index_name = 'routing_groups_one_system_default_key'
) = 0,
'CREATE UNIQUE INDEX routing_groups_one_system_default_key ON routing_groups (system_default_unique_guard)',
'DO 0'
);
PREPARE aether_routing_groups_default_unique_index_stmt
FROM @aether_routing_groups_default_unique_index_sql;
EXECUTE aether_routing_groups_default_unique_index_stmt;
DEALLOCATE PREPARE aether_routing_groups_default_unique_index_stmt;
SET @aether_routing_bindings_default_unique_index_sql := IF(
(
SELECT COUNT(*)
FROM information_schema.statistics
WHERE table_schema = DATABASE()
AND table_name = 'routing_group_bindings'
AND index_name = 'routing_group_bindings_subject_default_key'
) = 0,
'CREATE UNIQUE INDEX routing_group_bindings_subject_default_key ON routing_group_bindings (default_subject_type_guard, default_subject_id_guard)',
'DO 0'
);
PREPARE aether_routing_bindings_default_unique_index_stmt
FROM @aether_routing_bindings_default_unique_index_sql;
EXECUTE aether_routing_bindings_default_unique_index_stmt;
DEALLOCATE PREPARE aether_routing_bindings_default_unique_index_stmt;
@@ -4,10 +4,10 @@ use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
};
use aether_data_contracts::DataLayerError;
@@ -37,7 +37,6 @@ SELECT
pak.capabilities AS key_capabilities,
pak.internal_priority AS key_internal_priority,
pak.global_priority_by_format AS key_global_priority_by_format,
pak.last_used_at AS key_last_used_at_unix_secs,
m.id AS model_id,
m.global_model_id AS global_model_id,
gm.name AS global_model_name,
@@ -60,6 +59,9 @@ WHERE p.is_active = 1
AND gm.is_active = 1
"#;
const REQUESTED_MODEL_RAW_PAGE_SIZE: u32 = 256;
const REQUESTED_MODEL_RAW_SCAN_LIMIT: u32 = 2048;
#[derive(Debug, Clone)]
pub struct MysqlMinimalCandidateSelectionReadRepository {
pool: MysqlPool,
@@ -70,7 +72,52 @@ struct CandidateSelectionRow {
row: StoredMinimalCandidateSelectionRow,
provider_pool_enabled: bool,
key_auth_config: Option<String>,
key_last_used_at_unix_secs: Option<u64>,
}
#[derive(Debug)]
struct ExactPageAccumulator<T> {
rows: Vec<T>,
offset: usize,
limit: usize,
target_len: usize,
}
impl<T> ExactPageAccumulator<T> {
fn new(offset: u32, limit: u32) -> Self {
let offset = usize::try_from(offset).unwrap_or(usize::MAX);
let limit = usize::try_from(limit).unwrap_or(usize::MAX);
Self {
rows: Vec::new(),
offset,
limit,
target_len: offset.saturating_add(limit),
}
}
fn is_full(&self) -> bool {
self.rows.len() >= self.target_len
}
fn push_matching<I, F>(&mut self, rows: I, mut predicate: F)
where
I: IntoIterator<Item = T>,
F: FnMut(&T) -> bool,
{
let remaining = self.target_len.saturating_sub(self.rows.len());
self.rows.extend(
rows.into_iter()
.filter(|row| predicate(row))
.take(remaining),
);
}
fn into_page(self) -> Vec<T> {
self.rows
.into_iter()
.skip(self.offset)
.take(self.limit)
.collect()
}
}
impl MysqlMinimalCandidateSelectionReadRepository {
@@ -115,6 +162,263 @@ impl MysqlMinimalCandidateSelectionReadRepository {
let rows = self.load_rows_for_api_format(api_format).await?;
Ok(sort_rows(select_pool_rows(rows), true))
}
async fn load_selected_rows_for_api_format_page(
&self,
api_format: &str,
limit: u32,
offset: u32,
) -> Result<Vec<CandidateSelectionRow>, DataLayerError> {
let mut builder = api_format_page_query(api_format, limit, offset);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_candidate_selection_row).collect()
}
async fn selected_rows_for_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
if query.limit == 0 {
return Ok(Vec::new());
}
let target_len = query.offset.saturating_add(query.limit);
let mut raw_offset = 0_u32;
let mut selected = Vec::new();
while selected.len() < target_len as usize && raw_offset < REQUESTED_MODEL_RAW_SCAN_LIMIT {
let raw_limit =
target_len.min(REQUESTED_MODEL_RAW_SCAN_LIMIT.saturating_sub(raw_offset));
let rows = self
.load_selected_rows_for_api_format_page(&query.api_format, raw_limit, raw_offset)
.await?;
let raw_len = rows.len() as u32;
selected.extend(rows.into_iter().filter(|item| {
api_format_matches(&item.row.endpoint_api_format, &query.api_format)
&& item.row.key_supports_api_format(&query.api_format)
&& key_auth_channel_matches(item, &query.api_format)
}));
if raw_len < raw_limit || raw_len == 0 {
break;
}
let next_offset = raw_offset.saturating_add(raw_len);
if next_offset == raw_offset {
break;
}
raw_offset = next_offset;
}
let rows = selected.into_iter().map(|item| item.row).collect();
Ok(sort_rows(dedupe_candidate_selection_rows(rows), true)
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
}
}
fn api_format_page_query(
api_format: &str,
limit: u32,
offset: u32,
) -> QueryBuilder<'static, MySql> {
let mut builder = QueryBuilder::<MySql>::new("WITH candidate_rows AS (");
builder.push(CANDIDATE_SELECTION_COLUMNS);
push_candidate_api_format_filters(&mut builder, api_format);
push_selected_pool_rows(&mut builder);
builder.push(
r#"
ORDER BY
global_model_name ASC,
provider_priority ASC,
key_internal_priority ASC,
provider_id ASC,
endpoint_id ASC,
key_id ASC,
model_id ASC
LIMIT "#,
);
builder.push_bind(i64::from(limit));
builder.push(" OFFSET ");
builder.push_bind(i64::from(offset));
builder
}
fn requested_model_page_query(
api_format: &str,
requested_model_name: &str,
limit: u32,
offset: u32,
) -> QueryBuilder<'static, MySql> {
let mut builder = QueryBuilder::<MySql>::new("WITH candidate_rows AS (");
builder.push(CANDIDATE_SELECTION_COLUMNS);
push_candidate_api_format_filters(&mut builder, api_format);
push_requested_model_sql_filter(&mut builder, requested_model_name);
push_selected_pool_rows(&mut builder);
builder.push(
r#"
ORDER BY
global_model_name ASC,
provider_priority ASC,
key_internal_priority ASC,
provider_id ASC,
endpoint_id ASC,
key_id ASC,
model_id ASC
LIMIT "#,
);
builder.push_bind(i64::from(limit));
builder.push(" OFFSET ");
builder.push_bind(i64::from(offset));
builder
}
fn pool_key_group_query(query: &StoredPoolKeyCandidateRowsQuery) -> QueryBuilder<'static, MySql> {
let mut builder = QueryBuilder::<MySql>::new(CANDIDATE_SELECTION_COLUMNS);
push_candidate_api_format_filters(&mut builder, &query.api_format);
push_pool_key_group_filters(
&mut builder,
&query.provider_id,
&query.endpoint_id,
&query.model_id,
);
push_pool_key_order(&mut builder, &query.order);
builder.push(" LIMIT ");
builder.push_bind(i64::from(query.limit));
builder.push(" OFFSET ");
builder.push_bind(i64::from(query.offset));
builder
}
fn pool_key_group_by_key_ids_query(
query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
) -> QueryBuilder<'static, MySql> {
let mut builder = QueryBuilder::<MySql>::new(CANDIDATE_SELECTION_COLUMNS);
push_candidate_api_format_filters(&mut builder, &query.api_format);
push_pool_key_group_filters(
&mut builder,
&query.provider_id,
&query.endpoint_id,
&query.model_id,
);
builder.push(" AND pak.id IN (");
{
let mut separated = builder.separated(", ");
for key_id in &query.key_ids {
separated.push_bind(key_id.clone());
}
}
builder.push(") ORDER BY FIELD(pak.id, ");
{
let mut separated = builder.separated(", ");
for key_id in &query.key_ids {
separated.push_bind(key_id.clone());
}
}
builder.push(") ASC, pak.id ASC");
builder
}
fn push_candidate_api_format_filters(builder: &mut QueryBuilder<'_, MySql>, api_format: &str) {
let canonical_api_format = normalize_api_format(api_format);
let storage_aliases = sql_match_aliases(&api_format_aliases(&canonical_api_format));
let permission_aliases =
sql_match_aliases(&api_format_permission_aliases(&canonical_api_format));
builder.push(" AND LOWER(pe.api_format) IN (");
{
let mut separated = builder.separated(", ");
for alias in storage_aliases {
separated.push_bind(alias);
}
}
builder.push(") AND (pak.api_formats IS NULL OR TRIM(pak.api_formats) = ''");
for alias in permission_aliases {
builder.push(" OR JSON_SEARCH(LOWER(pak.api_formats), 'one', ");
builder.push_bind(alias);
builder.push(") IS NOT NULL");
}
builder.push(")");
push_key_auth_channel_filter(builder, &canonical_api_format);
}
fn push_requested_model_sql_filter(
builder: &mut QueryBuilder<'_, MySql>,
requested_model_name: &str,
) {
builder.push(" AND (gm.name = ");
builder.push_bind(requested_model_name.to_string());
builder.push(" OR m.provider_model_name = ");
builder.push_bind(requested_model_name.to_string());
builder.push(" OR (m.provider_model_mappings IS NOT NULL AND LOCATE(");
builder.push_bind(requested_model_name.to_string());
builder.push(", m.provider_model_mappings) > 0))");
}
fn push_selected_pool_rows(builder: &mut QueryBuilder<'_, MySql>) {
builder.push(
r#"
),
ranked_rows AS (
SELECT
candidate_rows.*,
CASE
WHEN JSON_EXTRACT(provider_config, '$.pool_advanced') IS NOT NULL
AND JSON_TYPE(JSON_EXTRACT(provider_config, '$.pool_advanced')) <> 'NULL'
THEN 1 ELSE 0
END AS provider_pool_enabled,
ROW_NUMBER() OVER (
PARTITION BY provider_id, endpoint_id, model_id
ORDER BY key_internal_priority ASC, key_id ASC
) AS pool_rank
FROM candidate_rows
),
selected_rows AS (
SELECT *
FROM ranked_rows
WHERE provider_pool_enabled = 0 OR pool_rank = 1
)
SELECT *
FROM selected_rows"#,
);
}
fn push_pool_key_group_filters(
builder: &mut QueryBuilder<'_, MySql>,
provider_id: &str,
endpoint_id: &str,
model_id: &str,
) {
builder.push(" AND p.id = ");
builder.push_bind(provider_id.to_string());
builder.push(" AND pe.id = ");
builder.push_bind(endpoint_id.to_string());
builder.push(" AND m.id = ");
builder.push_bind(model_id.to_string());
}
fn push_pool_key_order(builder: &mut QueryBuilder<'_, MySql>, order: &StoredPoolKeyCandidateOrder) {
match order {
StoredPoolKeyCandidateOrder::InternalPriority => {
builder.push(" ORDER BY pak.internal_priority ASC, pak.id ASC");
}
StoredPoolKeyCandidateOrder::Lru => {
builder.push(
" ORDER BY pak.last_used_at IS NOT NULL ASC, pak.last_used_at ASC, pak.internal_priority ASC, pak.id ASC",
);
}
StoredPoolKeyCandidateOrder::CacheAffinity => {
builder.push(
" ORDER BY pak.last_used_at IS NULL ASC, pak.last_used_at DESC, pak.internal_priority ASC, pak.id ASC",
);
}
StoredPoolKeyCandidateOrder::SingleAccount => {
builder.push(
" ORDER BY pak.internal_priority ASC, pak.last_used_at IS NULL ASC, pak.last_used_at DESC, pak.id ASC",
);
}
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
builder.push(" ORDER BY MD5(CONCAT(");
builder.push_bind(seed.clone());
builder.push(", ':', pak.id)) ASC, pak.id ASC");
}
}
}
#[async_trait]
@@ -126,6 +430,13 @@ impl MinimalCandidateSelectionReadRepository for MysqlMinimalCandidateSelectionR
self.selected_rows_for_api_format(api_format).await
}
async fn list_for_exact_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.selected_rows_for_api_format_page(query).await
}
async fn list_for_exact_api_format_and_global_model(
&self,
api_format: &str,
@@ -146,57 +457,76 @@ impl MinimalCandidateSelectionReadRepository for MysqlMinimalCandidateSelectionR
api_format: &str,
requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.list_for_exact_api_format_and_requested_model_page(
&StoredRequestedModelCandidateRowsQuery {
api_format: api_format.to_string(),
requested_model_name: requested_model_name.to_string(),
offset: 0,
limit: u32::MAX,
},
)
.await
let rows = self
.selected_rows_for_api_format(api_format)
.await?
.into_iter()
.filter(|row| row_matches_requested_model(row, requested_model_name, api_format))
.collect::<Vec<_>>();
Ok(sort_rows(rows, true))
}
async fn list_for_exact_api_format_and_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let rows = self
.selected_rows_for_api_format(&query.api_format)
.await?
if query.limit == 0 {
return Ok(Vec::new());
}
// The SQL model predicate is a coarse superset, so fill the exact page across raw windows.
let mut exact_page = ExactPageAccumulator::new(query.offset, query.limit);
let mut raw_offset = 0_u32;
while !exact_page.is_full() && raw_offset < REQUESTED_MODEL_RAW_SCAN_LIMIT {
let raw_limit = REQUESTED_MODEL_RAW_PAGE_SIZE
.min(REQUESTED_MODEL_RAW_SCAN_LIMIT.saturating_sub(raw_offset));
let mut builder = requested_model_page_query(
&query.api_format,
&query.requested_model_name,
raw_limit,
raw_offset,
);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let raw_len = u32::try_from(rows.len()).unwrap_or(u32::MAX);
let items = rows
.iter()
.map(map_candidate_selection_row)
.collect::<Result<Vec<_>, _>>()?;
exact_page.push_matching(items, |item| {
row_matches_requested_model(
&item.row,
&query.requested_model_name,
&query.api_format,
)
});
raw_offset = raw_offset.saturating_add(raw_len);
if raw_len < raw_limit || raw_len == 0 {
break;
}
}
let rows = exact_page
.into_page()
.into_iter()
.filter(|row| {
row_matches_requested_model(row, &query.requested_model_name, &query.api_format)
})
.collect::<Vec<_>>();
Ok(sort_rows(rows, true)
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
.map(|item| item.row)
.collect();
Ok(sort_rows(dedupe_candidate_selection_rows(rows), true))
}
async fn list_pool_key_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let rows = self
.load_rows_for_api_format(&query.api_format)
.await?
.into_iter()
.filter(|row| {
row.row.provider_id == query.provider_id
&& row.row.endpoint_id == query.endpoint_id
&& row.row.model_id == query.model_id
})
.collect::<Vec<_>>();
let mut rows = sort_pool_key_rows(rows, &query.order);
Ok(rows
.drain(..)
.skip(query.offset as usize)
.take(query.limit as usize)
.map(|item| item.row)
.collect())
if query.limit == 0 {
return Ok(Vec::new());
}
let mut builder = pool_key_group_query(query);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let rows = rows
.iter()
.map(map_candidate_selection_row)
.collect::<Result<Vec<_>, _>>()?;
Ok(dedupe_candidate_selection_rows(
rows.into_iter().map(|item| item.row).collect(),
))
}
async fn list_pool_key_rows_for_group_key_ids(
@@ -212,25 +542,23 @@ impl MinimalCandidateSelectionReadRepository for MysqlMinimalCandidateSelectionR
.enumerate()
.map(|(index, key_id)| (key_id.as_str(), index))
.collect::<BTreeMap<_, _>>();
let mut rows = self
.load_rows_for_api_format(&query.api_format)
.await?
let mut builder = pool_key_group_by_key_ids_query(query);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let mut rows = rows
.iter()
.map(map_candidate_selection_row)
.collect::<Result<Vec<_>, _>>()?
.into_iter()
.filter(|row| {
row.row.provider_id == query.provider_id
&& row.row.endpoint_id == query.endpoint_id
&& row.row.model_id == query.model_id
&& key_order.contains_key(row.row.key_id.as_str())
})
.map(|item| item.row)
.collect::<Vec<_>>();
rows = dedupe_candidate_selection_rows(rows);
rows.sort_by(|left, right| {
key_order
.get(left.key_id.as_str())
.cmp(&key_order.get(right.key_id.as_str()))
.then(left.key_id.cmp(&right.key_id))
});
Ok(dedupe_candidate_selection_rows(rows))
Ok(rows)
}
}
@@ -283,67 +611,6 @@ fn sort_rows(
rows
}
fn sort_pool_key_rows(
mut rows: Vec<CandidateSelectionRow>,
order: &StoredPoolKeyCandidateOrder,
) -> Vec<CandidateSelectionRow> {
rows.sort_by(|left, right| match order {
StoredPoolKeyCandidateOrder::InternalPriority => compare_pool_key_internal(left, right),
StoredPoolKeyCandidateOrder::Lru => left
.key_last_used_at_unix_secs
.cmp(&right.key_last_used_at_unix_secs)
.then_with(|| compare_pool_key_internal(left, right)),
StoredPoolKeyCandidateOrder::CacheAffinity => right
.key_last_used_at_unix_secs
.cmp(&left.key_last_used_at_unix_secs)
.then_with(|| compare_pool_key_internal(left, right)),
StoredPoolKeyCandidateOrder::SingleAccount => left
.row
.key_internal_priority
.cmp(&right.row.key_internal_priority)
.then_with(|| {
right
.key_last_used_at_unix_secs
.cmp(&left.key_last_used_at_unix_secs)
})
.then(left.row.key_id.cmp(&right.row.key_id)),
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
stable_pool_key_hash(seed.as_str(), left.row.key_id.as_str())
.cmp(&stable_pool_key_hash(
seed.as_str(),
right.row.key_id.as_str(),
))
.then(left.row.key_id.cmp(&right.row.key_id))
}
});
rows
}
fn compare_pool_key_internal(
left: &CandidateSelectionRow,
right: &CandidateSelectionRow,
) -> std::cmp::Ordering {
left.row
.key_internal_priority
.cmp(&right.row.key_internal_priority)
.then(left.row.key_id.cmp(&right.row.key_id))
}
fn stable_pool_key_hash(seed: &str, key_id: &str) -> u64 {
let mut hash = 0xcbf29ce484222325u64;
for byte in seed
.as_bytes()
.iter()
.copied()
.chain(std::iter::once(b':'))
.chain(key_id.as_bytes().iter().copied())
{
hash ^= u64::from(byte);
hash = hash.wrapping_mul(0x100000001b3);
}
hash
}
fn row_matches_requested_model(
row: &StoredMinimalCandidateSelectionRow,
requested_model_name: &str,
@@ -418,6 +685,57 @@ fn mapping_scope_matches(
})
}
fn push_key_auth_channel_filter(builder: &mut QueryBuilder<'_, MySql>, api_format: &str) {
builder.push(" AND ((LOWER(TRIM(p.provider_type)) = 'codex'");
builder.push(" AND LOWER(TRIM(pak.auth_type)) = 'oauth' AND ");
builder.push_bind(api_format.to_string());
builder.push(
" IN ('openai:responses', 'openai:responses:compact', 'openai:search', 'openai:image'))",
);
builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'chatgpt_web'");
builder.push(" AND LOWER(TRIM(pak.auth_type)) IN ('oauth', 'bearer') AND ");
builder.push_bind(api_format.to_string());
builder.push(" = 'openai:image')");
builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'claude_code'");
builder.push(" AND LOWER(TRIM(pak.auth_type)) = 'oauth' AND ");
builder.push_bind(api_format.to_string());
builder.push(" = 'claude:messages')");
builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'kiro' AND ");
builder.push_bind(api_format.to_string());
builder.push(" = 'claude:messages' AND (LOWER(TRIM(pak.auth_type)) = 'oauth'");
builder.push(" OR (LOWER(TRIM(pak.auth_type)) = 'bearer'");
builder.push(" AND pak.auth_config IS NOT NULL AND TRIM(pak.auth_config) <> '')))");
builder.push(" OR (LOWER(TRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')");
builder.push(" AND LOWER(TRIM(pak.auth_type)) = 'oauth' AND ");
builder.push_bind(api_format.to_string());
builder.push(" = 'gemini:generate_content')");
builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'grok'");
builder.push(" AND LOWER(TRIM(pak.auth_type)) = 'oauth' AND ");
builder.push_bind(api_format.to_string());
builder.push(" IN ('openai:chat', 'openai:responses', 'claude:messages', 'openai:image'))");
builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'windsurf'");
builder.push(" AND LOWER(TRIM(pak.auth_type)) IN ('oauth', 'api_key', 'bearer') AND ");
builder.push_bind(api_format.to_string());
builder.push(" = 'openai:chat')");
builder.push(" OR (LOWER(TRIM(p.provider_type)) = 'vertex_ai'");
builder.push(
" AND LOWER(TRIM(pak.auth_type)) IN ('api_key', 'service_account', 'vertex_ai') AND ",
);
builder.push_bind(api_format.to_string());
builder.push(" IN ('gemini:generate_content', 'gemini:embedding'))");
builder.push(
" OR (LOWER(TRIM(p.provider_type)) NOT IN ('chatgpt_web', 'claude_code', 'codex', 'gemini_cli', 'grok', 'vertex_ai', 'antigravity', 'kiro', 'windsurf') AND LOWER(TRIM(pak.auth_type)) <> 'oauth'))",
);
}
fn key_auth_channel_matches(row: &CandidateSelectionRow, api_format: &str) -> bool {
let provider_type = row.row.provider_type.trim().to_ascii_lowercase();
let auth_type = row.row.key_auth_type.trim().to_ascii_lowercase();
@@ -543,10 +861,6 @@ fn map_candidate_selection_row(row: &MySqlRow) -> Result<CandidateSelectionRow,
},
provider_pool_enabled,
key_auth_config: row.try_get("key_auth_config").map_sql_err()?,
key_last_used_at_unix_secs: row
.try_get::<Option<i64>, _>("key_last_used_at_unix_secs")
.map_sql_err()?
.and_then(|value| u64::try_from(value).ok()),
})
}
@@ -770,6 +1084,10 @@ fn api_format_aliases(api_format: &str) -> Vec<String> {
aether_ai_formats::api_format_storage_aliases(api_format)
}
fn api_format_permission_aliases(api_format: &str) -> Vec<String> {
aether_ai_formats::api_format_permission_storage_aliases(api_format)
}
fn normalize_api_format(api_format: &str) -> String {
aether_ai_formats::normalize_api_format_alias(api_format)
}
@@ -791,7 +1109,134 @@ fn sql_match_aliases(api_formats: &[String]) -> Vec<String> {
#[cfg(test)]
mod tests {
use super::{vertex_key_auth_channel_matches, MysqlMinimalCandidateSelectionReadRepository};
use super::{
api_format_page_query, pool_key_group_by_key_ids_query, pool_key_group_query,
requested_model_page_query, vertex_key_auth_channel_matches, ExactPageAccumulator,
MysqlMinimalCandidateSelectionReadRepository, REQUESTED_MODEL_RAW_SCAN_LIMIT,
};
use aether_data_contracts::repository::candidate_selection::{
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery,
};
#[test]
fn api_format_page_query_uses_portable_sql_pagination_and_stable_order() {
let query = api_format_page_query("openai:chat", 256, 512);
let sql = query.sql();
assert!(sql.contains("ROW_NUMBER() OVER ("));
assert!(sql.contains("JSON_SEARCH(LOWER(pak.api_formats), 'one', ?"));
assert!(sql.contains("WHERE provider_pool_enabled = 0 OR pool_rank = 1"));
assert!(sql.contains(
"ORDER BY\n global_model_name ASC,\n provider_priority ASC,\n key_internal_priority ASC,"
));
assert!(sql.contains("LIMIT ? OFFSET ?"));
}
#[test]
fn requested_model_page_query_filters_and_pages_before_fetch() {
let query = requested_model_page_query("openai:chat", "gpt-5", 256, 256);
let sql = query.sql();
assert!(sql.contains("LOWER(pe.api_format) IN ("));
assert!(sql.contains("JSON_SEARCH(LOWER(pak.api_formats), 'one', ?"));
assert!(sql.contains("LOWER(TRIM(p.provider_type)) = 'codex'"));
assert!(sql.contains("AND (gm.name = ? OR m.provider_model_name = ?"));
assert!(sql.contains("LOCATE(?, m.provider_model_mappings) > 0"));
assert!(sql.contains("ROW_NUMBER() OVER ("));
assert!(sql.contains("LIMIT ? OFFSET ?"));
}
#[test]
fn exact_page_accumulator_continues_after_coarse_false_positives() {
let mut accumulator = ExactPageAccumulator::new(1, 2);
accumulator.push_matching(vec![("coarse-1", false), ("coarse-2", false)], |row| row.1);
assert!(!accumulator.is_full());
accumulator.push_matching(
vec![
("exact-1", true),
("coarse-3", false),
("exact-2", true),
("exact-3", true),
],
|row| row.1,
);
assert!(accumulator.is_full());
assert_eq!(
accumulator.into_page(),
vec![("exact-2", true), ("exact-3", true)]
);
assert_eq!(REQUESTED_MODEL_RAW_SCAN_LIMIT, 2048);
}
#[test]
fn pool_group_query_pushes_group_filters_order_and_page_into_sql() {
let query = pool_key_group_query(&StoredPoolKeyCandidateRowsQuery {
api_format: "openai:chat".to_string(),
provider_id: "provider-1".to_string(),
endpoint_id: "endpoint-1".to_string(),
model_id: "model-1".to_string(),
selected_provider_model_name: "gpt-5".to_string(),
order: StoredPoolKeyCandidateOrder::Lru,
offset: 64,
limit: 64,
});
let sql = query.sql();
assert!(sql.contains("LOWER(pe.api_format) IN ("));
assert!(sql.contains("JSON_SEARCH(LOWER(pak.api_formats), 'one', ?"));
assert!(sql.contains("LOWER(TRIM(p.provider_type)) = 'codex'"));
assert!(sql.contains("AND p.id = ? AND pe.id = ? AND m.id = ?"));
assert!(sql.contains(
"ORDER BY pak.last_used_at IS NOT NULL ASC, pak.last_used_at ASC, pak.internal_priority ASC, pak.id ASC"
));
assert!(sql.contains("LIMIT ? OFFSET ?"));
}
#[test]
fn load_balance_pool_group_query_uses_seeded_sql_order_before_paging() {
let query = pool_key_group_query(&StoredPoolKeyCandidateRowsQuery {
api_format: "openai:chat".to_string(),
provider_id: "provider-1".to_string(),
endpoint_id: "endpoint-1".to_string(),
model_id: "model-1".to_string(),
selected_provider_model_name: "gpt-5".to_string(),
order: StoredPoolKeyCandidateOrder::LoadBalance {
seed: "request-1".to_string(),
},
offset: 128,
limit: 64,
});
let sql = query.sql();
assert!(sql.contains("AND p.id = ? AND pe.id = ? AND m.id = ?"));
assert!(
sql.contains("ORDER BY MD5(CONCAT(?, ':', pak.id)) ASC, pak.id ASC LIMIT ? OFFSET ?")
);
assert!(!sql.contains("ROW_NUMBER() OVER ("));
}
#[test]
fn pool_group_by_key_ids_query_filters_ids_and_preserves_requested_order() {
let query = pool_key_group_by_key_ids_query(&StoredPoolKeyCandidateRowsByKeyIdsQuery {
api_format: "openai:chat".to_string(),
provider_id: "provider-1".to_string(),
endpoint_id: "endpoint-1".to_string(),
model_id: "model-1".to_string(),
selected_provider_model_name: "gpt-5".to_string(),
key_ids: vec!["key-2".to_string(), "key-1".to_string()],
});
let sql = query.sql();
assert!(sql.contains("LOWER(pe.api_format) IN ("));
assert!(sql.contains("JSON_SEARCH(LOWER(pak.api_formats), 'one', ?"));
assert!(sql.contains("LOWER(TRIM(p.provider_type)) = 'codex'"));
assert!(sql.contains("AND p.id = ? AND pe.id = ? AND m.id = ?"));
assert!(sql.contains("AND pak.id IN (?, ?)"));
assert!(sql.contains("ORDER BY FIELD(pak.id, ?, ?) ASC, pak.id ASC"));
}
#[tokio::test]
async fn repository_builds_from_lazy_pool() {
@@ -1,6 +1,6 @@
use async_trait::async_trait;
use serde_json::Value;
use sqlx::{mysql::MySqlRow, Row};
use sqlx::{mysql::MySqlRow, Acquire, Row};
use aether_data_contracts::repository::routing_profiles::*;
use aether_data_contracts::DataLayerError;
@@ -56,24 +56,6 @@ impl MysqlRoutingGroupRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
async fn reload_group(&self, id: &str) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
self.find_routing_group(RoutingGroupLookupKey::Id(id)).await
}
async fn find_binding_by_id(
&self,
id: &str,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
let row = sqlx::query(&format!(
"{ROUTING_GROUP_BINDING_SELECT} WHERE id = ? LIMIT 1"
))
.bind(id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_binding_row).transpose()
}
}
#[async_trait]
@@ -170,6 +152,24 @@ impl RoutingGroupWriteRepository for MysqlRoutingGroupRepository {
record: CreateRoutingGroupRecord,
) -> Result<StoredRoutingGroup, DataLayerError> {
let group = StoredRoutingGroup::new(record)?;
let mut connection = self.pool.acquire().await.map_sql_err()?;
sqlx::query("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE")
.execute(&mut *connection)
.await
.map_sql_err()?;
let mut tx = connection.begin().await.map_sql_err()?;
sqlx::query("SELECT id FROM routing_groups ORDER BY id FOR UPDATE")
.fetch_all(&mut *tx)
.await
.map_sql_err()?;
if group.is_system_default {
sqlx::query(
"UPDATE routing_groups SET is_system_default = 0 WHERE is_system_default = 1",
)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
sqlx::query(
r#"
INSERT INTO routing_groups (
@@ -192,9 +192,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
.bind(group.created_at)
.bind(group.updated_at)
.bind(group.published_at)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(group)
}
@@ -203,10 +204,36 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
id: &str,
patch: UpdateRoutingGroupRecord,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
let Some(mut group) = self.reload_group(id).await? else {
let mut connection = self.pool.acquire().await.map_sql_err()?;
sqlx::query("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE")
.execute(&mut *connection)
.await
.map_sql_err()?;
let mut tx = connection.begin().await.map_sql_err()?;
sqlx::query("SELECT id FROM routing_groups ORDER BY id FOR UPDATE")
.fetch_all(&mut *tx)
.await
.map_sql_err()?;
let row = sqlx::query(&format!(
"{ROUTING_GROUP_SELECT} WHERE id = ? LIMIT 1 FOR UPDATE"
))
.bind(id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let Some(mut group) = row.as_ref().map(map_group_row).transpose()? else {
return Ok(None);
};
apply_group_patch(&mut group, patch)?;
if group.is_system_default {
sqlx::query(
"UPDATE routing_groups SET is_system_default = 0 WHERE is_system_default = 1 AND id <> ?",
)
.bind(id)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
sqlx::query(
r#"
UPDATE routing_groups
@@ -233,9 +260,10 @@ WHERE id = ?
.bind(group.updated_at)
.bind(group.published_at)
.bind(id)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(Some(group))
}
@@ -266,6 +294,30 @@ WHERE id = ?
record: CreateRoutingGroupBindingRecord,
) -> Result<StoredRoutingGroupBinding, DataLayerError> {
let binding = StoredRoutingGroupBinding::new(record)?;
let mut connection = self.pool.acquire().await.map_sql_err()?;
sqlx::query("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE")
.execute(&mut *connection)
.await
.map_sql_err()?;
let mut tx = connection.begin().await.map_sql_err()?;
sqlx::query("SELECT id FROM routing_group_bindings ORDER BY id FOR UPDATE")
.fetch_all(&mut *tx)
.await
.map_sql_err()?;
if binding.is_default {
sqlx::query(
r#"
UPDATE routing_group_bindings
SET is_default = 0
WHERE is_default = 1 AND subject_type = ? AND subject_id = ?
"#,
)
.bind(binding_subject_to_database(binding.subject_type))
.bind(&binding.subject_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
sqlx::query(
r#"
INSERT INTO routing_group_bindings (
@@ -283,9 +335,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?)
.bind(binding.allow_explicit_select)
.bind(binding.created_at)
.bind(binding.updated_at)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(binding)
}
@@ -306,10 +359,45 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?)
id: &str,
patch: UpdateRoutingGroupBindingRecord,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
let Some(mut binding) = self.find_binding_by_id(id).await? else {
let mut connection = self.pool.acquire().await.map_sql_err()?;
sqlx::query("SET TRANSACTION ISOLATION LEVEL SERIALIZABLE")
.execute(&mut *connection)
.await
.map_sql_err()?;
let mut tx = connection.begin().await.map_sql_err()?;
sqlx::query("SELECT id FROM routing_group_bindings ORDER BY id FOR UPDATE")
.fetch_all(&mut *tx)
.await
.map_sql_err()?;
let row = sqlx::query(&format!(
"{ROUTING_GROUP_BINDING_SELECT} WHERE id = ? LIMIT 1 FOR UPDATE"
))
.bind(id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let Some(mut binding) = row.as_ref().map(map_binding_row).transpose()? else {
return Ok(None);
};
apply_binding_patch(&mut binding, patch)?;
if binding.is_default {
sqlx::query(
r#"
UPDATE routing_group_bindings
SET is_default = 0
WHERE is_default = 1
AND subject_type = ?
AND subject_id = ?
AND id <> ?
"#,
)
.bind(binding_subject_to_database(binding.subject_type))
.bind(&binding.subject_id)
.bind(id)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
sqlx::query(
r#"
UPDATE routing_group_bindings
@@ -329,9 +417,10 @@ WHERE id = ?
.bind(binding.allow_explicit_select)
.bind(binding.updated_at)
.bind(id)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(Some(binding))
}
@@ -0,0 +1,36 @@
WITH ranked_defaults AS (
SELECT
id,
ROW_NUMBER() OVER (ORDER BY enabled DESC, updated_at DESC, id ASC) AS default_rank
FROM public.routing_groups
WHERE is_system_default = TRUE
)
UPDATE public.routing_groups AS routing_group
SET is_system_default = FALSE
FROM ranked_defaults
WHERE routing_group.id = ranked_defaults.id
AND ranked_defaults.default_rank > 1;
WITH ranked_defaults AS (
SELECT
id,
ROW_NUMBER() OVER (
PARTITION BY subject_type, subject_id
ORDER BY created_at ASC, id ASC
) AS default_rank
FROM public.routing_group_bindings
WHERE is_default = TRUE
)
UPDATE public.routing_group_bindings AS binding
SET is_default = FALSE
FROM ranked_defaults
WHERE binding.id = ranked_defaults.id
AND ranked_defaults.default_rank > 1;
CREATE UNIQUE INDEX IF NOT EXISTS routing_groups_one_system_default_key
ON public.routing_groups (is_system_default)
WHERE is_system_default = TRUE;
CREATE UNIQUE INDEX IF NOT EXISTS routing_group_bindings_subject_default_key
ON public.routing_group_bindings (subject_type, subject_id)
WHERE is_default = TRUE;
@@ -4,10 +4,10 @@ use sqlx::{PgPool, Row};
use std::collections::BTreeSet;
use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
};
use aether_data_contracts::DataLayerError;
@@ -764,6 +764,44 @@ impl SqlxMinimalCandidateSelectionReadRepository {
Ok(dedupe_candidate_selection_rows(rows))
}
pub async fn list_for_exact_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
if query.limit == 0 {
return Ok(Vec::new());
}
let canonical_api_format = normalize_api_format(&query.api_format);
let storage_aliases = api_format_aliases(&canonical_api_format);
let sql_match_aliases =
sql_match_aliases(&api_format_permission_aliases(&canonical_api_format));
let fetch_limit = i64::from(query.offset.saturating_add(query.limit));
let sql = format!("{LIST_FOR_EXACT_API_FORMAT_SQL}\nLIMIT $4\nOFFSET $5");
let mut rows = Vec::new();
for api_format in storage_aliases {
rows.extend(
Self::collect_query_rows(
sqlx::query(sql.as_str())
.bind(api_format)
.bind(sql_match_aliases.clone())
.bind(canonical_api_format.clone())
.bind(fetch_limit)
.bind(0_i64)
.fetch(&self.pool),
map_candidate_selection_row,
)
.await?,
);
}
let mut rows = dedupe_candidate_selection_rows(rows);
sort_candidate_selection_rows(&mut rows, true);
Ok(rows
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
}
pub async fn list_for_exact_api_format_and_global_model(
&self,
api_format: &str,
@@ -1109,6 +1147,24 @@ fn dedupe_candidate_selection_rows(
.collect()
}
fn sort_candidate_selection_rows(
rows: &mut [StoredMinimalCandidateSelectionRow],
include_global_model: bool,
) {
rows.sort_by(|left, right| {
let global_model_order = include_global_model
.then(|| left.global_model_name.cmp(&right.global_model_name))
.unwrap_or(std::cmp::Ordering::Equal);
global_model_order
.then(left.provider_priority.cmp(&right.provider_priority))
.then(left.key_internal_priority.cmp(&right.key_internal_priority))
.then(left.provider_id.cmp(&right.provider_id))
.then(left.endpoint_id.cmp(&right.endpoint_id))
.then(left.key_id.cmp(&right.key_id))
.then(left.model_id.cmp(&right.model_id))
});
}
#[async_trait]
impl MinimalCandidateSelectionReadRepository for SqlxMinimalCandidateSelectionReadRepository {
async fn list_for_exact_api_format(
@@ -1118,6 +1174,13 @@ impl MinimalCandidateSelectionReadRepository for SqlxMinimalCandidateSelectionRe
Self::list_for_exact_api_format(self, api_format).await
}
async fn list_for_exact_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Self::list_for_exact_api_format_page(self, query).await
}
async fn list_for_exact_api_format_and_global_model(
&self,
api_format: &str,
@@ -63,24 +63,6 @@ impl PostgresRoutingGroupRepository {
pub fn new(pool: PgPool) -> Self {
Self { pool }
}
async fn reload_group(&self, id: &str) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
self.find_routing_group(RoutingGroupLookupKey::Id(id)).await
}
async fn find_binding_by_id(
&self,
id: &str,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
let row = sqlx::query(&format!(
"{ROUTING_GROUP_BINDING_SELECT} WHERE id = $1 LIMIT 1"
))
.bind(id)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref().map(map_binding_row).transpose()
}
}
#[async_trait]
@@ -180,6 +162,19 @@ impl RoutingGroupWriteRepository for PostgresRoutingGroupRepository {
record: CreateRoutingGroupRecord,
) -> Result<StoredRoutingGroup, DataLayerError> {
let group = StoredRoutingGroup::new(record)?;
let mut tx = self.pool.begin().await.map_postgres_err()?;
if group.is_system_default {
sqlx::query("LOCK TABLE routing_groups IN SHARE ROW EXCLUSIVE MODE")
.execute(&mut *tx)
.await
.map_postgres_err()?;
sqlx::query(
"UPDATE routing_groups SET is_system_default = FALSE WHERE is_system_default = TRUE",
)
.execute(&mut *tx)
.await
.map_postgres_err()?;
}
sqlx::query(
r#"
INSERT INTO routing_groups (
@@ -199,9 +194,10 @@ VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)
.bind(group.created_at)
.bind(group.updated_at)
.bind(group.published_at)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_postgres_err()?;
tx.commit().await.map_postgres_err()?;
Ok(group)
}
@@ -210,10 +206,31 @@ VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)
id: &str,
patch: UpdateRoutingGroupRecord,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
let Some(mut group) = self.reload_group(id).await? else {
let mut tx = self.pool.begin().await.map_postgres_err()?;
sqlx::query("LOCK TABLE routing_groups IN SHARE ROW EXCLUSIVE MODE")
.execute(&mut *tx)
.await
.map_postgres_err()?;
let row = sqlx::query(&format!(
"{ROUTING_GROUP_SELECT} WHERE id = $1 LIMIT 1 FOR UPDATE"
))
.bind(id)
.fetch_optional(&mut *tx)
.await
.map_postgres_err()?;
let Some(mut group) = row.as_ref().map(map_group_row).transpose()? else {
return Ok(None);
};
apply_group_patch(&mut group, patch)?;
if group.is_system_default {
sqlx::query(
"UPDATE routing_groups SET is_system_default = FALSE WHERE is_system_default = TRUE AND id <> $1",
)
.bind(id)
.execute(&mut *tx)
.await
.map_postgres_err()?;
}
sqlx::query(
r#"
UPDATE routing_groups
@@ -237,9 +254,10 @@ WHERE id = $1
.bind(group.version)
.bind(group.updated_at)
.bind(group.published_at)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_postgres_err()?;
tx.commit().await.map_postgres_err()?;
Ok(Some(group))
}
@@ -270,6 +288,25 @@ WHERE id = $1
record: CreateRoutingGroupBindingRecord,
) -> Result<StoredRoutingGroupBinding, DataLayerError> {
let binding = StoredRoutingGroupBinding::new(record)?;
let mut tx = self.pool.begin().await.map_postgres_err()?;
if binding.is_default {
sqlx::query("LOCK TABLE routing_group_bindings IN SHARE ROW EXCLUSIVE MODE")
.execute(&mut *tx)
.await
.map_postgres_err()?;
sqlx::query(
r#"
UPDATE routing_group_bindings
SET is_default = FALSE
WHERE is_default = TRUE AND subject_type = $1 AND subject_id = $2
"#,
)
.bind(binding_subject_to_database(binding.subject_type))
.bind(&binding.subject_id)
.execute(&mut *tx)
.await
.map_postgres_err()?;
}
sqlx::query(
r#"
INSERT INTO routing_group_bindings (
@@ -287,9 +324,10 @@ VALUES ($1, $2, $3, $4, $5, $6, $7, $8)
.bind(binding.allow_explicit_select)
.bind(binding.created_at)
.bind(binding.updated_at)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_postgres_err()?;
tx.commit().await.map_postgres_err()?;
Ok(binding)
}
@@ -310,10 +348,40 @@ VALUES ($1, $2, $3, $4, $5, $6, $7, $8)
id: &str,
patch: UpdateRoutingGroupBindingRecord,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
let Some(mut binding) = self.find_binding_by_id(id).await? else {
let mut tx = self.pool.begin().await.map_postgres_err()?;
sqlx::query("LOCK TABLE routing_group_bindings IN SHARE ROW EXCLUSIVE MODE")
.execute(&mut *tx)
.await
.map_postgres_err()?;
let row = sqlx::query(&format!(
"{ROUTING_GROUP_BINDING_SELECT} WHERE id = $1 LIMIT 1 FOR UPDATE"
))
.bind(id)
.fetch_optional(&mut *tx)
.await
.map_postgres_err()?;
let Some(mut binding) = row.as_ref().map(map_binding_row).transpose()? else {
return Ok(None);
};
apply_binding_patch(&mut binding, patch)?;
if binding.is_default {
sqlx::query(
r#"
UPDATE routing_group_bindings
SET is_default = FALSE
WHERE is_default = TRUE
AND subject_type = $1
AND subject_id = $2
AND id <> $3
"#,
)
.bind(binding_subject_to_database(binding.subject_type))
.bind(&binding.subject_id)
.bind(id)
.execute(&mut *tx)
.await
.map_postgres_err()?;
}
sqlx::query(
r#"
UPDATE routing_group_bindings
@@ -333,9 +401,10 @@ WHERE id = $1
.bind(binding.is_default)
.bind(binding.allow_explicit_select)
.bind(binding.updated_at)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_postgres_err()?;
tx.commit().await.map_postgres_err()?;
Ok(Some(binding))
}
@@ -0,0 +1,38 @@
UPDATE routing_groups
SET is_system_default = 0
WHERE id IN (
SELECT id
FROM (
SELECT
id,
ROW_NUMBER() OVER (ORDER BY enabled DESC, updated_at DESC, id ASC) AS default_rank
FROM routing_groups
WHERE is_system_default = 1
) AS ranked_defaults
WHERE default_rank > 1
);
UPDATE routing_group_bindings
SET is_default = 0
WHERE id IN (
SELECT id
FROM (
SELECT
id,
ROW_NUMBER() OVER (
PARTITION BY subject_type, subject_id
ORDER BY created_at ASC, id ASC
) AS default_rank
FROM routing_group_bindings
WHERE is_default = 1
) AS ranked_defaults
WHERE default_rank > 1
);
CREATE UNIQUE INDEX IF NOT EXISTS routing_groups_one_system_default_key
ON routing_groups (is_system_default)
WHERE is_system_default = 1;
CREATE UNIQUE INDEX IF NOT EXISTS routing_group_bindings_subject_default_key
ON routing_group_bindings (subject_type, subject_id)
WHERE is_default = 1;
@@ -4,10 +4,10 @@ use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
};
use aether_data_contracts::DataLayerError;
@@ -68,6 +68,9 @@ WHERE p.is_active = 1
AND gm.is_active = 1
"#;
const REQUESTED_MODEL_RAW_PAGE_SIZE: u32 = 256;
const REQUESTED_MODEL_RAW_SCAN_LIMIT: u32 = 2048;
#[derive(Debug, Clone)]
pub struct SqliteMinimalCandidateSelectionReadRepository {
pool: SqlitePool,
@@ -77,7 +80,6 @@ pub struct SqliteMinimalCandidateSelectionReadRepository {
struct CandidateSelectionRow {
row: StoredMinimalCandidateSelectionRow,
key_auth_config: Option<String>,
key_last_used_at_unix_secs: Option<u64>,
}
#[derive(Debug, Clone, Copy)]
@@ -99,6 +101,58 @@ struct SqlPage {
offset: i64,
}
#[derive(Debug)]
struct ExactPageAccumulator<T> {
rows: Vec<T>,
offset: usize,
limit: usize,
target_len: usize,
}
impl<T> ExactPageAccumulator<T> {
fn new(offset: u32, limit: u32) -> Self {
let offset = usize::try_from(offset).unwrap_or(usize::MAX);
let limit = usize::try_from(limit).unwrap_or(usize::MAX);
Self {
rows: Vec::new(),
offset,
limit,
target_len: offset.saturating_add(limit),
}
}
fn is_full(&self) -> bool {
self.rows.len() >= self.target_len
}
fn push_matching<I, F>(&mut self, rows: I, mut predicate: F)
where
I: IntoIterator<Item = T>,
F: FnMut(&T) -> bool,
{
let remaining = self.target_len.saturating_sub(self.rows.len());
self.rows.extend(
rows.into_iter()
.filter(|row| predicate(row))
.take(remaining),
);
}
fn into_page(self) -> Vec<T> {
self.rows
.into_iter()
.skip(self.offset)
.take(self.limit)
.collect()
}
}
#[derive(Debug)]
struct RequestedModelRawPage {
rows: Vec<CandidateSelectionRow>,
raw_len: u32,
}
impl SqliteMinimalCandidateSelectionReadRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
@@ -148,44 +202,7 @@ impl SqliteMinimalCandidateSelectionReadRepository {
);
}
}
builder.push(
r#"
),
pool_rows AS (
SELECT candidate.*
FROM candidate_rows candidate
WHERE candidate.provider_pool_enabled = 1
AND NOT EXISTS (
SELECT 1
FROM candidate_rows other
WHERE other.provider_pool_enabled = 1
AND other.provider_id = candidate.provider_id
AND other.endpoint_id = candidate.endpoint_id
AND other.model_id = candidate.model_id
AND (
other.key_internal_priority < candidate.key_internal_priority
OR (
other.key_internal_priority = candidate.key_internal_priority
AND other.key_id < candidate.key_id
)
)
)
),
selected_rows AS (
SELECT * FROM candidate_rows WHERE provider_pool_enabled = 0
UNION ALL
SELECT * FROM pool_rows
)
SELECT * FROM selected_rows
"#,
);
push_selected_rows_order(&mut builder, order);
if let Some(page) = page {
builder.push(" LIMIT ");
builder.push_bind(page.limit);
builder.push(" OFFSET ");
builder.push_bind(page.offset);
}
push_selected_rows_query_tail(&mut builder, order, page);
let query_rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let mut items = query_rows
@@ -211,6 +228,42 @@ SELECT * FROM selected_rows
};
Ok(dedupe_candidate_selection_rows(rows))
}
async fn load_requested_model_raw_page(
&self,
api_format: &str,
requested_model_name: &str,
page: SqlPage,
) -> Result<RequestedModelRawPage, DataLayerError> {
let canonical_api_format = normalize_api_format(api_format);
let storage_aliases = api_format_aliases(&canonical_api_format);
let match_aliases =
sql_match_aliases(&api_format_permission_aliases(&canonical_api_format));
let mut builder = QueryBuilder::<Sqlite>::new("WITH candidate_rows AS (");
builder.push(CANDIDATE_SELECTION_COLUMNS);
push_candidate_sql_filters_for_aliases(
&mut builder,
&storage_aliases,
&match_aliases,
&canonical_api_format,
);
push_requested_model_sql_filter(&mut builder, requested_model_name, &match_aliases);
push_selected_rows_query_tail(&mut builder, SelectedRowsOrder::WithGlobalModel, Some(page));
let query_rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let raw_len = u32::try_from(query_rows.len()).unwrap_or(u32::MAX);
let mut rows = query_rows
.iter()
.map(map_candidate_selection_row)
.collect::<Result<Vec<_>, _>>()?;
rows.retain(|item| {
api_format_matches(&item.row.endpoint_api_format, &canonical_api_format)
&& item.row.key_supports_api_format(&canonical_api_format)
&& key_auth_channel_matches(item, &canonical_api_format)
});
Ok(RequestedModelRawPage { rows, raw_len })
}
}
#[async_trait]
@@ -222,6 +275,33 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection
self.selected_rows_for_api_format(api_format).await
}
async fn list_for_exact_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
if query.limit == 0 {
return Ok(Vec::new());
}
let fetch_limit = query.offset.saturating_add(query.limit);
let mut rows = self
.load_selected_rows_for_api_format(
&query.api_format,
SelectedRowsFilter::None,
SelectedRowsOrder::WithGlobalModel,
Some(SqlPage {
limit: i64::from(fetch_limit),
offset: 0,
}),
)
.await?;
sort_candidate_selection_rows(&mut rows, true);
Ok(rows
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
}
async fn list_for_exact_api_format_and_global_model(
&self,
api_format: &str,
@@ -254,28 +334,58 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.load_selected_rows_for_api_format(
&query.api_format,
SelectedRowsFilter::RequestedModel(&query.requested_model_name),
SelectedRowsOrder::WithGlobalModel,
Some(SqlPage {
limit: i64::from(query.limit.max(1)),
offset: i64::from(query.offset),
}),
)
.await
if query.limit == 0 {
return Ok(Vec::new());
}
let mut exact_page = ExactPageAccumulator::new(query.offset, query.limit);
let mut raw_offset = 0_u32;
while !exact_page.is_full() && raw_offset < REQUESTED_MODEL_RAW_SCAN_LIMIT {
let raw_limit = REQUESTED_MODEL_RAW_PAGE_SIZE
.min(REQUESTED_MODEL_RAW_SCAN_LIMIT.saturating_sub(raw_offset));
let raw_page = self
.load_requested_model_raw_page(
&query.api_format,
&query.requested_model_name,
SqlPage {
limit: i64::from(raw_limit),
offset: i64::from(raw_offset),
},
)
.await?;
exact_page.push_matching(raw_page.rows, |item| {
row_matches_requested_model(
&item.row,
&query.requested_model_name,
&query.api_format,
)
});
raw_offset = raw_offset.saturating_add(raw_page.raw_len);
if raw_page.raw_len < raw_limit || raw_page.raw_len == 0 {
break;
}
}
let rows = exact_page
.into_page()
.into_iter()
.map(|item| item.row)
.collect();
Ok(dedupe_candidate_selection_rows(rows))
}
async fn list_pool_key_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
if query.limit == 0 {
return Ok(Vec::new());
}
let canonical_api_format = normalize_api_format(&query.api_format);
let storage_aliases = api_format_aliases(&canonical_api_format);
let match_aliases =
sql_match_aliases(&api_format_permission_aliases(&canonical_api_format));
let mut rows = Vec::<CandidateSelectionRow>::new();
let page_in_sql = !matches!(query.order, StoredPoolKeyCandidateOrder::LoadBalance { .. });
for storage_api_format in storage_aliases {
let mut builder = QueryBuilder::<Sqlite>::new(CANDIDATE_SELECTION_COLUMNS);
@@ -286,15 +396,11 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection
builder.push_bind(&query.endpoint_id);
builder.push(" AND m.id = ");
builder.push_bind(&query.model_id);
if page_in_sql {
push_pool_key_order(&mut builder, &query.order);
builder.push(" LIMIT ");
builder.push_bind(i64::from(query.limit.max(1)));
builder.push(" OFFSET ");
builder.push_bind(i64::from(query.offset));
} else {
builder.push(" ORDER BY pak.id ASC");
}
push_pool_key_order(&mut builder, &query.order);
builder.push(" LIMIT ");
builder.push_bind(i64::from(query.limit));
builder.push(" OFFSET ");
builder.push_bind(i64::from(query.offset));
let query_rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let mut items = query_rows
@@ -309,20 +415,9 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection
rows.extend(items);
}
if page_in_sql {
Ok(dedupe_candidate_selection_rows(
rows.into_iter().map(|item| item.row).collect(),
))
} else {
Ok(dedupe_candidate_selection_rows(
sort_pool_key_rows(rows, &query.order)
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.map(|item| item.row)
.collect(),
))
}
Ok(dedupe_candidate_selection_rows(
rows.into_iter().map(|item| item.row).collect(),
))
}
async fn list_pool_key_rows_for_group_key_ids(
@@ -411,6 +506,19 @@ fn push_candidate_sql_filters(
push_key_auth_channel_sql_filter(builder, storage_api_format);
}
fn push_candidate_sql_filters_for_aliases(
builder: &mut QueryBuilder<'_, Sqlite>,
storage_api_formats: &[String],
match_aliases: &[String],
requested_api_format: &str,
) {
builder.push(" AND LOWER(COALESCE(pe.api_format, '')) IN (");
push_bind_list(builder, storage_api_formats);
builder.push(")");
push_key_api_format_sql_filter(builder, match_aliases);
push_key_auth_channel_sql_filter(builder, requested_api_format);
}
fn push_key_api_format_sql_filter(
builder: &mut QueryBuilder<'_, Sqlite>,
match_aliases: &[String],
@@ -626,6 +734,51 @@ fn push_requested_model_sql_filter(
);
}
fn push_selected_rows_query_tail(
builder: &mut QueryBuilder<'_, Sqlite>,
order: SelectedRowsOrder,
page: Option<SqlPage>,
) {
builder.push(
r#"
),
pool_rows AS (
SELECT candidate.*
FROM candidate_rows candidate
WHERE candidate.provider_pool_enabled = 1
AND NOT EXISTS (
SELECT 1
FROM candidate_rows other
WHERE other.provider_pool_enabled = 1
AND other.provider_id = candidate.provider_id
AND other.endpoint_id = candidate.endpoint_id
AND other.model_id = candidate.model_id
AND (
other.key_internal_priority < candidate.key_internal_priority
OR (
other.key_internal_priority = candidate.key_internal_priority
AND other.key_id < candidate.key_id
)
)
)
),
selected_rows AS (
SELECT * FROM candidate_rows WHERE provider_pool_enabled = 0
UNION ALL
SELECT * FROM pool_rows
)
SELECT * FROM selected_rows
"#,
);
push_selected_rows_order(builder, order);
if let Some(page) = page {
builder.push(" LIMIT ");
builder.push_bind(page.limit);
builder.push(" OFFSET ");
builder.push_bind(page.offset);
}
}
fn push_selected_rows_order(builder: &mut QueryBuilder<'_, Sqlite>, order: SelectedRowsOrder) {
builder.push(" ORDER BY ");
if matches!(order, SelectedRowsOrder::WithGlobalModel) {
@@ -660,12 +813,52 @@ fn push_pool_key_order(
);
}
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
let _ = seed;
builder.push(" ORDER BY pak.id ASC");
push_seeded_pool_key_order(builder, seed);
}
}
}
fn push_seeded_pool_key_order(builder: &mut QueryBuilder<'_, Sqlite>, seed: &str) {
const HEX_DIGITS: &[u8; 16] = b"0123456789abcdef";
const PLACEHOLDERS: &[u8; 16] = b"ghijklmnopqrstuv";
let mut digits_by_rank = HEX_DIGITS.map(char::from);
digits_by_rank.sort_by(|left, right| {
stable_pool_key_hash(seed, &left.to_string())
.cmp(&stable_pool_key_hash(seed, &right.to_string()))
.then(left.cmp(right))
});
let mut rank_by_digit = ['0'; 16];
for (rank, digit) in digits_by_rank.into_iter().enumerate() {
let digit_index = digit
.to_digit(16)
.expect("seeded pool-key rank input must be hexadecimal")
as usize;
rank_by_digit[digit_index] = char::from(HEX_DIGITS[rank]);
}
// SQLite has no built-in hash; placeholders avoid cascading replacements
// while remapping every key-id nibble to a seed-derived rank.
builder.push(" ORDER BY ");
for _ in 0..(HEX_DIGITS.len() + PLACEHOLDERS.len()) {
builder.push("replace(");
}
builder.push("lower(hex(pak.id))");
for (digit, placeholder) in HEX_DIGITS.iter().zip(PLACEHOLDERS) {
builder.push(format!(
", '{}', '{}')",
char::from(*digit),
char::from(*placeholder)
));
}
for (placeholder, rank) in PLACEHOLDERS.iter().zip(rank_by_digit) {
builder.push(format!(", '{}', ", char::from(*placeholder)));
builder.push_bind(rank.to_string());
builder.push(")");
}
builder.push(" ASC, pak.id ASC");
}
fn push_bind_list(builder: &mut QueryBuilder<'_, Sqlite>, values: &[String]) {
let mut separated = builder.separated(", ");
for value in values {
@@ -673,52 +866,6 @@ fn push_bind_list(builder: &mut QueryBuilder<'_, Sqlite>, values: &[String]) {
}
}
fn sort_pool_key_rows(
mut rows: Vec<CandidateSelectionRow>,
order: &StoredPoolKeyCandidateOrder,
) -> Vec<CandidateSelectionRow> {
rows.sort_by(|left, right| match order {
StoredPoolKeyCandidateOrder::InternalPriority => compare_pool_key_internal(left, right),
StoredPoolKeyCandidateOrder::Lru => left
.key_last_used_at_unix_secs
.cmp(&right.key_last_used_at_unix_secs)
.then_with(|| compare_pool_key_internal(left, right)),
StoredPoolKeyCandidateOrder::CacheAffinity => right
.key_last_used_at_unix_secs
.cmp(&left.key_last_used_at_unix_secs)
.then_with(|| compare_pool_key_internal(left, right)),
StoredPoolKeyCandidateOrder::SingleAccount => left
.row
.key_internal_priority
.cmp(&right.row.key_internal_priority)
.then_with(|| {
right
.key_last_used_at_unix_secs
.cmp(&left.key_last_used_at_unix_secs)
})
.then(left.row.key_id.cmp(&right.row.key_id)),
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
stable_pool_key_hash(seed.as_str(), left.row.key_id.as_str())
.cmp(&stable_pool_key_hash(
seed.as_str(),
right.row.key_id.as_str(),
))
.then(left.row.key_id.cmp(&right.row.key_id))
}
});
rows
}
fn compare_pool_key_internal(
left: &CandidateSelectionRow,
right: &CandidateSelectionRow,
) -> std::cmp::Ordering {
left.row
.key_internal_priority
.cmp(&right.row.key_internal_priority)
.then(left.row.key_id.cmp(&right.row.key_id))
}
fn stable_pool_key_hash(seed: &str, key_id: &str) -> u64 {
let mut hash = 0xcbf29ce484222325u64;
for byte in seed
@@ -875,6 +1022,24 @@ fn dedupe_candidate_selection_rows(
.collect()
}
fn sort_candidate_selection_rows(
rows: &mut [StoredMinimalCandidateSelectionRow],
include_global_model: bool,
) {
rows.sort_by(|left, right| {
let global_model_order = include_global_model
.then(|| left.global_model_name.cmp(&right.global_model_name))
.unwrap_or(std::cmp::Ordering::Equal);
global_model_order
.then(left.provider_priority.cmp(&right.provider_priority))
.then(left.key_internal_priority.cmp(&right.key_internal_priority))
.then(left.provider_id.cmp(&right.provider_id))
.then(left.endpoint_id.cmp(&right.endpoint_id))
.then(left.key_id.cmp(&right.key_id))
.then(left.model_id.cmp(&right.model_id))
});
}
fn map_candidate_selection_row(row: &SqliteRow) -> Result<CandidateSelectionRow, DataLayerError> {
let _provider_config = parse_json(row.try_get("provider_config").ok().flatten())?;
let global_model_config = parse_json(row.try_get("global_model_config").ok().flatten())?;
@@ -931,10 +1096,6 @@ fn map_candidate_selection_row(row: &SqliteRow) -> Result<CandidateSelectionRow,
model_is_available: row.try_get("model_is_available").map_sql_err()?,
},
key_auth_config: row.try_get("key_auth_config").map_sql_err()?,
key_last_used_at_unix_secs: row
.try_get::<Option<i64>, _>("key_last_used_at_unix_secs")
.map_sql_err()?
.and_then(|value| u64::try_from(value).ok()),
})
}
@@ -1177,8 +1338,9 @@ fn sql_match_aliases(api_formats: &[String]) -> Vec<String> {
#[cfg(test)]
mod tests {
use super::{
push_key_auth_channel_sql_filter, vertex_key_auth_channel_matches,
SqliteMinimalCandidateSelectionReadRepository,
push_key_auth_channel_sql_filter, push_pool_key_order, vertex_key_auth_channel_matches,
ExactPageAccumulator, SqliteMinimalCandidateSelectionReadRepository,
REQUESTED_MODEL_RAW_SCAN_LIMIT,
};
use crate::run_migrations;
use aether_data_contracts::repository::candidate_selection::{
@@ -1216,6 +1378,57 @@ mod tests {
assert!(vertex_clause.contains("gemini:embedding"));
}
#[test]
fn exact_page_accumulator_continues_after_coarse_false_positives() {
let mut accumulator = ExactPageAccumulator::new(1, 2);
accumulator.push_matching(vec![("coarse-1", false), ("coarse-2", false)], |row| row.1);
assert!(!accumulator.is_full());
accumulator.push_matching(
vec![
("exact-1", true),
("coarse-3", false),
("exact-2", true),
("exact-3", true),
],
|row| row.1,
);
assert!(accumulator.is_full());
assert_eq!(
accumulator.into_page(),
vec![("exact-2", true), ("exact-3", true)]
);
assert_eq!(REQUESTED_MODEL_RAW_SCAN_LIMIT, 2048);
}
#[test]
fn load_balance_pool_key_order_is_seeded_and_pageable_in_sql() {
let sql_for_seed = |seed: &str| {
let mut builder =
sqlx::QueryBuilder::<sqlx::Sqlite>::new("SELECT pak.id FROM provider_api_keys pak");
push_pool_key_order(
&mut builder,
&StoredPoolKeyCandidateOrder::LoadBalance {
seed: seed.to_string(),
},
);
builder.push(" LIMIT ");
builder.push_bind(64_i64);
builder.push(" OFFSET ");
builder.push_bind(128_i64);
builder.sql().to_string()
};
let first_seed_sql = sql_for_seed("seed-a");
let second_seed_sql = sql_for_seed("seed-b");
assert!(first_seed_sql.contains("lower(hex(pak.id))"));
assert!(first_seed_sql.contains("ASC, pak.id ASC LIMIT ? OFFSET ?"));
assert_eq!(first_seed_sql, second_seed_sql);
assert_eq!(first_seed_sql.matches('?').count(), 18);
}
#[tokio::test]
async fn sqlite_repository_reads_candidate_selection_rows() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
@@ -1336,6 +1549,97 @@ mod tests {
assert_eq!(search_rows[0].endpoint_api_format, "openai:search");
}
#[tokio::test]
async fn sqlite_requested_model_page_crosses_coarse_false_positive_windows() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_requested_model_pagination(&pool).await;
let repository = SqliteMinimalCandidateSelectionReadRepository::new(pool);
let rows = repository
.list_for_exact_api_format_and_requested_model_page(
&StoredRequestedModelCandidateRowsQuery {
api_format: "openai:chat".to_string(),
requested_model_name: "sqlite-page-target".to_string(),
offset: 1,
limit: 1,
},
)
.await
.expect("requested model page should cross the coarse-only window");
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].model_id, "model-pagination-exact-1");
}
#[tokio::test]
async fn sqlite_load_balance_pool_key_pages_use_stable_seeded_order() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_migrations(&pool)
.await
.expect("sqlite migrations should run");
seed_candidate_selection(&pool).await;
let repository = SqliteMinimalCandidateSelectionReadRepository::new(pool);
let load_page = |seed: &str, offset, limit| StoredPoolKeyCandidateRowsQuery {
api_format: "openai:chat".to_string(),
provider_id: "provider-1".to_string(),
endpoint_id: "endpoint-1".to_string(),
model_id: "model-1".to_string(),
selected_provider_model_name: "provider-model".to_string(),
order: StoredPoolKeyCandidateOrder::LoadBalance {
seed: seed.to_string(),
},
offset,
limit,
};
let seed_a_first = repository
.list_pool_key_rows_for_group(&load_page("seed-a", 0, 1))
.await
.expect("first load-balance page should load");
let seed_a_second = repository
.list_pool_key_rows_for_group(&load_page("seed-a", 1, 1))
.await
.expect("second load-balance page should load");
let seed_a_replay = repository
.list_pool_key_rows_for_group(&load_page("seed-a", 0, 2))
.await
.expect("replayed load-balance window should load");
let seed_b = repository
.list_pool_key_rows_for_group(&load_page("seed-b", 0, 2))
.await
.expect("alternate load-balance seed should load");
let seed_a_pages = seed_a_first
.iter()
.chain(&seed_a_second)
.map(|row| row.key_id.as_str())
.collect::<Vec<_>>();
let seed_a_replay = seed_a_replay
.iter()
.map(|row| row.key_id.as_str())
.collect::<Vec<_>>();
let seed_b = seed_b
.iter()
.map(|row| row.key_id.as_str())
.collect::<Vec<_>>();
assert_eq!(seed_a_pages, vec!["key-1", "key-2"]);
assert_eq!(seed_a_pages, seed_a_replay);
assert_eq!(seed_b, vec!["key-2", "key-1"]);
}
async fn seed_candidate_selection(pool: &sqlx::SqlitePool) {
sqlx::query(
r#"
@@ -1455,4 +1759,94 @@ VALUES (
.await
.expect("candidate selection rows should seed");
}
async fn seed_requested_model_pagination(pool: &sqlx::SqlitePool) {
sqlx::query(
r#"
INSERT INTO providers (
id, name, provider_type, provider_priority, is_active, created_at, updated_at
)
VALUES (
'provider-pagination', 'Pagination Provider', 'custom', 10, 1, 1, 1
);
INSERT INTO provider_endpoints (
id, provider_id, name, base_url, api_format, is_active, created_at, updated_at
)
VALUES (
'endpoint-pagination', 'provider-pagination', 'Pagination Endpoint',
'https://example.test', 'openai:chat', 1, 1, 1
);
INSERT INTO provider_api_keys (
id, provider_id, name, auth_type, api_formats, internal_priority,
is_active, created_at, updated_at
)
VALUES (
'key-pagination', 'provider-pagination', 'Pagination Key', 'api_key',
'["openai:chat"]', 10, 1, 1, 1
);
WITH RECURSIVE sequence(value) AS (
SELECT 0
UNION ALL
SELECT value + 1 FROM sequence WHERE value < 255
)
INSERT INTO global_models (
id, name, is_active, created_at, updated_at
)
SELECT
printf('global-pagination-false-%03d', value),
printf('a-pagination-false-%03d', value),
1, 1, 1
FROM sequence;
INSERT INTO global_models (
id, name, is_active, created_at, updated_at
)
VALUES
('global-pagination-exact-0', 'z-pagination-exact-0', 1, 1, 1),
('global-pagination-exact-1', 'z-pagination-exact-1', 1, 1, 1);
WITH RECURSIVE sequence(value) AS (
SELECT 0
UNION ALL
SELECT value + 1 FROM sequence WHERE value < 255
)
INSERT INTO models (
id, provider_id, global_model_id, provider_model_name, provider_model_mappings,
is_active, is_available, created_at, updated_at
)
SELECT
printf('model-pagination-false-%03d', value),
'provider-pagination',
printf('global-pagination-false-%03d', value),
'upstream-false',
'[{"name":"sqlite-page-target-noise","api_formats":["openai:chat"],"priority":1}]',
1, 1, 1, 1
FROM sequence;
INSERT INTO models (
id, provider_id, global_model_id, provider_model_name, provider_model_mappings,
is_active, is_available, created_at, updated_at
)
VALUES
(
'model-pagination-exact-0', 'provider-pagination', 'global-pagination-exact-0',
'upstream-exact-0',
'[{"name":"sqlite-page-target","api_formats":["openai:chat"],"priority":1}]',
1, 1, 1, 1
),
(
'model-pagination-exact-1', 'provider-pagination', 'global-pagination-exact-1',
'upstream-exact-1',
'[{"name":"sqlite-page-target","api_formats":["openai:chat"],"priority":1}]',
1, 1, 1, 1
);
"#,
)
.execute(pool)
.await
.expect("requested model pagination rows should seed");
}
}
@@ -56,24 +56,6 @@ impl SqliteRoutingGroupRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn reload_group(&self, id: &str) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
self.find_routing_group(RoutingGroupLookupKey::Id(id)).await
}
async fn find_binding_by_id(
&self,
id: &str,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
let row = sqlx::query(&format!(
"{ROUTING_GROUP_BINDING_SELECT} WHERE id = ? LIMIT 1"
))
.bind(id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_binding_row).transpose()
}
}
#[async_trait]
@@ -170,6 +152,19 @@ impl RoutingGroupWriteRepository for SqliteRoutingGroupRepository {
record: CreateRoutingGroupRecord,
) -> Result<StoredRoutingGroup, DataLayerError> {
let group = StoredRoutingGroup::new(record)?;
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("UPDATE routing_groups SET is_system_default = is_system_default WHERE 0")
.execute(&mut *tx)
.await
.map_sql_err()?;
if group.is_system_default {
sqlx::query(
"UPDATE routing_groups SET is_system_default = 0 WHERE is_system_default = 1",
)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
sqlx::query(
r#"
INSERT INTO routing_groups (
@@ -192,9 +187,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
.bind(group.created_at)
.bind(group.updated_at)
.bind(group.published_at)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(group)
}
@@ -203,10 +199,29 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
id: &str,
patch: UpdateRoutingGroupRecord,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
let Some(mut group) = self.reload_group(id).await? else {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("UPDATE routing_groups SET is_system_default = is_system_default WHERE 0")
.execute(&mut *tx)
.await
.map_sql_err()?;
let row = sqlx::query(&format!("{ROUTING_GROUP_SELECT} WHERE id = ? LIMIT 1"))
.bind(id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let Some(mut group) = row.as_ref().map(map_group_row).transpose()? else {
return Ok(None);
};
apply_group_patch(&mut group, patch)?;
if group.is_system_default {
sqlx::query(
"UPDATE routing_groups SET is_system_default = 0 WHERE is_system_default = 1 AND id <> ?",
)
.bind(id)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
sqlx::query(
r#"
UPDATE routing_groups
@@ -233,9 +248,10 @@ WHERE id = ?
.bind(group.updated_at)
.bind(group.published_at)
.bind(id)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(Some(group))
}
@@ -266,6 +282,25 @@ WHERE id = ?
record: CreateRoutingGroupBindingRecord,
) -> Result<StoredRoutingGroupBinding, DataLayerError> {
let binding = StoredRoutingGroupBinding::new(record)?;
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("UPDATE routing_group_bindings SET is_default = is_default WHERE 0")
.execute(&mut *tx)
.await
.map_sql_err()?;
if binding.is_default {
sqlx::query(
r#"
UPDATE routing_group_bindings
SET is_default = 0
WHERE is_default = 1 AND subject_type = ? AND subject_id = ?
"#,
)
.bind(binding_subject_to_database(binding.subject_type))
.bind(&binding.subject_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
sqlx::query(
r#"
INSERT INTO routing_group_bindings (
@@ -283,9 +318,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?)
.bind(binding.allow_explicit_select)
.bind(binding.created_at)
.bind(binding.updated_at)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(binding)
}
@@ -306,10 +342,40 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?)
id: &str,
patch: UpdateRoutingGroupBindingRecord,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
let Some(mut binding) = self.find_binding_by_id(id).await? else {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("UPDATE routing_group_bindings SET is_default = is_default WHERE 0")
.execute(&mut *tx)
.await
.map_sql_err()?;
let row = sqlx::query(&format!(
"{ROUTING_GROUP_BINDING_SELECT} WHERE id = ? LIMIT 1"
))
.bind(id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let Some(mut binding) = row.as_ref().map(map_binding_row).transpose()? else {
return Ok(None);
};
apply_binding_patch(&mut binding, patch)?;
if binding.is_default {
sqlx::query(
r#"
UPDATE routing_group_bindings
SET is_default = 0
WHERE is_default = 1
AND subject_type = ?
AND subject_id = ?
AND id <> ?
"#,
)
.bind(binding_subject_to_database(binding.subject_type))
.bind(&binding.subject_id)
.bind(id)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
sqlx::query(
r#"
UPDATE routing_group_bindings
@@ -329,9 +395,10 @@ WHERE id = ?
.bind(binding.allow_explicit_select)
.bind(binding.updated_at)
.bind(id)
.execute(&self.pool)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(Some(binding))
}
@@ -526,4 +593,255 @@ mod tests {
1
);
}
#[tokio::test]
async fn sqlite_keeps_system_and_subject_defaults_unique() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
let repository = SqliteRoutingGroupRepository::new(pool);
for (id, is_system_default) in [("group-1", true), ("group-2", true), ("group-3", false)] {
repository
.create_routing_group(group_record(id, is_system_default))
.await
.expect("group should create");
}
assert_eq!(system_default_ids(&repository).await, vec!["group-2"]);
repository
.update_routing_group(
"group-1",
UpdateRoutingGroupRecord {
is_system_default: Some(true),
updated_at: 2,
..UpdateRoutingGroupRecord::default()
},
)
.await
.expect("group should update");
assert_eq!(system_default_ids(&repository).await, vec!["group-1"]);
repository
.create_routing_group_binding(binding_record("binding-1", "group-1", "subject-1", true))
.await
.expect("binding should create");
repository
.create_routing_group_binding(binding_record("binding-2", "group-2", "subject-1", true))
.await
.expect("binding should create");
repository
.create_routing_group_binding(binding_record("binding-3", "group-3", "subject-2", true))
.await
.expect("binding should create");
assert_eq!(
default_binding_ids(&repository, "subject-1").await,
vec!["binding-2"]
);
assert_eq!(
default_binding_ids(&repository, "subject-2").await,
vec!["binding-3"]
);
repository
.update_routing_group_binding(
"binding-1",
UpdateRoutingGroupBindingRecord {
is_default: Some(true),
updated_at: 2,
..UpdateRoutingGroupBindingRecord::default()
},
)
.await
.expect("binding should update");
assert_eq!(
default_binding_ids(&repository, "subject-1").await,
vec!["binding-1"]
);
assert_eq!(
default_binding_ids(&repository, "subject-2").await,
vec!["binding-3"]
);
repository
.update_routing_group_binding(
"binding-3",
UpdateRoutingGroupBindingRecord {
subject_id: Some("subject-1".to_string()),
updated_at: 3,
..UpdateRoutingGroupBindingRecord::default()
},
)
.await
.expect("binding should move");
assert_eq!(
default_binding_ids(&repository, "subject-1").await,
vec!["binding-3"]
);
assert!(default_binding_ids(&repository, "subject-2")
.await
.is_empty());
}
#[tokio::test]
async fn sqlite_repair_migration_resolves_existing_duplicate_defaults() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
let repository = SqliteRoutingGroupRepository::new(pool.clone());
sqlx::raw_sql(
r#"
DROP INDEX routing_groups_one_system_default_key;
DROP INDEX routing_group_bindings_subject_default_key;
"#,
)
.execute(&pool)
.await
.expect("unique indexes should be removable to simulate a pre-repair database");
for id in ["group-1", "group-2", "group-3"] {
repository
.create_routing_group(group_record(id, false))
.await
.expect("group should create");
}
sqlx::query(
r#"
UPDATE routing_groups
SET is_system_default = 1,
enabled = CASE id WHEN 'group-3' THEN 0 ELSE 1 END,
updated_at = CASE id
WHEN 'group-1' THEN 1
WHEN 'group-2' THEN 2
ELSE 3
END
"#,
)
.execute(&pool)
.await
.expect("duplicate system defaults should seed");
for (id, subject_id) in [
("binding-3", "subject-1"),
("binding-2", "subject-1"),
("binding-1", "subject-1"),
("binding-4", "subject-2"),
] {
repository
.create_routing_group_binding(binding_record(id, "group-1", subject_id, false))
.await
.expect("binding should create");
}
sqlx::query(
r#"
UPDATE routing_group_bindings
SET is_default = 1,
created_at = CASE id WHEN 'binding-3' THEN 2 ELSE 1 END
"#,
)
.execute(&pool)
.await
.expect("duplicate binding defaults should seed");
let repair_migration =
include_str!("../migrations/20260727000000_repair_routing_default_uniqueness.sql");
for _ in 0..2 {
sqlx::raw_sql(repair_migration)
.execute(&pool)
.await
.expect("repair migration should be idempotent");
}
assert_eq!(system_default_ids(&repository).await, vec!["group-2"]);
assert_eq!(
default_binding_ids(&repository, "subject-1").await,
vec!["binding-1"]
);
assert_eq!(
default_binding_ids(&repository, "subject-2").await,
vec!["binding-4"]
);
sqlx::query("UPDATE routing_groups SET is_system_default = 1 WHERE id = 'group-3'")
.execute(&pool)
.await
.expect_err("database should reject a second system default");
sqlx::query("UPDATE routing_group_bindings SET is_default = 1 WHERE id = 'binding-2'")
.execute(&pool)
.await
.expect_err("database should reject a second default for the same subject");
}
fn group_record(id: &str, is_system_default: bool) -> CreateRoutingGroupRecord {
CreateRoutingGroupRecord {
id: id.to_string(),
name: id.to_string(),
description: None,
enabled: true,
is_system_default,
config_json: json!({}),
version: 1,
created_at: 1,
updated_at: 1,
published_at: None,
}
}
fn binding_record(
id: &str,
group_id: &str,
subject_id: &str,
is_default: bool,
) -> CreateRoutingGroupBindingRecord {
CreateRoutingGroupBindingRecord {
id: id.to_string(),
group_id: group_id.to_string(),
subject_type: RoutingGroupBindingSubject::ApiKey,
subject_id: subject_id.to_string(),
is_default,
allow_explicit_select: true,
created_at: 1,
updated_at: 1,
}
}
async fn system_default_ids(repository: &SqliteRoutingGroupRepository) -> Vec<String> {
repository
.list_routing_groups()
.await
.expect("groups should list")
.into_iter()
.filter(|group| group.is_system_default)
.map(|group| group.id)
.collect()
}
async fn default_binding_ids(
repository: &SqliteRoutingGroupRepository,
subject_id: &str,
) -> Vec<String> {
repository
.list_routing_group_bindings(&RoutingGroupBindingQuery {
group_id: None,
subject_type: Some(RoutingGroupBindingSubject::ApiKey),
subject_id: Some(subject_id.to_string()),
})
.await
.expect("bindings should list")
.into_iter()
.filter(|binding| binding.is_default)
.map(|binding| binding.id)
.collect()
}
}
@@ -2,7 +2,8 @@ mod types;
pub use types::{
MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
StoredApiFormatCandidateRowsQuery, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
};
@@ -89,6 +89,13 @@ pub struct StoredRequestedModelCandidateRowsQuery {
pub limit: u32,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredApiFormatCandidateRowsQuery {
pub api_format: String,
pub offset: u32,
pub limit: u32,
}
impl StoredMinimalCandidateSelectionRow {
pub fn supports_streaming(&self) -> bool {
self.model_supports_streaming
@@ -119,6 +126,19 @@ pub trait MinimalCandidateSelectionReadRepository: Send + Sync {
api_format: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
async fn list_for_exact_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError> {
Ok(self
.list_for_exact_api_format(&query.api_format)
.await?
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
}
async fn list_for_exact_api_format_and_global_model(
&self,
api_format: &str,
@@ -17,7 +17,10 @@ pub fn merge_score_reason_patch(
}
pub fn score_with_delta(score: f64, delta_basis_points: Option<i32>) -> f64 {
let delta = delta_basis_points.unwrap_or_default() as f64 / 10_000.0;
let Some(delta_basis_points) = delta_basis_points else {
return score;
};
let delta = delta_basis_points as f64 / 10_000.0;
(score + delta).clamp(0.0, 1.0)
}
@@ -46,3 +49,19 @@ pub fn u64_opt_from_i64(
) -> Result<Option<u64>, crate::DataLayerError> {
value.map(|value| u64_from_i64(value, field)).transpose()
}
#[cfg(test)]
mod tests {
use super::score_with_delta;
#[test]
fn score_without_delta_is_unchanged() {
assert_eq!(score_with_delta(1_000.0, None), 1_000.0);
}
#[test]
fn score_with_delta_remains_bounded() {
assert_eq!(score_with_delta(0.95, Some(1_000)), 1.0);
assert_eq!(score_with_delta(0.05, Some(-1_000)), 0.0);
}
}
@@ -3,9 +3,10 @@ use std::sync::RwLock;
use async_trait::async_trait;
use super::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredRequestedModelCandidateRowsQuery,
};
use crate::DataLayerError;
@@ -61,6 +62,19 @@ impl MinimalCandidateSelectionReadRepository for InMemoryMinimalCandidateSelecti
Ok(rows)
}
async fn list_for_exact_api_format_page(
&self,
query: &StoredApiFormatCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
Ok(self
.list_for_exact_api_format(&query.api_format)
.await?
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
}
async fn list_for_exact_api_format_and_global_model(
&self,
api_format: &str,
@@ -351,8 +365,9 @@ fn key_auth_channel_matches(row: &StoredMinimalCandidateSelectionRow, api_format
mod tests {
use super::InMemoryMinimalCandidateSelectionReadRepository;
use crate::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
};
@@ -697,6 +712,27 @@ mod tests {
assert_eq!(rows[1].provider_id, "provider-2");
}
#[tokio::test]
async fn lists_exact_api_format_in_stable_pages() {
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
sample_row("provider-3", "openai:chat", "gpt-4.1", 30),
sample_row("provider-1", "openai:chat", "gpt-4.1", 10),
sample_row("provider-2", "openai:chat", "gpt-4.1", 20),
]);
let page = repository
.list_for_exact_api_format_page(&StoredApiFormatCandidateRowsQuery {
api_format: "openai:chat".to_string(),
offset: 1,
limit: 1,
})
.await
.expect("API-format page should load");
assert_eq!(page.len(), 1);
assert_eq!(page[0].provider_id, "provider-2");
}
#[tokio::test]
async fn list_pool_key_rows_for_group_returns_requested_page_only() {
let mut rows = Vec::new();
@@ -3,9 +3,10 @@ mod memory;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
StoredApiFormatCandidateRowsQuery, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlMinimalCandidateSelectionReadRepository;
@@ -1,6 +1,7 @@
use std::collections::BTreeMap;
use std::sync::RwLock;
use aether_data_contracts::repository::routing_profiles::{apply_binding_patch, apply_group_patch};
use async_trait::async_trait;
use super::{
@@ -149,10 +150,13 @@ impl RoutingGroupWriteRepository for InMemoryRoutingGroupRepository {
record: CreateRoutingGroupRecord,
) -> Result<StoredRoutingGroup, DataLayerError> {
let group = StoredRoutingGroup::new(record)?;
self.groups
.write()
.expect("routing group repository lock")
.insert(group.id.clone(), group.clone());
let mut groups = self.groups.write().expect("routing group repository lock");
if group.is_system_default {
for existing in groups.values_mut() {
existing.is_system_default = false;
}
}
groups.insert(group.id.clone(), group.clone());
Ok(group)
}
@@ -162,42 +166,17 @@ impl RoutingGroupWriteRepository for InMemoryRoutingGroupRepository {
patch: UpdateRoutingGroupRecord,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
let mut groups = self.groups.write().expect("routing group repository lock");
let Some(group) = groups.get_mut(id) else {
let Some(mut group) = groups.get(id).cloned() else {
return Ok(None);
};
if let Some(name) = patch.name {
if name.trim().is_empty() {
return Err(DataLayerError::InvalidInput(
"routing_groups.name is empty".to_string(),
));
apply_group_patch(&mut group, patch)?;
if group.is_system_default {
for existing in groups.values_mut() {
existing.is_system_default = false;
}
group.name = name;
}
if let Some(description) = patch.description {
group.description = description;
}
if let Some(enabled) = patch.enabled {
group.enabled = enabled;
}
if let Some(is_system_default) = patch.is_system_default {
group.is_system_default = is_system_default;
}
if let Some(config_json) = patch.config_json {
if !config_json.is_object() {
return Err(DataLayerError::InvalidInput(
"routing_groups.config_json must be a JSON object".to_string(),
));
}
group.config_json = config_json;
}
if let Some(version) = patch.version {
group.version = version.max(1);
}
if let Some(published_at) = patch.published_at {
group.published_at = published_at;
}
group.updated_at = patch.updated_at;
Ok(Some(group.clone()))
groups.insert(id.to_string(), group.clone());
Ok(Some(group))
}
async fn delete_routing_group(&self, id: &str) -> Result<bool, DataLayerError> {
@@ -214,10 +193,19 @@ impl RoutingGroupWriteRepository for InMemoryRoutingGroupRepository {
record: CreateRoutingGroupBindingRecord,
) -> Result<StoredRoutingGroupBinding, DataLayerError> {
let binding = StoredRoutingGroupBinding::new(record)?;
self.bindings
let mut bindings = self
.bindings
.write()
.expect("routing group binding repository lock")
.insert(binding.id.clone(), binding.clone());
.expect("routing group binding repository lock");
if binding.is_default {
for existing in bindings.values_mut().filter(|existing| {
existing.subject_type == binding.subject_type
&& existing.subject_id == binding.subject_id
}) {
existing.is_default = false;
}
}
bindings.insert(binding.id.clone(), binding.clone());
Ok(binding)
}
@@ -239,36 +227,20 @@ impl RoutingGroupWriteRepository for InMemoryRoutingGroupRepository {
.bindings
.write()
.expect("routing group binding repository lock");
let Some(binding) = bindings.get_mut(id) else {
let Some(mut binding) = bindings.get(id).cloned() else {
return Ok(None);
};
if let Some(group_id) = patch.group_id {
if group_id.trim().is_empty() {
return Err(DataLayerError::InvalidInput(
"routing_group_bindings.group_id is empty".to_string(),
));
apply_binding_patch(&mut binding, patch)?;
if binding.is_default {
for existing in bindings.values_mut().filter(|existing| {
existing.subject_type == binding.subject_type
&& existing.subject_id == binding.subject_id
}) {
existing.is_default = false;
}
binding.group_id = group_id;
}
if let Some(subject_type) = patch.subject_type {
binding.subject_type = subject_type;
}
if let Some(subject_id) = patch.subject_id {
if subject_id.trim().is_empty() {
return Err(DataLayerError::InvalidInput(
"routing_group_bindings.subject_id is empty".to_string(),
));
}
binding.subject_id = subject_id;
}
if let Some(is_default) = patch.is_default {
binding.is_default = is_default;
}
if let Some(allow_explicit_select) = patch.allow_explicit_select {
binding.allow_explicit_select = allow_explicit_select;
}
binding.updated_at = patch.updated_at;
Ok(Some(binding.clone()))
bindings.insert(id.to_string(), binding.clone());
Ok(Some(binding))
}
async fn create_routing_group_version(
@@ -347,4 +319,153 @@ mod tests {
1
);
}
#[tokio::test]
async fn keeps_system_and_subject_defaults_unique() {
let repository = InMemoryRoutingGroupRepository::default();
for (id, is_system_default) in [("group-1", true), ("group-2", true), ("group-3", false)] {
repository
.create_routing_group(group_record(id, is_system_default))
.await
.expect("group should store");
}
assert_eq!(system_default_ids(&repository).await, vec!["group-2"]);
repository
.update_routing_group(
"group-1",
UpdateRoutingGroupRecord {
is_system_default: Some(true),
updated_at: 2,
..UpdateRoutingGroupRecord::default()
},
)
.await
.expect("group should update");
assert_eq!(system_default_ids(&repository).await, vec!["group-1"]);
repository
.create_routing_group_binding(binding_record("binding-1", "group-1", "subject-1", true))
.await
.expect("binding should store");
repository
.create_routing_group_binding(binding_record("binding-2", "group-2", "subject-1", true))
.await
.expect("binding should store");
repository
.create_routing_group_binding(binding_record("binding-3", "group-3", "subject-2", true))
.await
.expect("binding should store");
assert_eq!(
default_binding_ids(&repository, "subject-1").await,
vec!["binding-2"]
);
assert_eq!(
default_binding_ids(&repository, "subject-2").await,
vec!["binding-3"]
);
repository
.update_routing_group_binding(
"binding-1",
UpdateRoutingGroupBindingRecord {
is_default: Some(true),
updated_at: 2,
..UpdateRoutingGroupBindingRecord::default()
},
)
.await
.expect("binding should update");
assert_eq!(
default_binding_ids(&repository, "subject-1").await,
vec!["binding-1"]
);
assert_eq!(
default_binding_ids(&repository, "subject-2").await,
vec!["binding-3"]
);
repository
.update_routing_group_binding(
"binding-3",
UpdateRoutingGroupBindingRecord {
subject_id: Some("subject-1".to_string()),
updated_at: 3,
..UpdateRoutingGroupBindingRecord::default()
},
)
.await
.expect("binding should move");
assert_eq!(
default_binding_ids(&repository, "subject-1").await,
vec!["binding-3"]
);
assert!(default_binding_ids(&repository, "subject-2")
.await
.is_empty());
}
fn group_record(id: &str, is_system_default: bool) -> CreateRoutingGroupRecord {
CreateRoutingGroupRecord {
id: id.to_string(),
name: id.to_string(),
description: None,
enabled: true,
is_system_default,
config_json: json!({}),
version: 1,
created_at: 1,
updated_at: 1,
published_at: None,
}
}
fn binding_record(
id: &str,
group_id: &str,
subject_id: &str,
is_default: bool,
) -> CreateRoutingGroupBindingRecord {
CreateRoutingGroupBindingRecord {
id: id.to_string(),
group_id: group_id.to_string(),
subject_type: RoutingGroupBindingSubject::ApiKey,
subject_id: subject_id.to_string(),
is_default,
allow_explicit_select: true,
created_at: 1,
updated_at: 1,
}
}
async fn system_default_ids(repository: &InMemoryRoutingGroupRepository) -> Vec<String> {
repository
.list_routing_groups()
.await
.expect("groups should list")
.into_iter()
.filter(|group| group.is_system_default)
.map(|group| group.id)
.collect()
}
async fn default_binding_ids(
repository: &InMemoryRoutingGroupRepository,
subject_id: &str,
) -> Vec<String> {
repository
.list_routing_group_bindings(&RoutingGroupBindingQuery {
group_id: None,
subject_type: Some(RoutingGroupBindingSubject::ApiKey),
subject_id: Some(subject_id.to_string()),
})
.await
.expect("bindings should list")
.into_iter()
.filter(|binding| binding.is_default)
.map(|binding| binding.id)
.collect()
}
}
+272 -41
View File
@@ -360,21 +360,12 @@ fn build_pool_sort_vectors<Candidate>(
let lru_ranks = lru_rank_indices(items, false);
let cache_affinity_ranks = lru_rank_indices(items, true);
if lru_enabled {
for item in items {
let key_id = item.item.facts.key_id.clone();
vectors
.entry(key_id.clone())
.or_default()
.push(*lru_ranks.get(&key_id).unwrap_or(&0));
}
}
for preset in presets {
let ranks = match preset.preset.as_str() {
"cache_affinity" => cache_affinity_ranks.clone(),
"priority_first" => priority_first_ranks(items, &lru_ranks),
"single_account" => single_account_ranks(items),
"free_team_first" => plan_ranks(items, &lru_ranks, preset.mode.as_deref()),
"plus_first" => plan_ranks(items, &lru_ranks, Some("plus_only")),
"pro_first" => plan_ranks(items, &lru_ranks, Some("pro_only")),
"free_first" => plan_ranks(items, &lru_ranks, Some("free_only")),
@@ -396,6 +387,16 @@ fn build_pool_sort_vectors<Candidate>(
}
}
if lru_enabled {
for item in items {
let key_id = item.item.facts.key_id.clone();
vectors
.entry(key_id.clone())
.or_default()
.push(*lru_ranks.get(&key_id).unwrap_or(&0));
}
}
vectors
}
@@ -415,13 +416,13 @@ fn lru_rank_indices<Candidate>(
fn priority_first_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
_lru_ranks: &BTreeMap<String, usize>,
lru_ranks: &BTreeMap<String, usize>,
) -> BTreeMap<String, usize> {
let scores = collect_metric_scores(items, |item| {
Some(f64::from(item.item.facts.key_internal_priority))
});
if !score_map_has_variation(&scores) {
return neutral_rank_indices(items);
return lru_ranks.clone();
}
rank_indices_from_score_map(items, &scores, false)
}
@@ -457,7 +458,7 @@ fn single_account_ranks<Candidate>(
fn plan_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
_lru_ranks: &BTreeMap<String, usize>,
lru_ranks: &BTreeMap<String, usize>,
mode: Option<&str>,
) -> BTreeMap<String, usize> {
let scores = items
@@ -473,14 +474,14 @@ fn plan_ranks<Candidate>(
})
.collect::<BTreeMap<_, _>>();
if !score_map_has_variation(&scores) {
return neutral_rank_indices(items);
return lru_ranks.clone();
}
rank_indices_from_score_map(items, &scores, false)
}
fn health_first_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
_lru_ranks: &BTreeMap<String, usize>,
lru_ranks: &BTreeMap<String, usize>,
) -> BTreeMap<String, usize> {
let scores = collect_metric_scores(items, |item| {
item.item
@@ -489,60 +490,63 @@ fn health_first_ranks<Candidate>(
.map(|score| 1.0 - score.clamp(0.0, 1.0))
});
if !score_map_has_signal(&scores) {
return neutral_rank_indices(items);
return lru_ranks.clone();
}
rank_indices_from_score_map(items, &scores, false)
}
fn latency_first_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
_lru_ranks: &BTreeMap<String, usize>,
lru_ranks: &BTreeMap<String, usize>,
) -> BTreeMap<String, usize> {
let scores = collect_metric_scores(items, |item| item.item.key_context.latency_avg_ms);
if !score_map_has_signal(&scores) {
return neutral_rank_indices(items);
return lru_ranks.clone();
}
rank_indices_from_score_map(items, &scores, false)
}
fn cost_first_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
_lru_ranks: &BTreeMap<String, usize>,
lru_ranks: &BTreeMap<String, usize>,
cost_limit_per_key_tokens: Option<u64>,
) -> BTreeMap<String, usize> {
let has_cost_signal = items.iter().any(|item| item.cost_usage > 0);
let scores = collect_metric_scores(items, |item| {
cost_penalty(item, cost_limit_per_key_tokens).or(item.item.key_context.quota_usage_ratio)
cost_penalty(item, cost_limit_per_key_tokens, has_cost_signal)
.or(item.item.key_context.quota_usage_ratio)
});
if !score_map_has_signal(&scores) {
return neutral_rank_indices(items);
return lru_ranks.clone();
}
rank_indices_from_score_map(items, &scores, false)
}
fn quota_balanced_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
_lru_ranks: &BTreeMap<String, usize>,
lru_ranks: &BTreeMap<String, usize>,
cost_limit_per_key_tokens: Option<u64>,
) -> BTreeMap<String, usize> {
let has_cost_signal = items.iter().any(|item| item.cost_usage > 0);
let scores = collect_metric_scores(items, |item| {
item.item
.key_context
.quota_usage_ratio
.or_else(|| cost_penalty(item, cost_limit_per_key_tokens))
.or_else(|| cost_penalty(item, cost_limit_per_key_tokens, has_cost_signal))
});
if !score_map_has_signal(&scores) {
return neutral_rank_indices(items);
return lru_ranks.clone();
}
rank_indices_from_score_map(items, &scores, false)
}
fn recent_refresh_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
_lru_ranks: &BTreeMap<String, usize>,
lru_ranks: &BTreeMap<String, usize>,
) -> BTreeMap<String, usize> {
let scores = collect_metric_scores(items, |item| item.item.key_context.quota_reset_seconds);
if !score_map_has_signal(&scores) {
return neutral_rank_indices(items);
return lru_ranks.clone();
}
rank_indices_from_score_map(items, &scores, false)
}
@@ -653,29 +657,33 @@ fn rank_indices_from_score_map<Candidate>(
.then(left.2.cmp(&right.2))
});
decorated
.into_iter()
.enumerate()
.map(|(rank, (_, _, _, key_id))| (key_id, rank))
.collect()
}
fn neutral_rank_indices<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
) -> BTreeMap<String, usize> {
items
.iter()
.map(|item| (item.item.facts.key_id.clone(), 0))
.collect()
let mut ranks = BTreeMap::new();
let mut previous_score = None::<(bool, f64)>;
let mut rank = 0;
for (index, (missing, score, _, key_id)) in decorated.into_iter().enumerate() {
if previous_score
.as_ref()
.is_some_and(|previous| previous.0 != missing || previous.1 != score)
{
rank = index;
}
previous_score = Some((missing, score));
ranks.insert(key_id, rank);
}
ranks
}
fn cost_penalty<Candidate>(
item: &PoolGroupCandidateOrdering<Candidate>,
cost_limit_per_key_tokens: Option<u64>,
has_cost_signal: bool,
) -> Option<f64> {
if item.cost_usage == 0 {
if !has_cost_signal {
return None;
}
if item.cost_usage == 0 {
return Some(0.0);
}
if let Some(limit) = cost_limit_per_key_tokens.filter(|limit| *limit > 0) {
return Some((item.cost_usage as f64 / limit as f64).clamp(0.0, 1.0));
@@ -1006,6 +1014,223 @@ mod tests {
);
}
#[test]
fn cost_first_treats_zero_usage_as_the_lowest_cost() {
let unused = sample_candidate("provider-pool", "endpoint-1", "key-unused", 10, true)
.with_presets(vec![PoolSchedulingPreset {
preset: "cost_first".to_string(),
enabled: true,
mode: None,
}]);
let used = sample_candidate("provider-pool", "endpoint-1", "key-used", 10, true)
.with_presets(vec![PoolSchedulingPreset {
preset: "cost_first".to_string(),
enabled: true,
mode: None,
}]);
let runtime_by_provider = BTreeMap::from([(
"provider-pool".to_string(),
PoolRuntimeState {
cost_window_usage_by_key: BTreeMap::from([("key-used".to_string(), 100)]),
lru_score_by_key: BTreeMap::from([
("key-used".to_string(), 0.0),
("key-unused".to_string(), 100.0),
]),
..PoolRuntimeState::default()
},
)]);
let outcome = run_pool_scheduler(vec![used, unused], &runtime_by_provider, "seed");
assert!(outcome.skipped_candidates.is_empty());
assert_eq!(
outcome
.candidates
.iter()
.map(|item| item.candidate.as_str())
.collect::<Vec<_>>(),
vec!["key-unused", "key-used"]
);
}
#[test]
fn cost_first_falls_back_to_quota_without_cost_signal() {
let high_quota_usage =
sample_candidate("provider-pool", "endpoint-1", "key-high-quota", 10, true)
.with_presets(vec![PoolSchedulingPreset {
preset: "cost_first".to_string(),
enabled: true,
mode: None,
}])
.with_quota_usage(0.8);
let low_quota_usage =
sample_candidate("provider-pool", "endpoint-1", "key-low-quota", 10, true)
.with_presets(vec![PoolSchedulingPreset {
preset: "cost_first".to_string(),
enabled: true,
mode: None,
}])
.with_quota_usage(0.2);
let runtime_by_provider = BTreeMap::from([(
"provider-pool".to_string(),
PoolRuntimeState {
lru_score_by_key: BTreeMap::from([
("key-high-quota".to_string(), 0.0),
("key-low-quota".to_string(), 100.0),
]),
..PoolRuntimeState::default()
},
)]);
let outcome = run_pool_scheduler(
vec![high_quota_usage, low_quota_usage],
&runtime_by_provider,
"seed",
);
assert_eq!(
outcome
.candidates
.iter()
.map(|item| item.candidate.as_str())
.collect::<Vec<_>>(),
vec!["key-low-quota", "key-high-quota"]
);
}
#[test]
fn cost_first_falls_back_to_lru_without_cost_or_quota_signal() {
let key_recent = sample_candidate("provider-pool", "endpoint-1", "key-recent", 10, true)
.with_presets(vec![PoolSchedulingPreset {
preset: "cost_first".to_string(),
enabled: true,
mode: None,
}]);
let key_old = sample_candidate("provider-pool", "endpoint-1", "key-old", 10, true)
.with_presets(vec![PoolSchedulingPreset {
preset: "cost_first".to_string(),
enabled: true,
mode: None,
}]);
let runtime_by_provider = BTreeMap::from([(
"provider-pool".to_string(),
PoolRuntimeState {
lru_score_by_key: BTreeMap::from([
("key-recent".to_string(), 100.0),
("key-old".to_string(), 0.0),
]),
..PoolRuntimeState::default()
},
)]);
let outcome = run_pool_scheduler(vec![key_recent, key_old], &runtime_by_provider, "seed");
assert_eq!(
outcome
.candidates
.iter()
.map(|item| item.candidate.as_str())
.collect::<Vec<_>>(),
vec!["key-old", "key-recent"]
);
}
#[test]
fn legacy_free_team_first_preserves_all_modes() {
let scheduled = |mode: &str, key_plans: &[(&str, &str)]| {
let candidates = key_plans
.iter()
.map(|(key_id, plan)| {
sample_candidate("provider-pool", "endpoint-1", key_id, 10, true)
.with_presets(vec![PoolSchedulingPreset {
preset: "free_team_first".to_string(),
enabled: true,
mode: Some(mode.to_string()),
}])
.with_plan(plan)
})
.collect::<Vec<_>>();
run_pool_scheduler(candidates, &BTreeMap::new(), "seed")
.candidates
.into_iter()
.map(|item| item.candidate)
.collect::<Vec<_>>()
};
assert_eq!(
scheduled("free_only", &[("key-team", "team"), ("key-free", "free")]),
["key-free", "key-team"]
);
assert_eq!(
scheduled("team_only", &[("key-free", "free"), ("key-team", "team")]),
["key-team", "key-free"]
);
assert_eq!(
scheduled(
"both",
&[
("key-plus", "plus"),
("key-team", "team"),
("key-free", "free")
],
),
["key-team", "key-free", "key-plus"]
);
}
#[test]
fn strategy_ties_use_lru_as_the_final_tiebreaker() {
let recent_free =
sample_candidate("provider-pool", "endpoint-1", "key-recent-free", 10, true)
.with_presets(vec![PoolSchedulingPreset {
preset: "free_first".to_string(),
enabled: true,
mode: None,
}])
.with_plan("free");
let plus = sample_candidate("provider-pool", "endpoint-1", "key-plus", 10, true)
.with_presets(vec![PoolSchedulingPreset {
preset: "free_first".to_string(),
enabled: true,
mode: None,
}])
.with_plan("plus");
let old_free = sample_candidate("provider-pool", "endpoint-1", "key-old-free", 10, true)
.with_presets(vec![PoolSchedulingPreset {
preset: "free_first".to_string(),
enabled: true,
mode: None,
}])
.with_plan("free");
let runtime_by_provider = BTreeMap::from([(
"provider-pool".to_string(),
PoolRuntimeState {
lru_score_by_key: BTreeMap::from([
("key-recent-free".to_string(), 200.0),
("key-old-free".to_string(), 10.0),
("key-plus".to_string(), 0.0),
]),
..PoolRuntimeState::default()
},
)]);
let outcome = run_pool_scheduler(
vec![recent_free, plus, old_free],
&runtime_by_provider,
"seed",
);
assert!(outcome.skipped_candidates.is_empty());
assert_eq!(
outcome
.candidates
.iter()
.map(|item| item.candidate.as_str())
.collect::<Vec<_>>(),
["key-old-free", "key-recent-free", "key-plus"]
);
}
#[test]
fn pool_scheduler_applies_distribution_mode_before_strategy_presets() {
let key_cache_hit =
@@ -1276,6 +1501,7 @@ mod tests {
fn with_cost_limit(self, limit: u64) -> Self;
fn with_presets(self, presets: Vec<PoolSchedulingPreset>) -> Self;
fn with_plan(self, plan: &str) -> Self;
fn with_quota_usage(self, ratio: f64) -> Self;
}
impl TestCandidateExt for PoolCandidateInput<String> {
@@ -1297,5 +1523,10 @@ mod tests {
self.key_context.plan_tier = Some(plan.to_string());
self
}
fn with_quota_usage(mut self, ratio: f64) -> Self {
self.key_context.quota_usage_ratio = Some(ratio);
self
}
}
}
+23
View File
@@ -604,6 +604,10 @@ mod tests {
.iter()
.find(|item| item["name"] == "recent_refresh")
.expect("recent_refresh should exist");
let legacy_free_team = items
.iter()
.find(|item| item["name"] == "free_team_first")
.expect("legacy free_team_first should remain configurable");
assert_eq!(
free_first["providers"],
@@ -613,6 +617,14 @@ mod tests {
recent_refresh["providers"],
json!(["codex", "grok", "kiro", "windsurf"])
);
assert_eq!(free_first["default_enabled"], json!(false));
assert_eq!(recent_refresh["default_enabled"], json!(false));
assert_eq!(
recent_refresh["default_enabled_providers"],
json!(["codex", "windsurf"])
);
assert_eq!(legacy_free_team["default_mode"], json!("both"));
assert_eq!(legacy_free_team["modes"].as_array().map(Vec::len), Some(3));
}
#[test]
@@ -657,6 +669,17 @@ mod tests {
["cache_affinity", "recent_refresh"]
);
let legacy_free_team = service.normalize_scheduling_presets(
"codex",
&[PoolSchedulingPreset {
preset: "free_team_first".to_string(),
enabled: true,
mode: Some("team_only".to_string()),
}],
);
assert_eq!(legacy_free_team[0].preset, "free_team_first");
assert_eq!(legacy_free_team[0].mode.as_deref(), Some("team_only"));
let unsupported = service.normalize_scheduling_presets(
"chatgpt_web",
&[PoolSchedulingPreset {
+33 -1
View File
@@ -110,6 +110,7 @@ pub fn build_admin_pool_scheduling_presets_payload() -> Value {
"依据窗口成本/Token 用量,缺失时回退配额使用率",
&service,
),
legacy_free_team_first_preset_payload(&service),
provider_pool_preset_payload(
"free_first",
"Free 优先",
@@ -201,6 +202,24 @@ pub fn build_admin_pool_scheduling_presets_payload() -> Value {
])
}
fn legacy_free_team_first_preset_payload(service: &ProviderPoolService) -> Value {
let mut payload = provider_pool_preset_payload(
"free_team_first",
"Free/Team 优先",
"兼容旧配置:优先消耗 Free、Team 或两者",
Some(ProviderPoolCapability::PlanTier),
"依据 plan_type,保留旧 free_only/team_only/both 语义",
service,
);
payload["modes"] = json!([
{"value": "free_only", "label": "Free"},
{"value": "team_only", "label": "Team"},
{"value": "both", "label": "Free + Team"}
]);
payload["default_mode"] = json!("both");
payload
}
fn provider_pool_preset_payload(
name: &'static str,
label: &'static str,
@@ -212,11 +231,24 @@ fn provider_pool_preset_payload(
let providers = capability
.map(|capability| service.provider_types_for_capability(capability))
.unwrap_or_default();
let default_enabled_providers = service
.provider_types()
.filter(|provider_type| {
service
.adapter(provider_type)
.default_scheduling_presets()
.iter()
.any(|preset| preset.enabled && preset.preset.eq_ignore_ascii_case(name))
})
.map(str::to_string)
.collect::<Vec<_>>();
json!({
"name": name,
"label": label,
"description": description,
"providers": providers,
"default_enabled": name == "cache_affinity",
"default_enabled_providers": default_enabled_providers,
"modes": Value::Null,
"default_mode": Value::Null,
"mutex_group": provider_pool_preset_mutex_group(name),
@@ -226,7 +258,7 @@ fn provider_pool_preset_payload(
fn provider_pool_supports_preset(adapter: &dyn ProviderPoolAdapter, preset: &str) -> bool {
match preset {
"free_first" | "plus_first" | "pro_first" | "team_first" => adapter
"free_first" | "free_team_first" | "plus_first" | "pro_first" | "team_first" => adapter
.capabilities()
.supports(ProviderPoolCapability::PlanTier),
"recent_refresh" => adapter
+3 -1
View File
@@ -33,4 +33,6 @@ pub use trace::{
RoutingCandidateTrace, RoutingDecisionTrace, RoutingPatchSummary, RoutingPoolExpansionTrace,
RoutingRuntimeFacts,
};
pub use validation::{validate_routing_group_config, RoutingValidationError};
pub use validation::{
validate_routing_group_config, RoutingValidationError, MAX_ROUTING_ALLOWED_KEYS,
};
+4
View File
@@ -11,6 +11,7 @@ use crate::conditions::RoutingConditionContext;
use crate::model::{RoutingGroupConfig, RoutingModelPolicy, RoutingPoolPolicyOverride};
use crate::mutations::{validate_header_patch, validate_json_patch_operations, MutationPlan};
use crate::ranking::RankingOverlay;
use crate::validation::validate_routing_group_config;
#[derive(Debug, Error, Clone, PartialEq, Eq)]
pub enum RoutingPolicyError {
@@ -68,6 +69,9 @@ pub fn resolve_routing_policy(
config: &RoutingGroupConfig,
input: RoutingPolicyInput<'_>,
) -> Result<ResolvedRoutingPolicy, RoutingPolicyError> {
validate_routing_group_config(config)
.map_err(|error| RoutingPolicyError::InvalidConfig(error.to_string()))?;
if !model_allowed(&config.allowed_models, input.requested_model)
&& !model_allowed(&config.allowed_models, input.resolved_model)
{
+17 -7
View File
@@ -40,6 +40,13 @@ impl RankingOverlay {
.unwrap_or(fallback)
}
pub fn pool_priority(&self, provider_id: &str, fallback: i32) -> i32 {
self.pool_priority_overrides
.get(provider_id)
.copied()
.unwrap_or(fallback)
}
pub fn provider_priority_or_unspecified(&self, provider_id: &str) -> i32 {
self.provider_priority_overrides
.get(provider_id)
@@ -100,15 +107,18 @@ pub fn rank_vector_for_candidate(
) -> RoutingCandidateRankVector {
RoutingCandidateRankVector {
provider_priority_before: facts.provider_priority,
provider_priority_after: overlay.provider_priority_or_unspecified(&facts.provider_id),
provider_priority_after: overlay
.provider_priority(&facts.provider_id, facts.provider_priority),
key_priority_before: facts.key_priority,
key_priority_after: match facts.candidate_kind {
CandidateKind::Provider => facts
.key_id
.as_deref()
.map(|key_id| overlay.key_priority_or_unspecified(key_id))
.unwrap_or(ROUTING_PRIORITY_UNSPECIFIED),
CandidateKind::PoolGroup => overlay.pool_priority_or_unspecified(&facts.provider_id),
.map(|key_id| overlay.key_priority(key_id, facts.key_priority))
.unwrap_or(facts.key_priority),
CandidateKind::PoolGroup => {
overlay.pool_priority(&facts.provider_id, facts.key_priority)
}
},
}
}
@@ -142,7 +152,7 @@ mod tests {
}
#[test]
fn rank_vector_marks_missing_routing_priorities_unspecified() {
fn rank_vector_falls_back_to_existing_priorities() {
let facts = RoutingCandidateFacts {
candidate_kind: CandidateKind::Provider,
provider_id: "provider-a".to_string(),
@@ -154,8 +164,8 @@ mod tests {
};
let vector = rank_vector_for_candidate(&RankingOverlay::default(), &facts);
assert_eq!(vector.provider_priority_after, ROUTING_PRIORITY_UNSPECIFIED);
assert_eq!(vector.key_priority_after, ROUTING_PRIORITY_UNSPECIFIED);
assert_eq!(vector.provider_priority_after, 10);
assert_eq!(vector.key_priority_after, 20);
}
#[test]
+440 -2
View File
@@ -2,9 +2,29 @@ use std::collections::BTreeSet;
use thiserror::Error;
use crate::model::RoutingGroupConfig;
use crate::model::{RoutingGroupConfig, RoutingPoolPolicyOverride};
use crate::mutations::{validate_header_patch, validate_json_patch_operations};
use crate::RoutingAction;
use crate::{RoutingAction, RoutingRulePhase};
pub const MAX_ROUTING_ALLOWED_KEYS: usize = 512;
const ROUTING_POOL_PRESETS: &[&str] = &[
"lru",
"cache_affinity",
"load_balance",
"single_account",
"priority_first",
"free_team_first",
"free_first",
"team_first",
"plus_first",
"pro_first",
"health_first",
"latency_first",
"cost_first",
"quota_balanced",
"recent_refresh",
];
#[derive(Debug, Error, Clone, PartialEq, Eq)]
pub enum RoutingValidationError {
@@ -16,6 +36,34 @@ pub enum RoutingValidationError {
EmptyModelSelector,
#[error("invalid mutation action: {0}")]
InvalidMutation(String),
#[error("routing rule {rule_id} uses unsupported {action} action in provider_request phase")]
ProviderRequestActionNotAllowed {
rule_id: String,
action: &'static str,
},
#[error("routing key selector {selector} contains {count} entries; maximum is {max}")]
TooManyAllowedKeys {
selector: String,
count: usize,
max: usize,
},
#[error("routing pool policy {selector} has an empty provider id")]
EmptyPoolProviderId { selector: String },
#[error("routing pool policy {selector} uses unsupported preset: {preset}")]
UnsupportedPoolPreset { selector: String, preset: String },
#[error("routing pool policy {selector} contains duplicate preset: {preset}")]
DuplicatePoolPreset { selector: String, preset: String },
#[error("routing pool policy {selector} preset {preset} has invalid mode: {mode}")]
InvalidPoolPresetMode {
selector: String,
preset: String,
mode: String,
},
#[error("routing pool policy {selector} enables mutually exclusive distribution presets: {presets:?}")]
ConflictingPoolDistributionPresets {
selector: String,
presets: Vec<String>,
},
}
pub fn validate_routing_group_config(
@@ -26,6 +74,22 @@ pub fn validate_routing_group_config(
if model_policy.model.trim().is_empty() {
return Err(RoutingValidationError::EmptyModelSelector);
}
validate_allowed_key_count(
format!("model:{}", model_policy.model.trim()),
model_policy.allowed_keys.len(),
)?;
for (provider_id, override_policy) in &model_policy.pool_policy_overrides {
let model_selector = format!("model:{}", model_policy.model.trim());
if provider_id.trim().is_empty() {
return Err(RoutingValidationError::EmptyPoolProviderId {
selector: model_selector,
});
}
validate_pool_policy_override(
format!("{model_selector}:provider:{}", provider_id.trim()),
override_policy,
)?;
}
}
for rule in &config.rules {
if rule.id.trim().is_empty() {
@@ -35,6 +99,17 @@ pub fn validate_routing_group_config(
return Err(RoutingValidationError::DuplicateRuleId(rule.id.clone()));
}
for action in &rule.actions {
if rule.phase == RoutingRulePhase::ProviderRequest
&& !matches!(
action,
RoutingAction::JsonPatchBody { .. } | RoutingAction::PatchHeaders { .. }
)
{
return Err(RoutingValidationError::ProviderRequestActionNotAllowed {
rule_id: rule.id.clone(),
action: routing_action_name(action),
});
}
match action {
RoutingAction::JsonPatchBody { patch } => {
validate_json_patch_operations(patch).map_err(|error| {
@@ -46,9 +121,372 @@ pub fn validate_routing_group_config(
RoutingValidationError::InvalidMutation(error.to_string())
})?;
}
RoutingAction::RestrictKeys { key_ids } => {
validate_allowed_key_count(format!("rule:{}", rule.id), key_ids.len())?;
}
_ => {}
}
}
}
Ok(())
}
fn validate_pool_policy_override(
selector: String,
override_policy: &RoutingPoolPolicyOverride,
) -> Result<(), RoutingValidationError> {
let mut seen = BTreeSet::new();
let mut enabled_distribution_presets = Vec::new();
for preset in &override_policy.scheduling_presets {
let normalized = preset.preset.trim().to_ascii_lowercase();
if !ROUTING_POOL_PRESETS.contains(&normalized.as_str()) {
return Err(RoutingValidationError::UnsupportedPoolPreset {
selector,
preset: normalized,
});
}
if !seen.insert(normalized.clone()) {
return Err(RoutingValidationError::DuplicatePoolPreset {
selector,
preset: normalized,
});
}
if let Some(mode) = preset.mode.as_deref() {
let mode = mode.trim().to_ascii_lowercase();
if !routing_pool_preset_mode_valid(&normalized, &mode) {
return Err(RoutingValidationError::InvalidPoolPresetMode {
selector,
preset: normalized,
mode,
});
}
}
if preset.enabled && routing_pool_distribution_preset(&normalized) {
enabled_distribution_presets.push(normalized);
}
}
if enabled_distribution_presets.len() > 1 {
return Err(RoutingValidationError::ConflictingPoolDistributionPresets {
selector,
presets: enabled_distribution_presets,
});
}
Ok(())
}
fn routing_pool_preset_mode_valid(preset: &str, mode: &str) -> bool {
match preset {
"free_team_first" => matches!(mode, "free_only" | "team_only" | "both"),
"free_first" => mode == "free_only",
"team_first" => mode == "team_only",
"plus_first" => mode == "plus_only",
"pro_first" => mode == "pro_only",
_ => false,
}
}
fn routing_pool_distribution_preset(preset: &str) -> bool {
matches!(
preset,
"lru" | "cache_affinity" | "load_balance" | "single_account"
)
}
fn validate_allowed_key_count(
selector: String,
count: usize,
) -> Result<(), RoutingValidationError> {
if count > MAX_ROUTING_ALLOWED_KEYS {
return Err(RoutingValidationError::TooManyAllowedKeys {
selector,
count,
max: MAX_ROUTING_ALLOWED_KEYS,
});
}
Ok(())
}
fn routing_action_name(action: &RoutingAction) -> &'static str {
match action {
RoutingAction::RestrictModels { .. } => "restrict_models",
RoutingAction::RestrictProviders { .. } => "restrict_providers",
RoutingAction::RestrictKeys { .. } => "restrict_keys",
RoutingAction::SetScheduling { .. } => "set_scheduling",
RoutingAction::SetProviderPriority { .. } => "set_provider_priority",
RoutingAction::SetKeyPriority { .. } => "set_key_priority",
RoutingAction::JsonPatchBody { .. } => "json_patch_body",
RoutingAction::PatchHeaders { .. } => "patch_headers",
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use crate::{
RoutingCondition, RoutingGroupConfig, RoutingHeaderPatch, RoutingJsonPatchOperation,
RoutingModelPolicy, RoutingRule,
};
use super::*;
fn provider_request_config(actions: Vec<RoutingAction>) -> RoutingGroupConfig {
RoutingGroupConfig {
rules: vec![RoutingRule {
id: "provider-rule".to_string(),
priority: 0,
enabled: true,
phase: RoutingRulePhase::ProviderRequest,
conditions: RoutingCondition::default(),
actions,
stop_processing: false,
}],
..RoutingGroupConfig::default()
}
}
#[test]
fn rejects_non_mutation_actions_in_provider_request_phase() {
let actions = [
(
RoutingAction::RestrictModels {
models: vec!["gpt-5".to_string()],
},
"restrict_models",
),
(
RoutingAction::RestrictProviders {
provider_ids: vec!["provider-1".to_string()],
},
"restrict_providers",
),
(
RoutingAction::RestrictKeys {
key_ids: vec!["key-1".to_string()],
},
"restrict_keys",
),
(
RoutingAction::SetScheduling {
priority_mode: None,
scheduling_mode: None,
keep_priority_on_conversion: Some(true),
},
"set_scheduling",
),
(
RoutingAction::SetProviderPriority {
provider_id: "provider-1".to_string(),
priority: 1,
},
"set_provider_priority",
),
(
RoutingAction::SetKeyPriority {
key_id: "key-1".to_string(),
priority: 1,
},
"set_key_priority",
),
];
for (action, action_name) in actions {
let error = validate_routing_group_config(&provider_request_config(vec![action]))
.expect_err("provider_request must only accept mutations");
assert_eq!(
error,
RoutingValidationError::ProviderRequestActionNotAllowed {
rule_id: "provider-rule".to_string(),
action: action_name,
}
);
}
}
#[test]
fn accepts_mutation_actions_in_provider_request_phase() {
let config = provider_request_config(vec![
RoutingAction::JsonPatchBody {
patch: vec![RoutingJsonPatchOperation::Add {
path: "/metadata/routed".to_string(),
value: json!(true),
}],
},
RoutingAction::PatchHeaders {
patch: vec![RoutingHeaderPatch::Set {
name: "x-routing-profile".to_string(),
value: "provider-rule".to_string(),
}],
},
]);
validate_routing_group_config(&config)
.expect("provider_request mutation actions should remain valid");
}
#[test]
fn rejects_oversized_model_allowed_key_selector() {
let config = RoutingGroupConfig {
model_policies: vec![RoutingModelPolicy {
model: "gpt-5".to_string(),
allowed_keys: key_ids(MAX_ROUTING_ALLOWED_KEYS + 1),
..RoutingModelPolicy::default()
}],
..RoutingGroupConfig::default()
};
assert_eq!(
validate_routing_group_config(&config),
Err(RoutingValidationError::TooManyAllowedKeys {
selector: "model:gpt-5".to_string(),
count: MAX_ROUTING_ALLOWED_KEYS + 1,
max: MAX_ROUTING_ALLOWED_KEYS,
})
);
}
#[test]
fn rejects_oversized_rule_allowed_key_selector() {
let config = RoutingGroupConfig {
rules: vec![RoutingRule {
id: "restrict-keys".to_string(),
priority: 0,
enabled: true,
phase: RoutingRulePhase::ClientRequest,
conditions: RoutingCondition::default(),
actions: vec![RoutingAction::RestrictKeys {
key_ids: key_ids(MAX_ROUTING_ALLOWED_KEYS + 1),
}],
stop_processing: false,
}],
..RoutingGroupConfig::default()
};
assert_eq!(
validate_routing_group_config(&config),
Err(RoutingValidationError::TooManyAllowedKeys {
selector: "rule:restrict-keys".to_string(),
count: MAX_ROUTING_ALLOWED_KEYS + 1,
max: MAX_ROUTING_ALLOWED_KEYS,
})
);
}
#[test]
fn accepts_allowed_key_selectors_at_the_limit() {
let config = RoutingGroupConfig {
model_policies: vec![RoutingModelPolicy {
model: "gpt-5".to_string(),
allowed_keys: key_ids(MAX_ROUTING_ALLOWED_KEYS),
..RoutingModelPolicy::default()
}],
..RoutingGroupConfig::default()
};
validate_routing_group_config(&config)
.expect("allowed key selectors at the scan limit should remain valid");
}
#[test]
fn rejects_unknown_pool_override_preset() {
let config = pool_override_config(vec![pool_preset("typo_priority", true, None)]);
assert_eq!(
validate_routing_group_config(&config),
Err(RoutingValidationError::UnsupportedPoolPreset {
selector: "model:gpt-5:provider:provider-1".to_string(),
preset: "typo_priority".to_string(),
})
);
}
#[test]
fn rejects_duplicate_pool_override_presets() {
let config = pool_override_config(vec![
pool_preset("health_first", true, None),
pool_preset(" HEALTH_FIRST ", false, None),
]);
assert_eq!(
validate_routing_group_config(&config),
Err(RoutingValidationError::DuplicatePoolPreset {
selector: "model:gpt-5:provider:provider-1".to_string(),
preset: "health_first".to_string(),
})
);
}
#[test]
fn rejects_invalid_pool_override_mode() {
let config =
pool_override_config(vec![pool_preset("free_team_first", true, Some("pro_only"))]);
assert_eq!(
validate_routing_group_config(&config),
Err(RoutingValidationError::InvalidPoolPresetMode {
selector: "model:gpt-5:provider:provider-1".to_string(),
preset: "free_team_first".to_string(),
mode: "pro_only".to_string(),
})
);
}
#[test]
fn rejects_conflicting_pool_override_distribution_modes() {
let config = pool_override_config(vec![
pool_preset("lru", true, None),
pool_preset("cache_affinity", true, None),
]);
assert_eq!(
validate_routing_group_config(&config),
Err(RoutingValidationError::ConflictingPoolDistributionPresets {
selector: "model:gpt-5:provider:provider-1".to_string(),
presets: vec!["lru".to_string(), "cache_affinity".to_string()],
})
);
}
#[test]
fn accepts_valid_pool_override_modes_and_disabled_alternatives() {
let config = pool_override_config(vec![
pool_preset("cache_affinity", true, None),
pool_preset("lru", false, None),
pool_preset("free_team_first", true, Some("team_only")),
]);
validate_routing_group_config(&config).expect("valid pool override should pass");
}
fn pool_override_config(
scheduling_presets: Vec<crate::RoutingSchedulingPreset>,
) -> RoutingGroupConfig {
RoutingGroupConfig {
model_policies: vec![RoutingModelPolicy {
model: "gpt-5".to_string(),
pool_policy_overrides: std::collections::BTreeMap::from([(
"provider-1".to_string(),
RoutingPoolPolicyOverride { scheduling_presets },
)]),
..RoutingModelPolicy::default()
}],
..RoutingGroupConfig::default()
}
}
fn pool_preset(
preset: &str,
enabled: bool,
mode: Option<&str>,
) -> crate::RoutingSchedulingPreset {
crate::RoutingSchedulingPreset {
preset: preset.to_string(),
enabled,
mode: mode.map(str::to_string),
}
}
fn key_ids(count: usize) -> Vec<String> {
(0..count).map(|index| format!("key-{index}")).collect()
}
}
+122 -16
View File
@@ -1,3 +1,4 @@
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use crate::SchedulerMinimalCandidateSelectionCandidate;
@@ -9,6 +10,26 @@ pub struct SchedulerAffinityTarget {
pub key_id: String,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SchedulerAffinityScope {
pub routing_group_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub routing_group_version: Option<i64>,
}
impl SchedulerAffinityScope {
pub fn new(routing_group_id: impl Into<String>, routing_group_version: Option<i64>) -> Self {
Self {
routing_group_id: routing_group_id.into(),
routing_group_version,
}
}
pub fn is_valid(&self) -> bool {
!self.routing_group_id.trim().is_empty()
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct ClientSessionAffinity {
pub client_family: Option<String>,
@@ -56,6 +77,22 @@ pub fn build_scheduler_affinity_cache_key_for_api_key_id_with_client_session(
api_format: &str,
global_model_name: &str,
client_session_affinity: Option<&ClientSessionAffinity>,
) -> Option<String> {
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,
None,
)
}
pub fn build_scheduler_affinity_cache_key_for_api_key_id_with_client_session_and_scope(
api_key_id: &str,
api_format: &str,
global_model_name: &str,
client_session_affinity: Option<&ClientSessionAffinity>,
affinity_scope: Option<&SchedulerAffinityScope>,
) -> Option<String> {
let api_key_id = api_key_id.trim();
if api_key_id.is_empty() {
@@ -67,36 +104,63 @@ pub fn build_scheduler_affinity_cache_key_for_api_key_id_with_client_session(
return None;
}
let legacy_key = format!("scheduler_affinity:{api_key_id}:{api_format}:{global_model_name}");
let Some(client_session_affinity) = client_session_affinity else {
return Some(legacy_key);
};
let Some(session_key) = client_session_affinity
.session_key
.as_deref()
let affinity_scope = affinity_scope.filter(|scope| scope.is_valid());
let session_key = client_session_affinity
.and_then(|affinity| affinity.session_key.as_deref())
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Some(legacy_key);
};
.filter(|value| !value.is_empty());
if session_key.is_none() && affinity_scope.is_none() {
return Some(format!(
"scheduler_affinity:{api_key_id}:{api_format}:{global_model_name}"
));
}
let client_family = client_session_affinity
.client_family
.as_deref()
.and_then(|affinity| affinity.client_family.as_deref())
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_ascii_lowercase)
.unwrap_or_else(|| "generic".to_string());
let session_hash = hash_session_key(session_key);
let session_hash = match affinity_scope {
Some(scope) => hash_scoped_session_key(session_key, scope),
None => hash_session_key(session_key.expect("session key should exist without scope")),
};
Some(format!(
"scheduler_affinity:v2:{api_key_id}:{api_format}:{global_model_name}:{client_family}:{session_hash}"
))
}
fn hash_scoped_session_key(
session_key: Option<&str>,
affinity_scope: &SchedulerAffinityScope,
) -> String {
let mut hasher = Sha256::new();
hasher.update(b"scheduler-affinity-scope-v1\0");
hasher.update(affinity_scope.routing_group_id.trim().as_bytes());
hasher.update(b"\0");
match affinity_scope.routing_group_version {
Some(version) => hasher.update(version.to_string().as_bytes()),
None => hasher.update(b"unversioned"),
}
hasher.update(b"\0");
match session_key {
Some(session_key) => {
hasher.update(b"session\0");
hasher.update(session_key.as_bytes());
}
None => hasher.update(b"api-key-scope"),
}
hex_digest(hasher.finalize())
}
fn hash_session_key(session_key: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(session_key.as_bytes());
let digest = hasher.finalize();
hex_digest(hasher.finalize())
}
fn hex_digest(digest: impl AsRef<[u8]>) -> String {
let digest = digest.as_ref();
digest.iter().map(|byte| format!("{byte:02x}")).collect()
}
@@ -142,8 +206,9 @@ mod tests {
use super::{
build_scheduler_affinity_cache_key_for_api_key_id,
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,
candidate_affinity_hash, candidate_key, matches_affinity_target, ClientSessionAffinity,
SchedulerAffinityTarget,
SchedulerAffinityScope, SchedulerAffinityTarget,
};
use crate::SchedulerMinimalCandidateSelectionCandidate;
@@ -262,6 +327,47 @@ mod tests {
assert_ne!(left_key, other_client_key);
}
#[test]
fn routing_scoped_affinity_keys_split_groups_and_versions() {
let affinity =
ClientSessionAffinity::new(Some("generic".to_string()), Some("session-a".to_string()));
let group_one_v1 = SchedulerAffinityScope::new("group-1", Some(1));
let group_one_v2 = SchedulerAffinityScope::new("group-1", Some(2));
let group_two_v1 = SchedulerAffinityScope::new("group-2", Some(1));
let key_for = |scope: &SchedulerAffinityScope| {
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 affinity key should build")
};
assert_ne!(key_for(&group_one_v1), key_for(&group_one_v2));
assert_ne!(key_for(&group_one_v1), key_for(&group_two_v1));
}
#[test]
fn routing_scoped_affinity_key_isolated_without_client_session() {
let group_one = SchedulerAffinityScope::new("group-1", Some(1));
let group_two = SchedulerAffinityScope::new("group-2", Some(1));
let key_for = |scope: &SchedulerAffinityScope| {
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 affinity key should build")
};
assert!(key_for(&group_one).starts_with("scheduler_affinity:v2:"));
assert_ne!(key_for(&group_one), key_for(&group_two));
}
#[test]
fn affinity_hash_is_candidate_specific() {
let left = sample_candidate("1");
+4 -2
View File
@@ -9,8 +9,10 @@ mod request_candidate;
pub use affinity::{
build_scheduler_affinity_cache_key_for_api_key_id,
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session, candidate_affinity_hash,
candidate_key, matches_affinity_target, 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,
candidate_affinity_hash, candidate_key, matches_affinity_target, ClientSessionAffinity,
SchedulerAffinityScope, SchedulerAffinityTarget,
};
pub use auth::{
api_format_matches_allowed_value, auth_constraints_allow_api_format,
+2
View File
@@ -103,6 +103,8 @@ export interface PoolPresetMeta {
label: string
description: string
providers: string[]
default_enabled?: boolean
default_enabled_providers?: string[]
modes?: PoolPresetModeMeta[] | null
default_mode?: string | null
mutex_group?: string | null
@@ -530,7 +530,8 @@ import { CircleHelp } from 'lucide-vue-next'
import { Dialog, Button, Input, Label, Switch, Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from '@/components/ui'
import { useToast } from '@/composables/useToast'
import { parseApiError } from '@/utils/errorParser'
import { updateProvider } from '@/api/endpoints'
import { getProvider, updateProvider } from '@/api/endpoints'
import { mergePoolAdvancedPatch } from '@/features/pool/utils/poolSchedulingDialog'
import {
buildPoolCooldownFieldLayout,
buildPoolHealthToggleCards,
@@ -557,6 +558,7 @@ const emit = defineEmits<{
const { success, error: showError } = useToast()
const loading = ref(false)
let dialogRevision = 0
const isClaudeCode = computed(() => {
return (props.providerType || '').trim().toLowerCase() === 'claude_code'
@@ -652,7 +654,9 @@ function updateHealthToggleValue(key: PoolHealthToggleKey, value: boolean): void
}
}
watch(() => props.modelValue, (open) => {
watch([() => props.modelValue, () => props.providerId], ([open]) => {
dialogRevision += 1
loading.value = false
if (!open) return
const cfg = props.currentConfig
@@ -698,12 +702,26 @@ watch(() => props.modelValue, (open) => {
})
async function handleSave() {
const providerId = props.providerId
const revision = dialogRevision
loading.value = true
try {
const latestProvider = await getProvider(providerId)
if (!props.modelValue || props.providerId !== providerId || dialogRevision !== revision) return
const latestAdvanced = (latestProvider as Record<string, unknown>).pool_advanced
const existingPoolAdvanced = mergePoolAdvancedPatch(latestAdvanced, {})
const latestScoreRules = typeof existingPoolAdvanced.score_rules === 'object'
&& existingPoolAdvanced.score_rules !== null
? existingPoolAdvanced.score_rules as Record<string, unknown>
: {}
const latestScoreWeights = typeof latestScoreRules.weights === 'object'
&& latestScoreRules.weights !== null
? latestScoreRules.weights as Record<string, unknown>
: {}
const scoreRules = {
...(props.currentConfig?.score_rules ?? {}),
...latestScoreRules,
weights: {
...(props.currentConfig?.score_rules?.weights ?? {}),
...latestScoreWeights,
manual_priority: form.value.score_weight_manual_priority ?? undefined,
health: form.value.score_weight_health ?? undefined,
probe_freshness: form.value.score_weight_probe_freshness ?? undefined,
@@ -717,31 +735,8 @@ async function handleSave() {
request_failure_penalty: form.value.request_failure_penalty ?? undefined,
probe_failure_cooldown_threshold: form.value.probe_failure_cooldown_threshold ?? undefined,
}
const existingPoolAdvanced: Record<string, unknown> = { ...(props.currentConfig ?? {}) }
for (const key of [
'probing_target_percent',
'probing_target_count',
'probing_active_target_percent',
'probing_active_target_count',
'active_probe_target_percent',
'active_probe_target_count',
'probing_interval_minutes',
'account_self_check_method',
'self_check_method',
'account_self_check_request',
'self_check_request',
'health_policy_enabled',
'sticky_session_ttl_seconds',
'global_priority',
'cost_window_seconds',
'cost_limit_per_key_tokens',
'cost_soft_threshold_percent',
]) {
delete existingPoolAdvanced[key]
}
// scheduling_presets
const poolAdvanced: Record<string, unknown> = {
...existingPoolAdvanced,
const poolAdvanced = mergePoolAdvancedPatch(existingPoolAdvanced, {
rate_limit_cooldown_seconds: form.value.rate_limit_cooldown_seconds ?? undefined,
overload_cooldown_seconds: form.value.overload_cooldown_seconds ?? undefined,
batch_concurrency: form.value.batch_concurrency ?? undefined,
@@ -760,7 +755,7 @@ async function handleSave() {
auto_remove_banned_keys: form.value.auto_remove_banned_keys,
auto_remove_quota_exhausted_keys: form.value.auto_remove_quota_exhausted_keys,
skip_exhausted_accounts: form.value.skip_exhausted_accounts,
}
})
const payload: Parameters<typeof updateProvider>[1] = {
pool_advanced: poolAdvanced as PoolAdvancedConfig,
@@ -776,14 +771,17 @@ async function handleSave() {
cli_only_enabled: cf.cli_only_enabled,
}
}
const updatedProvider = await updateProvider(props.providerId, payload)
const updatedProvider = await updateProvider(providerId, payload)
if (!props.modelValue || props.providerId !== providerId || dialogRevision !== revision) return
success('高级设置已保存')
emit('saved', updatedProvider)
emit('update:modelValue', false)
} catch (err) {
showError(parseApiError(err))
} finally {
loading.value = false
if (dialogRevision === revision) {
loading.value = false
}
}
}
</script>
@@ -218,9 +218,13 @@ import { GripVertical } from 'lucide-vue-next'
import { Dialog, Button, Switch } from '@/components/ui'
import { useToast } from '@/composables/useToast'
import { parseApiError } from '@/utils/errorParser'
import { updateProvider } from '@/api/endpoints'
import { getProvider, updateProvider } from '@/api/endpoints'
import { getPoolSchedulingPresets } from '@/api/endpoints/pool'
import { moveStrategyItem } from '@/features/pool/utils/poolSchedulingDialog'
import {
mergePoolAdvancedPatch,
moveStrategyItem,
normalizeMutexSelection,
} from '@/features/pool/utils/poolSchedulingDialog'
import type { PoolPresetMeta } from '@/api/endpoints/pool'
import type {
PoolAdvancedConfig,
@@ -267,6 +271,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
mutex_group: DISTRIBUTION_GROUP,
evidence_hint: '依据 LRU 时间戳(最近使用优先,与 LRU 轮转相反)',
providers: [],
default_enabled: true,
modes: null,
default_mode: null,
},
@@ -300,12 +305,25 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
modes: null,
default_mode: null,
},
{
name: 'free_team_first',
label: 'Free/Team 优先',
description: '兼容旧配置:优先消耗 Free、Team 或两者',
evidence_hint: '依据 plan_type,保留旧 free_only/team_only/both 语义',
providers: ['codex', 'grok', 'kiro', 'windsurf'],
modes: [
{ value: 'free_only', label: 'Free' },
{ value: 'team_only', label: 'Team' },
{ value: 'both', label: 'Free + Team' },
],
default_mode: 'both',
},
{
name: 'free_first',
label: 'Free 优先',
description: '优先消耗 Free 账号(依赖 plan_type',
evidence_hint: '依据 plan_typeFree 账号优先调度)',
providers: ['codex', 'kiro'],
providers: ['codex', 'grok', 'kiro', 'windsurf'],
modes: null,
default_mode: null,
},
@@ -314,7 +332,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
label: 'Team 优先',
description: '优先消耗 Team 账号(依赖 plan_type',
evidence_hint: '依据 plan_typeTeam 账号优先调度)',
providers: ['codex', 'kiro'],
providers: ['codex', 'grok', 'kiro', 'windsurf'],
modes: null,
default_mode: null,
},
@@ -323,7 +341,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
label: 'Plus 优先',
description: '优先消耗 Plus 账号(依赖 plan_type',
evidence_hint: '依据 plan_typePlus 账号优先调度)',
providers: ['codex', 'kiro'],
providers: ['codex', 'grok', 'kiro', 'windsurf'],
modes: null,
default_mode: null,
},
@@ -332,7 +350,7 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
label: 'Pro 优先',
description: '优先消耗 Pro 账号(依赖 plan_type',
evidence_hint: '依据 plan_typePro 账号优先调度)',
providers: ['codex', 'kiro'],
providers: ['codex', 'grok', 'kiro', 'windsurf'],
modes: null,
default_mode: null,
},
@@ -350,7 +368,8 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
label: '额度刷新优先',
description: '优先选即将刷新额度的账号',
evidence_hint: '依据账号额度重置倒计时(next_reset / reset_seconds',
providers: ['codex', 'kiro'],
providers: ['codex', 'grok', 'kiro', 'windsurf'],
default_enabled_providers: ['codex', 'windsurf'],
modes: null,
default_mode: null,
},
@@ -392,10 +411,9 @@ const FALLBACK_PRESET_DEFS: PoolPresetMeta[] = [
},
]
const DEFAULT_ENABLED_PRESETS = new Set(['cache_affinity', 'recent_refresh'])
const { success, error: showError } = useToast()
const loading = ref(false)
let dialogRevision = 0
const presetDefs = ref<PoolPresetMeta[]>([])
const presetDefsLoaded = ref(false)
const loadingPresetDefs = ref(false)
@@ -443,11 +461,16 @@ function normalizePresetDefs(defs: PoolPresetMeta[]): PoolPresetMeta[] {
.filter(mode => Boolean(mode.value))
: null
const defaultMode = normalizeMode(raw.default_mode)
const defaultEnabledProviders = Array.isArray(raw.default_enabled_providers)
? raw.default_enabled_providers.map(p => normalizeProviderType(p)).filter(Boolean)
: []
ordered.push({
name,
label: String(raw.label ?? '').trim() || name,
description: String(raw.description ?? '').trim(),
providers,
default_enabled: raw.default_enabled === true,
default_enabled_providers: defaultEnabledProviders,
modes: modes && modes.length > 0 ? modes : null,
default_mode: defaultMode,
mutex_group: normalizeMutexGroup(raw.mutex_group),
@@ -478,11 +501,11 @@ async function ensurePresetDefsLoaded(): Promise<void> {
const normalized = normalizePresetDefs(Array.isArray(remoteDefs) ? remoteDefs : [])
if (normalized.length > 0) {
presetDefs.value = normalized
presetDefsLoaded.value = true
}
} catch (err) {
showError(parseApiError(err))
} finally {
presetDefsLoaded.value = true
loadingPresetDefs.value = false
}
}
@@ -530,7 +553,15 @@ function buildPresetListItem(def: PoolPresetMeta, enabled: boolean, mode?: unkno
}
function buildDefaultPresetList(): PresetListItem[] {
return getPresetDefs().map(def => buildPresetListItem(def, DEFAULT_ENABLED_PRESETS.has(def.name)))
const providerType = normalizeProviderType(props.providerType)
return getPresetDefs().map((def) => {
const providerDefaults = Array.isArray(def.default_enabled_providers)
? def.default_enabled_providers.map(normalizeProviderType)
: []
const enabled = def.default_enabled === true
|| (Boolean(providerType) && providerDefaults.includes(providerType))
return buildPresetListItem(def, enabled)
})
}
function isNewFormatPresetItem(item: unknown): item is SchedulingPresetItem {
@@ -641,7 +672,7 @@ function loadFromConfig(cfg: PoolAdvancedConfig | null): PresetListItem[] {
}
insertMissingByPreferredOrder(ordered, seen, defs, defsByName)
return reorderDistributionGroup(ordered)
return reorderDistributionGroup(normalizeMutexSelection(ordered))
}
const legacyPresets = rawPresets as string[]
@@ -664,33 +695,7 @@ function loadFromConfig(cfg: PoolAdvancedConfig | null): PresetListItem[] {
}
insertMissingByPreferredOrder(ordered, seen, defs, defsByName)
return reorderDistributionGroup(ordered)
}
function normalizeMutexSelection(items: PresetListItem[]): PresetListItem[] {
const next = [...items]
const groups = new Map<string, number[]>()
next.forEach((item, index) => {
if (!item.mutexGroup) return
if (!groups.has(item.mutexGroup)) groups.set(item.mutexGroup, [])
groups.get(item.mutexGroup)?.push(index)
})
for (const indexes of groups.values()) {
if (indexes.length <= 1) continue
const enabledApplicable = indexes.find(index => {
const item = next[index]
return item.enabled && item.applicable
})
const firstApplicable = indexes.find(index => next[index].applicable)
const winner = enabledApplicable ?? firstApplicable ?? indexes[0]
indexes.forEach((index) => {
next[index].enabled = index === winner && next[index].applicable
})
}
return next
return reorderDistributionGroup(normalizeMutexSelection(ordered))
}
function togglePreset(index: number, enabled: boolean) {
@@ -814,15 +819,26 @@ function handleDrop(dropIndex: number) {
dragOverIndex.value = null
}
watch(() => props.modelValue, async (open) => {
watch([() => props.modelValue, () => props.providerId], async ([open]) => {
const revision = ++dialogRevision
loading.value = false
if (!open) return
await ensurePresetDefsLoaded()
if (!props.modelValue || dialogRevision !== revision) return
presetList.value = normalizeMutexSelection(loadFromConfig(props.currentConfig))
})
async function handleSave() {
const providerId = props.providerId
const revision = dialogRevision
loading.value = true
try {
await ensurePresetDefsLoaded()
if (!props.modelValue || props.providerId !== providerId || dialogRevision !== revision) return
if (!presetDefsLoaded.value) {
showError('调度策略元数据加载失败,请重试')
return
}
presetList.value = normalizeMutexSelection(presetList.value)
const schedulingPresets: SchedulingPresetItem[] = presetList.value.map(item => {
const result: SchedulingPresetItem = {
@@ -835,14 +851,17 @@ async function handleSave() {
return result
})
const mergedAdvanced: Record<string, unknown> = {
...(props.currentConfig ?? {}),
const latestProvider = await getProvider(providerId)
if (!props.modelValue || props.providerId !== providerId || dialogRevision !== revision) return
const latestAdvanced = (latestProvider as Record<string, unknown>).pool_advanced
const mergedAdvanced = mergePoolAdvancedPatch(latestAdvanced, {
scheduling_presets: schedulingPresets,
}
})
const payload: Parameters<typeof updateProvider>[1] = {
pool_advanced: mergedAdvanced as PoolAdvancedConfig,
}
const updatedProvider = await updateProvider(props.providerId, payload)
const updatedProvider = await updateProvider(providerId, payload)
if (!props.modelValue || props.providerId !== providerId || dialogRevision !== revision) return
success('号池调度已保存')
emit('saved', updatedProvider)
@@ -850,7 +869,9 @@ async function handleSave() {
} catch (err) {
showError(parseApiError(err))
} finally {
loading.value = false
if (dialogRevision === revision) {
loading.value = false
}
}
}
</script>
@@ -1,22 +1,27 @@
import { describe, expect, it } from 'vitest'
import { moveStrategyItem } from '@/features/pool/utils/poolSchedulingDialog'
import {
mergePoolAdvancedPatch,
moveStrategyItem,
normalizeMutexSelection,
} from '@/features/pool/utils/poolSchedulingDialog'
interface TestPresetItem {
preset: string
mutexGroup: string | null
enabled: boolean
applicable: boolean
}
function buildItems(): TestPresetItem[] {
return [
{ preset: 'cache_affinity', mutexGroup: 'distribution_mode', enabled: false },
{ preset: 'lru', mutexGroup: 'distribution_mode', enabled: true },
{ preset: 'single_account', mutexGroup: 'distribution_mode', enabled: false },
{ preset: 'load_balance', mutexGroup: 'distribution_mode', enabled: false },
{ preset: 'recent_refresh', mutexGroup: null, enabled: true },
{ preset: 'quota_balanced', mutexGroup: null, enabled: false },
{ preset: 'priority_first', mutexGroup: null, enabled: true },
{ preset: 'cache_affinity', mutexGroup: 'distribution_mode', enabled: false, applicable: true },
{ preset: 'lru', mutexGroup: 'distribution_mode', enabled: true, applicable: true },
{ preset: 'single_account', mutexGroup: 'distribution_mode', enabled: false, applicable: true },
{ preset: 'load_balance', mutexGroup: 'distribution_mode', enabled: false, applicable: true },
{ preset: 'recent_refresh', mutexGroup: null, enabled: true, applicable: true },
{ preset: 'quota_balanced', mutexGroup: null, enabled: false, applicable: true },
{ preset: 'priority_first', mutexGroup: null, enabled: true, applicable: true },
]
}
@@ -69,4 +74,49 @@ describe('poolSchedulingDialog', () => {
expect(moved.map(item => item.preset)).toEqual(original.map(item => item.preset))
})
it('keeps the first enabled distribution from the saved order', () => {
const items = buildItems()
items[0].enabled = true
items[1].enabled = false
items[3].enabled = true
const savedOrder = [items[3], items[0], ...items.slice(1, 3), ...items.slice(4)]
const normalized = normalizeMutexSelection(savedOrder)
expect(normalized.find(item => item.preset === 'load_balance')?.enabled).toBe(true)
expect(normalized.find(item => item.preset === 'cache_affinity')?.enabled).toBe(false)
})
it('does not invent a distribution mode when all are disabled', () => {
const items = buildItems().map(item => ({
...item,
enabled: item.mutexGroup ? false : item.enabled,
}))
const normalized = normalizeMutexSelection(items)
expect(normalized.filter(item => item.mutexGroup).every(item => !item.enabled)).toBe(true)
})
it('preserves pool fields that the current dialog does not edit', () => {
const merged = mergePoolAdvancedPatch({
sticky_session_ttl_seconds: 900,
cost_window_seconds: 7200,
cost_limit_per_key_tokens: 100_000,
probing_target_percent: 25,
global_priority: 7,
}, {
score_top_n: 256,
})
expect(merged).toEqual({
sticky_session_ttl_seconds: 900,
cost_window_seconds: 7200,
cost_limit_per_key_tokens: 100_000,
probing_target_percent: 25,
global_priority: 7,
score_top_n: 256,
})
})
})
@@ -2,6 +2,47 @@ export interface SchedulingDialogPresetLike {
mutexGroup: string | null
}
export interface SchedulingDialogSelectablePresetLike extends SchedulingDialogPresetLike {
enabled: boolean
applicable: boolean
}
export function mergePoolAdvancedPatch(
current: unknown,
patch: Record<string, unknown>,
): Record<string, unknown> {
const currentRecord = typeof current === 'object' && current !== null && !Array.isArray(current)
? current as Record<string, unknown>
: {}
return {
...currentRecord,
...patch,
}
}
export function normalizeMutexSelection<T extends SchedulingDialogSelectablePresetLike>(
items: readonly T[],
): T[] {
const next = items.map(item => ({ ...item }))
const groups = new Map<string, number[]>()
next.forEach((item, index) => {
if (!item.mutexGroup) return
const indexes = groups.get(item.mutexGroup) ?? []
indexes.push(index)
groups.set(item.mutexGroup, indexes)
})
for (const indexes of groups.values()) {
const winner = indexes.find(index => next[index].enabled && next[index].applicable)
indexes.forEach((index) => {
next[index].enabled = winner !== undefined && index === winner && next[index].applicable
})
}
return next
}
export function moveStrategyItem<T extends SchedulingDialogPresetLike>(
items: readonly T[],
itemIndex: number,
+35 -12
View File
@@ -25,7 +25,7 @@
@import="showImportDialog = true"
@scheduling="openSchedulingDialog"
@demand-metrics="showDemandMetricsDialog = true"
@advanced="showAdvancedDialog = true"
@advanced="openAdvancedDialog"
@toggle-select-all="toggleAllFilteredPoolKeys"
@batch-action="openAccountBatchDialog"
@refresh="refreshCurrentPage"
@@ -1368,13 +1368,12 @@ async function loadOverview(options: { cacheTtlMs?: number, silent?: boolean } =
}
async function handleSchedulingSaved(updatedProvider: ProviderWithEndpointsSummary) {
if (!selectedProviderId.value || updatedProvider.id !== selectedProviderId.value) return
//
if (selectedProviderId.value && updatedProvider.id === selectedProviderId.value) {
if (selectedProviderData.value) {
Object.assign(selectedProviderData.value, updatedProvider)
} else {
selectedProviderData.value = updatedProvider
}
if (selectedProviderData.value) {
Object.assign(selectedProviderData.value, updatedProvider)
} else {
selectedProviderData.value = updatedProvider
}
showSchedulingDialog.value = false
showAdvancedDialog.value = false
@@ -1405,10 +1404,13 @@ const selectedProviderClaudeConfig = computed(() => {
return (selectedProviderData.value as Record<string, unknown> | null)?.claude_code_advanced as ClaudeCodeAdvancedConfig | null ?? null
})
const DEFAULT_ENABLED_PRESETS = new Set(['cache_affinity', 'recent_refresh'])
function defaultEnabledPresetCount(providerType: string): number {
return ['codex', 'windsurf'].includes(providerType) ? 2 : 1
}
const DEFAULT_PRESET_LABELS: Record<string, string> = {
lru: 'LRU',
free_team_first: 'Free/Team',
free_first: 'Free',
team_first: 'Team',
plus_first: 'Plus',
@@ -1546,7 +1548,7 @@ const poolSchedulingLabel = computed(() => {
const cfg = selectedProviderConfig.value
// No pool_advanced config at all: use default enabled presets count
if (!cfg) return `${DEFAULT_ENABLED_PRESETS.size} 维度`
if (!cfg) return `${defaultEnabledPresetCount(selectedProviderType.value)} 维度`
const presets = Array.isArray(cfg.scheduling_presets) ? cfg.scheduling_presets : []
const presetLabels = presetLabelsByName.value
@@ -1582,7 +1584,7 @@ const poolSchedulingLabel = computed(() => {
if (lruEnabled && stickyEnabled) return 'LRU + 粘性'
if (lruEnabled) return 'LRU'
if (!cfg.scheduling_mode && (cfg.lru_enabled === null || cfg.lru_enabled === undefined)) {
return `${DEFAULT_ENABLED_PRESETS.size} 维度`
return `${defaultEnabledPresetCount(selectedProviderType.value)} 维度`
}
if (stickyEnabled) return '粘性'
return '随机'
@@ -1712,6 +1714,8 @@ async function selectProvider(
hasHydratedInitialProviderSelection = true
selectedProviderId.value = id
selectedProviderData.value = null
showSchedulingDialog.value = false
showAdvancedDialog.value = false
resetPoolKeySelection(true)
providerDrawerOpen.value = false
editingKeyDetail.value = null
@@ -2873,8 +2877,27 @@ const showAccountBatchDialog = ref(false)
const pendingAccountBatchAction = ref<PoolBatchActionValue | null>(null)
const togglingProviderStatus = ref(false)
function openSchedulingDialog() {
showSchedulingDialog.value = true
async function ensureSelectedProviderDetail(): Promise<boolean> {
const providerId = selectedProviderId.value
if (!providerId) return false
if (selectedProviderData.value?.id !== providerId) {
await loadProviderData(providerId, { preserveOnError: true })
}
if (selectedProviderData.value?.id === providerId) return true
showWarning('Provider 详情尚未加载,无法编辑调度配置')
return false
}
async function openSchedulingDialog() {
if (await ensureSelectedProviderDetail()) {
showSchedulingDialog.value = true
}
}
async function openAdvancedDialog() {
if (await ensureSelectedProviderDetail()) {
showAdvancedDialog.value = true
}
}
function openAccountBatchDialog(action: PoolBatchActionValue = 'refresh_quota'): void {