mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-12 14:10:19 +08:00
fix(routing): harden routed pool scheduling
This commit is contained in:
@@ -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")
|
||||
|
||||
+1
@@ -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
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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";
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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();
|
||||
|
||||
+115
@@ -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))
|
||||
}
|
||||
|
||||
|
||||
+36
@@ -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))
|
||||
}
|
||||
|
||||
|
||||
+38
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
|
||||
@@ -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)
|
||||
{
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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");
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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_type(Free 账号优先调度)',
|
||||
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_type(Team 账号优先调度)',
|
||||
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_type(Plus 账号优先调度)',
|
||||
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_type(Pro 账号优先调度)',
|
||||
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,
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user