Merge remote-tracking branch 'origin/aether-rust-pioneer' into payment-billing-plans

# Conflicts:
#	crates/aether-data/src/lifecycle/bootstrap/postgres.rs
#	crates/aether-data/src/lifecycle/migrate/tests.rs
This commit is contained in:
Entropy.Xu
2026-05-13 01:29:18 +08:00
185 changed files with 14980 additions and 3638 deletions

View File

@@ -3,6 +3,7 @@ pub(crate) use crate::handlers::admin::{
build_internal_control_error_response, create_provider_oauth_catalog_key,
find_duplicate_provider_oauth_key, maybe_build_local_admin_pool_response,
maybe_build_local_admin_response, provider_oauth_runtime_endpoint_for_provider,
provider_type_supports_quota_refresh, reconcile_admin_fixed_provider_template_endpoints,
refresh_antigravity_provider_quota_locally, refresh_chatgpt_web_provider_quota_locally,
refresh_codex_provider_quota_locally, refresh_kiro_provider_quota_locally,
refresh_provider_oauth_account_state_after_update, update_existing_provider_oauth_catalog_key,

View File

@@ -20,8 +20,8 @@ pub(crate) use self::finalize::internal::{
SyncToStreamBridgeOutcome,
};
pub(crate) use self::planner::{
build_gemini_stream_plan_from_decision, build_gemini_sync_plan_from_decision,
build_local_gemini_files_stream_attempt_source_for_kind,
apply_local_runtime_candidate_terminal_reason, build_gemini_stream_plan_from_decision,
build_gemini_sync_plan_from_decision, build_local_gemini_files_stream_attempt_source_for_kind,
build_local_gemini_files_stream_plan_and_reports_for_kind,
build_local_gemini_files_sync_attempt_source_for_kind,
build_local_gemini_files_sync_plan_and_reports_for_kind,
@@ -43,16 +43,20 @@ pub(crate) use self::planner::{
build_local_video_sync_plan_and_reports_for_kind,
build_openai_responses_stream_plan_from_decision,
build_openai_responses_sync_plan_from_decision, build_passthrough_sync_plan_from_decision,
build_standard_family_stream_attempt_source, build_standard_family_stream_plan_and_reports,
build_standard_family_sync_attempt_source, build_standard_family_sync_plan_and_reports,
build_standard_stream_plan_from_decision, build_standard_sync_plan_from_decision,
build_provider_key_pool_score_upsert, build_standard_family_stream_attempt_source,
build_standard_family_stream_plan_and_reports, build_standard_family_sync_attempt_source,
build_standard_family_sync_plan_and_reports, build_standard_stream_plan_from_decision,
build_standard_sync_plan_from_decision, candidate_auth_channel_skip_reason,
extract_pool_sticky_session_token, maybe_build_stream_decision_payload,
maybe_build_stream_plan_payload, maybe_build_sync_decision_payload,
maybe_build_sync_plan_payload, planner_is_matching_stream_request,
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,
set_local_openai_chat_execution_exhausted_diagnostic,
set_local_openai_image_execution_exhausted_diagnostic, CandidateFailureDiagnostic,
CandidateFailureDiagnosticKind, GatewayAuthApiKeySnapshot, GatewayProviderTransportSnapshot,
LocalExecutionAttemptSource, LocalResolvedOAuthRequestAuth, PlannerAppState,
CandidateFailureDiagnosticKind, EligibleLocalExecutionCandidate, GatewayAuthApiKeySnapshot,
GatewayProviderTransportSnapshot, LocalExecutionAttemptSource, LocalExecutionCandidateKind,
LocalResolvedOAuthRequestAuth, PlannerAppState, SkippedLocalExecutionCandidate,
};
pub(crate) use self::pure::*;
pub(crate) use self::transport::{

View File

@@ -4,14 +4,17 @@ use aether_ai_serving::{
run_ai_available_candidate_persistence, run_ai_candidate_materialization,
run_ai_skipped_candidate_persistence, AiAvailableCandidatePersistencePort,
AiCandidateMaterializationOutcome, AiCandidateMaterializationPort,
AiSkippedCandidatePersistencePort,
AiCandidatePreselectionOutcome, AiSkippedCandidatePersistencePort,
};
use aether_dispatch_core::{DispatchSequence, DispatchSequenceItem};
use aether_scheduler_core::{ClientSessionAffinity, SchedulerMinimalCandidateSelectionCandidate};
use async_trait::async_trait;
use serde_json::Value;
use std::collections::VecDeque;
use std::convert::Infallible;
use std::sync::Arc;
use std::time::Duration;
use tokio::time::Instant;
use tracing::warn;
use uuid::Uuid;
@@ -29,12 +32,16 @@ use crate::ai_serving::planner::runtime_miss::record_local_runtime_candidate_ski
use crate::ai_serving::planner::CandidateFailureDiagnostic;
use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState};
use crate::clock::current_unix_ms;
use crate::dispatch::refs::dispatch_ref_for_local_candidate;
use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value;
use crate::orchestration::{local_attempt_slot_count, ExecutionAttemptIdentity};
use crate::scheduler::candidate::API_KEY_CONCURRENCY_LIMIT_SKIP_REASON;
use crate::scheduler::config::{read_scheduler_ordering_config, SchedulerSchedulingMode};
use crate::{AppState, GatewayError};
const POOL_KEY_RETRY_INDEX_STRIDE: u32 = 100;
const AUTH_API_KEY_CONCURRENCY_WAIT_BUDGET: Duration = Duration::from_millis(100);
const AUTH_API_KEY_CONCURRENCY_RETRY_DELAY: Duration = Duration::from_millis(10);
#[derive(Debug, Clone)]
pub(crate) struct LocalExecutionCandidateAttempt {
@@ -61,12 +68,12 @@ pub(crate) trait LocalExecutionAttemptSource<T>: Send {
enum LocalExecutionCandidateAttemptSourceItem<'a> {
Static {
attempts: VecDeque<LocalExecutionCandidateAttempt>,
attempts: DispatchSequence<LocalExecutionCandidateAttempt>,
},
Pool {
cursor: PoolKeyCursor<'a>,
candidate_index: u32,
pending_attempts: VecDeque<LocalExecutionCandidateAttempt>,
pending_attempts: DispatchSequence<LocalExecutionCandidateAttempt>,
},
RequestedModelPage {
cursor: Box<RequestedModelAttemptPageCursor<'a>>,
@@ -80,7 +87,7 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
let mut items = VecDeque::new();
if !attempts.is_empty() {
items.push_back(LocalExecutionCandidateAttemptSourceItem::Static {
attempts: VecDeque::from(attempts),
attempts: dispatch_sequence_from_attempts(attempts),
});
}
Self { items }
@@ -91,8 +98,8 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
let front = self.items.front_mut()?;
match front {
LocalExecutionCandidateAttemptSourceItem::Static { attempts } => {
if let Some(attempt) = attempts.pop_front() {
if attempts.is_empty() {
if let Some(attempt) = next_attempt_from_dispatch_sequence(attempts) {
if dispatch_sequence_exhausted(attempts) {
self.items.pop_front();
}
return Some(attempt);
@@ -104,7 +111,7 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
candidate_index,
pending_attempts,
} => {
if let Some(attempt) = pending_attempts.pop_front() {
if let Some(attempt) = next_attempt_from_dispatch_sequence(pending_attempts) {
return Some(attempt);
}
let Some(candidate) = cursor.next_key().await else {
@@ -113,9 +120,12 @@ impl<'a> LocalExecutionCandidateAttemptSource<'a> {
self.items.pop_front();
continue;
};
*pending_attempts = build_unpersisted_local_execution_candidate_attempts(
candidate,
*candidate_index,
*pending_attempts = dispatch_sequence_from_attempts(
build_unpersisted_local_execution_candidate_attempts(
candidate,
*candidate_index,
)
.into(),
);
}
LocalExecutionCandidateAttemptSourceItem::RequestedModelPage { cursor } => {
@@ -311,10 +321,7 @@ where
}
fn build_extra_data(&self, candidate: &Self::Candidate) -> Option<Self::ExtraData> {
ai_candidate_extra_data_with_ranking(
(self.build_extra_data)(candidate),
candidate.ranking.as_ref(),
)
available_candidate_extra_data_with_dispatch_ref(candidate, &self.build_extra_data)
}
fn generate_candidate_id(&self) -> String {
@@ -506,15 +513,6 @@ where
.map(decorate_skipped_candidate)
.collect::<Vec<_>>();
let candidate_count = candidates.len() + skipped_candidate_count;
if persistence_policy.skipped.record_runtime_miss_diagnostic {
for skipped_candidate in &skipped_candidates {
record_local_runtime_candidate_skip_reason(
state.app(),
trace_id,
skipped_candidate.skip_reason,
);
}
}
if scheduler_cache_affinity_enabled {
remember_first_local_candidate_affinity(
@@ -526,6 +524,14 @@ where
&candidates,
);
}
persist_skipped_local_execution_candidates_with_context(
state.app(),
trace_id,
persistence_policy.skipped,
u32::try_from(candidates.len()).unwrap_or(u32::MAX),
skipped_candidates,
)
.await;
let (items, _) = build_logical_candidate_items(
state,
@@ -566,7 +572,9 @@ fn build_logical_candidate_items<'a>(
candidate_index,
);
if !attempts.is_empty() {
items.push_back(LocalExecutionCandidateAttemptSourceItem::Static { attempts });
items.push_back(LocalExecutionCandidateAttemptSourceItem::Static {
attempts: dispatch_sequence_from_attempts(attempts.into()),
});
}
}
LocalExecutionCandidateKind::PoolGroup => {
@@ -585,7 +593,7 @@ fn build_logical_candidate_items<'a>(
items.push_back(LocalExecutionCandidateAttemptSourceItem::Pool {
cursor,
candidate_index,
pending_attempts: VecDeque::new(),
pending_attempts: DispatchSequence::new(Vec::new()),
});
}
}
@@ -646,6 +654,10 @@ where
required_capabilities: required_capabilities.cloned(),
sticky_session_token: sticky_session_token.map(str::to_string),
request_auth_channel: request_auth_channel.map(str::to_string),
skipped_user_id: persistence_policy.skipped.user_id.to_string(),
skipped_api_key_id: persistence_policy.skipped.api_key_id.to_string(),
skipped_required_capabilities: persistence_policy.skipped.required_capabilities.cloned(),
skipped_error_context: persistence_policy.skipped.error_context,
record_runtime_miss_diagnostic,
resolution_mode,
decorate_skipped_candidate,
@@ -655,6 +667,7 @@ where
next_candidate_index: 0,
remembered_affinity: false,
scheduler_cache_affinity_enabled,
auth_api_key_concurrency_wait_deadline: None,
};
cursor.load_next_page().await;
let candidate_count = cursor.candidate_count;
@@ -682,6 +695,10 @@ struct RequestedModelAttemptPageCursor<'a> {
required_capabilities: Option<Value>,
sticky_session_token: Option<String>,
request_auth_channel: Option<String>,
skipped_user_id: String,
skipped_api_key_id: String,
skipped_required_capabilities: Option<Value>,
skipped_error_context: &'static str,
record_runtime_miss_diagnostic: bool,
resolution_mode: LocalCandidateResolutionMode,
decorate_skipped_candidate: DecorateSkippedCandidateFn<'a>,
@@ -691,6 +708,7 @@ struct RequestedModelAttemptPageCursor<'a> {
next_candidate_index: u32,
remembered_affinity: bool,
scheduler_cache_affinity_enabled: bool,
auth_api_key_concurrency_wait_deadline: Option<Instant>,
}
impl<'a> RequestedModelAttemptPageCursor<'a> {
@@ -720,6 +738,15 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
}
};
if page_is_exact_auth_api_key_concurrency_limited(&page) {
if self.wait_for_auth_api_key_concurrency_retry().await {
continue;
}
self.persist_final_auth_api_key_concurrency_skips(page.skipped_candidates)
.await;
return false;
}
let (candidates, resolved_skipped) =
resolve_and_rank_logical_local_execution_candidates(
self.state,
@@ -740,18 +767,10 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
.chain(resolved_skipped)
.map(|skipped| (self.decorate_skipped_candidate)(skipped))
.collect::<Vec<_>>();
let skipped_candidate_count = skipped_candidates.len();
self.candidate_count = self
.candidate_count
.saturating_add(candidates.len() + skipped_candidates.len());
if self.record_runtime_miss_diagnostic {
for skipped_candidate in &skipped_candidates {
record_local_runtime_candidate_skip_reason(
self.state.app(),
&self.trace_id,
skipped_candidate.skip_reason,
);
}
}
.saturating_add(candidates.len() + skipped_candidate_count);
if self.scheduler_cache_affinity_enabled
&& !self.remembered_affinity
&& !candidates.is_empty()
@@ -776,13 +795,90 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
Some(&self.requested_model),
self.request_auth_channel.as_deref(),
);
self.next_candidate_index = next_candidate_index;
self.next_candidate_index = next_candidate_index
.saturating_add(u32::try_from(skipped_candidate_count).unwrap_or(u32::MAX));
if !items.is_empty() {
self.pending_items = items;
return true;
}
let skipped_starting_candidate_index = next_candidate_index;
let skipped_persistence = LocalSkippedCandidatePersistenceContext {
user_id: self.skipped_user_id.as_str(),
api_key_id: self.skipped_api_key_id.as_str(),
required_capabilities: self.skipped_required_capabilities.as_ref(),
error_context: self.skipped_error_context,
record_runtime_miss_diagnostic: self.record_runtime_miss_diagnostic,
};
persist_skipped_local_execution_candidates_with_context(
self.state.app(),
&self.trace_id,
skipped_persistence,
skipped_starting_candidate_index,
skipped_candidates,
)
.await;
}
}
async fn wait_for_auth_api_key_concurrency_retry(&mut self) -> bool {
let now = Instant::now();
let deadline = *self
.auth_api_key_concurrency_wait_deadline
.get_or_insert(now + AUTH_API_KEY_CONCURRENCY_WAIT_BUDGET);
if now >= deadline {
return false;
}
let sleep_duration =
AUTH_API_KEY_CONCURRENCY_RETRY_DELAY.min(deadline.saturating_duration_since(now));
tokio::time::sleep(sleep_duration).await;
self.page_cursor.restart_scan();
true
}
async fn persist_final_auth_api_key_concurrency_skips(
&mut self,
skipped_candidates: Vec<SkippedLocalExecutionCandidate>,
) {
let skipped_candidates = skipped_candidates
.into_iter()
.map(|skipped| (self.decorate_skipped_candidate)(skipped))
.collect::<Vec<_>>();
let skipped_candidate_count = skipped_candidates.len();
self.candidate_count = self.candidate_count.saturating_add(skipped_candidate_count);
let skipped_persistence = LocalSkippedCandidatePersistenceContext {
user_id: self.skipped_user_id.as_str(),
api_key_id: self.skipped_api_key_id.as_str(),
required_capabilities: self.skipped_required_capabilities.as_ref(),
error_context: self.skipped_error_context,
record_runtime_miss_diagnostic: self.record_runtime_miss_diagnostic,
};
persist_skipped_local_execution_candidates_with_context(
self.state.app(),
&self.trace_id,
skipped_persistence,
self.next_candidate_index,
skipped_candidates,
)
.await;
self.next_candidate_index = self
.next_candidate_index
.saturating_add(u32::try_from(skipped_candidate_count).unwrap_or(u32::MAX));
}
}
fn page_is_exact_auth_api_key_concurrency_limited(
page: &AiCandidatePreselectionOutcome<
SchedulerMinimalCandidateSelectionCandidate,
SkippedLocalExecutionCandidate,
>,
) -> bool {
page.candidates.is_empty()
&& !page.skipped_candidates.is_empty()
&& page
.skipped_candidates
.iter()
.all(|skipped| skipped.skip_reason == API_KEY_CONCURRENCY_LIMIT_SKIP_REASON)
}
async fn pop_attempt_from_items(
@@ -792,8 +888,8 @@ async fn pop_attempt_from_items(
let front = items.front_mut()?;
match front {
LocalExecutionCandidateAttemptSourceItem::Static { attempts } => {
if let Some(attempt) = attempts.pop_front() {
if attempts.is_empty() {
if let Some(attempt) = next_attempt_from_dispatch_sequence(attempts) {
if dispatch_sequence_exhausted(attempts) {
items.pop_front();
}
return Some(attempt);
@@ -805,7 +901,7 @@ async fn pop_attempt_from_items(
candidate_index,
pending_attempts,
} => {
if let Some(attempt) = pending_attempts.pop_front() {
if let Some(attempt) = next_attempt_from_dispatch_sequence(pending_attempts) {
return Some(attempt);
}
let Some(candidate) = cursor.next_key().await else {
@@ -814,9 +910,12 @@ async fn pop_attempt_from_items(
items.pop_front();
continue;
};
*pending_attempts = build_unpersisted_local_execution_candidate_attempts(
candidate,
*candidate_index,
*pending_attempts = dispatch_sequence_from_attempts(
build_unpersisted_local_execution_candidate_attempts(
candidate,
*candidate_index,
)
.into(),
);
}
LocalExecutionCandidateAttemptSourceItem::RequestedModelPage { .. } => {
@@ -1005,7 +1104,7 @@ where
{
let attempt_slots = local_attempt_slot_count(&candidate.transport).max(1);
let extra_data = ai_candidate_extra_data_with_ranking(
build_extra_data(&candidate),
available_candidate_base_extra_data_with_dispatch_ref(&candidate, build_extra_data),
candidate.ranking.as_ref(),
);
let should_persist = should_persist_available_local_candidate(&candidate);
@@ -1057,6 +1156,70 @@ where
attempts
}
fn available_candidate_extra_data_with_dispatch_ref<F>(
candidate: &EligibleLocalExecutionCandidate,
build_extra_data: &F,
) -> Option<Value>
where
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync,
{
ai_candidate_extra_data_with_ranking(
available_candidate_base_extra_data_with_dispatch_ref(candidate, build_extra_data),
candidate.ranking.as_ref(),
)
}
fn available_candidate_base_extra_data_with_dispatch_ref<F>(
candidate: &EligibleLocalExecutionCandidate,
build_extra_data: &F,
) -> Option<Value>
where
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync,
{
let dispatch_ref = serde_json::to_value(dispatch_ref_for_local_candidate(candidate)).ok()?;
let mut object = match build_extra_data(candidate) {
Some(Value::Object(object)) => object,
Some(value) => {
let mut object = serde_json::Map::new();
object.insert("extra".to_string(), value);
object
}
None => serde_json::Map::new(),
};
object.insert("dispatch_ref".to_string(), dispatch_ref);
Some(Value::Object(object))
}
fn dispatch_sequence_from_attempts(
attempts: Vec<LocalExecutionCandidateAttempt>,
) -> DispatchSequence<LocalExecutionCandidateAttempt> {
DispatchSequence::new(
attempts
.into_iter()
.map(|attempt| DispatchSequenceItem {
candidate_index: attempt.candidate_index,
retry_index: attempt.retry_index,
candidate: attempt,
mark: aether_dispatch_core::DispatchSequenceMark::Pending,
})
.collect(),
)
}
fn next_attempt_from_dispatch_sequence(
sequence: &mut DispatchSequence<LocalExecutionCandidateAttempt>,
) -> Option<LocalExecutionCandidateAttempt> {
let attempt = sequence.next()?.candidate.clone();
let _ = sequence.mark_succeeded();
Some(attempt)
}
fn dispatch_sequence_exhausted(
sequence: &mut DispatchSequence<LocalExecutionCandidateAttempt>,
) -> bool {
sequence.next().is_none()
}
fn build_unpersisted_local_execution_candidate_attempts(
candidate: EligibleLocalExecutionCandidate,
candidate_index: u32,
@@ -1457,6 +1620,16 @@ mod tests {
assert_eq!(stored.len(), 1);
assert_eq!(stored[0].key_id.as_deref(), Some("normal-key"));
assert_eq!(stored[0].candidate_index, 2);
assert_eq!(
stored[0]
.extra_data
.as_ref()
.and_then(|value| value.get("dispatch_ref"))
.and_then(|value| value.get("SingleKey"))
.and_then(|value| value.get("key"))
.and_then(|value| value.get("key_id")),
Some(&json!("normal-key"))
);
}
#[test]
@@ -1562,6 +1735,16 @@ mod tests {
assert_eq!(stored.len(), 1);
assert_eq!(stored[0].key_id.as_deref(), Some("normal-key"));
assert_eq!(stored[0].candidate_index, 1);
assert_eq!(
stored[0]
.extra_data
.as_ref()
.and_then(|value| value.get("dispatch_ref"))
.and_then(|value| value.get("SingleKey"))
.and_then(|value| value.get("key"))
.and_then(|value| value.get("key_id")),
Some(&json!("normal-key"))
);
}
#[test]
@@ -1643,15 +1826,26 @@ mod tests {
Some(&json!("cached_affinity"))
);
assert_eq!(extra_data.get("demoted_by"), Some(&json!("cross_format")));
assert_eq!(
extra_data
.get("dispatch_ref")
.and_then(|value| value.get("SingleKey"))
.and_then(|value| value.get("key"))
.and_then(|value| value.get("key_id")),
Some(&json!("ranked-key"))
);
}
#[tokio::test]
async fn dynamic_attempt_source_does_not_drain_unexecuted_single_keys() {
let mut source = LocalExecutionCandidateAttemptSource {
items: VecDeque::from([LocalExecutionCandidateAttemptSourceItem::Static {
attempts: build_unpersisted_local_execution_candidate_attempts(
sample_eligible("normal-key", None),
0,
attempts: dispatch_sequence_from_attempts(
build_unpersisted_local_execution_candidate_attempts(
sample_eligible("normal-key", None),
0,
)
.into(),
),
}]),
};

View File

@@ -21,7 +21,6 @@ use crate::ai_serving::{
use crate::orchestration::LocalExecutionCandidateMetadata;
use super::candidate_ranking::rank_eligible_local_execution_candidates;
use super::pool_scheduler::apply_local_execution_pool_scheduler;
#[derive(Debug, Clone, PartialEq)]
pub(crate) struct EligibleLocalExecutionCandidate {
@@ -61,7 +60,6 @@ struct GatewayLocalCandidateResolutionPort<'a> {
auth_snapshot: Option<&'a GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&'a ClientSessionAffinity>,
required_capabilities: Option<&'a serde_json::Value>,
sticky_session_token: Option<&'a str>,
request_auth_channel: Option<&'a str>,
}
@@ -182,14 +180,7 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
&self,
candidates: Vec<Self::Eligible>,
) -> Result<(Vec<Self::Eligible>, Vec<Self::Skipped>), Self::Error> {
Ok(apply_local_execution_pool_scheduler(
self.state,
candidates,
self.sticky_session_token,
self.requested_model,
self.request_auth_channel,
)
.await)
Ok((candidates, Vec::new()))
}
}
@@ -201,7 +192,7 @@ pub(crate) async fn resolve_and_rank_local_execution_candidates(
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&serde_json::Value>,
sticky_session_token: Option<&str>,
_sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>,
) -> (
Vec<EligibleLocalExecutionCandidate>,
@@ -216,7 +207,7 @@ pub(crate) async fn resolve_and_rank_local_execution_candidates(
auth_snapshot,
client_session_affinity,
required_capabilities,
sticky_session_token,
None,
request_auth_channel,
AiCandidateResolutionMode::Standard,
)
@@ -231,7 +222,7 @@ pub(crate) async fn resolve_and_rank_local_execution_candidates_without_transpor
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&serde_json::Value>,
sticky_session_token: Option<&str>,
_sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>,
) -> (
Vec<EligibleLocalExecutionCandidate>,
@@ -246,7 +237,7 @@ pub(crate) async fn resolve_and_rank_local_execution_candidates_without_transpor
auth_snapshot,
client_session_affinity,
required_capabilities,
sticky_session_token,
None,
request_auth_channel,
AiCandidateResolutionMode::WithoutTransportPairGate,
)
@@ -261,7 +252,7 @@ pub(crate) async fn resolve_and_rank_logical_local_execution_candidates(
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&serde_json::Value>,
sticky_session_token: Option<&str>,
_sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>,
mode: AiCandidateResolutionMode,
) -> (
@@ -276,7 +267,7 @@ pub(crate) async fn resolve_and_rank_logical_local_execution_candidates(
auth_snapshot,
client_session_affinity,
required_capabilities,
sticky_session_token,
None,
request_auth_channel,
mode,
false,
@@ -292,7 +283,7 @@ async fn resolve_and_rank_local_execution_candidates_with_mode(
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&serde_json::Value>,
sticky_session_token: Option<&str>,
_sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>,
mode: AiCandidateResolutionMode,
) -> (
@@ -307,10 +298,10 @@ async fn resolve_and_rank_local_execution_candidates_with_mode(
auth_snapshot,
client_session_affinity,
required_capabilities,
sticky_session_token,
None,
request_auth_channel,
mode,
true,
false,
)
.await
}
@@ -324,7 +315,7 @@ async fn resolve_and_rank_local_execution_candidates_with_pool_expansion(
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&serde_json::Value>,
sticky_session_token: Option<&str>,
_sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>,
mode: AiCandidateResolutionMode,
expand_pool_groups: bool,
@@ -339,7 +330,6 @@ async fn resolve_and_rank_local_execution_candidates_with_pool_expansion(
auth_snapshot,
client_session_affinity,
required_capabilities,
sticky_session_token,
request_auth_channel,
};

View File

@@ -323,6 +323,16 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
Ok(None)
}
pub(crate) fn restart_scan(&mut self) {
self.format_index = 0;
self.requested_name_indexes.clear();
self.requested_name_offsets.clear();
self.scanned_rows_by_format.clear();
self.resolved_global_model_names.clear();
self.fallback_scanned_api_formats.clear();
self.seen_candidate_keys.clear();
}
async fn next_page_for_api_format(
&mut self,
candidate_api_format: &str,
@@ -735,7 +745,9 @@ mod tests {
use crate::data::GatewayDataState;
use crate::AppState;
use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelectionReadRepository;
use aether_data_contracts::repository::candidate_selection::MinimalCandidateSelectionReadRepository;
use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredProviderModelMapping,
};
use std::sync::Arc;
fn unrestricted_auth_snapshot() -> GatewayAuthApiKeySnapshot {
@@ -800,6 +812,57 @@ mod tests {
}
}
fn opg_deepseek_row(
endpoint_id: &str,
api_format: &str,
key_id: &str,
key_name: &str,
key_allowed_models: Vec<&str>,
key_internal_priority: i32,
) -> StoredMinimalCandidateSelectionRow {
StoredMinimalCandidateSelectionRow {
provider_id: "provider-opg".to_string(),
provider_name: "OpenCode Go".to_string(),
provider_type: "custom".to_string(),
provider_priority: 1,
provider_is_active: true,
endpoint_id: endpoint_id.to_string(),
endpoint_api_format: api_format.to_string(),
endpoint_api_family: None,
endpoint_kind: Some("chat".to_string()),
endpoint_is_active: true,
key_id: key_id.to_string(),
key_name: key_name.to_string(),
key_auth_type: "api_key".to_string(),
key_is_active: true,
key_api_formats: Some(vec![api_format.to_string()]),
key_allowed_models: Some(
key_allowed_models
.into_iter()
.map(ToOwned::to_owned)
.collect(),
),
key_capabilities: None,
key_internal_priority,
key_global_priority_by_format: None,
model_id: "model-opg-deepseek-v4-pro".to_string(),
global_model_id: "global-model-deepseek-v4-pro".to_string(),
global_model_name: "deepseek-v4-pro".to_string(),
global_model_mappings: None,
global_model_supports_streaming: Some(true),
model_provider_model_name: "deepseek-v4-pro".to_string(),
model_provider_model_mappings: Some(vec![StoredProviderModelMapping {
name: "deepseek-v4-pro".to_string(),
priority: 1,
api_formats: None,
endpoint_ids: Some(vec!["endpoint-opg-openai".to_string()]),
}]),
model_supports_streaming: Some(true),
model_is_active: true,
model_is_available: true,
}
}
#[tokio::test]
async fn paged_preselection_falls_back_to_format_scan_for_directive_mapping_match() {
let repository: Arc<dyn MinimalCandidateSelectionReadRepository> =
@@ -844,4 +907,60 @@ mod tests {
"gpt-5-upstream"
);
}
#[tokio::test]
async fn claude_request_uses_cross_format_key_when_same_provider_messages_key_lacks_model() {
let repository: Arc<dyn MinimalCandidateSelectionReadRepository> =
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed([
opg_deepseek_row(
"endpoint-opg-claude",
"claude:messages",
"key-opg-messages",
"OPG Key Messages",
vec!["glm-5", "glm-5.1", "minimax-m2.5", "minimax-m2.7"],
1,
),
opg_deepseek_row(
"endpoint-opg-openai",
"openai:chat",
"key-opg-completions",
"OPG Key Completions",
vec!["deepseek-v4-pro", "glm-5", "glm-5.1", "minimax-m2.7"],
10,
),
]));
let data_state =
GatewayDataState::with_minimal_candidate_selection_reader_for_tests(repository);
let app = AppState::new()
.expect("gateway state should build")
.with_data_state_for_tests(data_state);
let auth_snapshot = unrestricted_auth_snapshot();
let mut cursor = LocalCandidatePreselectionPageCursor::new(
PlannerAppState::new(&app),
"claude:messages",
"deepseek-v4-pro",
false,
None,
&auth_snapshot,
None,
true,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
)
.await;
let page = cursor
.next_page()
.await
.expect("preselection should succeed")
.expect("openai chat candidate should be found via conversion");
assert_eq!(page.skipped_candidates.len(), 0);
assert_eq!(page.candidates.len(), 1);
assert_eq!(page.candidates[0].endpoint_api_format, "openai:chat");
assert_eq!(page.candidates[0].key_name, "OPG Key Completions");
assert_eq!(
page.candidates[0].selected_provider_model_name,
"deepseek-v4-pro"
);
}
}

View File

@@ -16,6 +16,7 @@ mod materialization_policy;
mod passthrough;
mod plan_builders;
mod pool_scheduler;
pub(crate) mod pool_scores;
mod report_context;
mod route;
mod runtime_miss;
@@ -25,6 +26,10 @@ mod standard;
mod state;
pub(crate) use self::candidate_materialization::LocalExecutionAttemptSource;
pub(crate) use self::candidate_resolution::{
candidate_auth_channel_skip_reason, read_candidate_transport_snapshot,
EligibleLocalExecutionCandidate, LocalExecutionCandidateKind, SkippedLocalExecutionCandidate,
};
pub(crate) use self::passthrough::{
build_local_same_format_stream_attempt_source, build_local_same_format_stream_plan_and_reports,
build_local_same_format_sync_attempt_source, build_local_same_format_sync_plan_and_reports,
@@ -36,7 +41,13 @@ pub(crate) use self::plan_builders::{
build_standard_stream_plan_from_decision, build_standard_sync_plan_from_decision,
AiStreamAttempt, AiSyncAttempt,
};
pub(crate) use self::pool_scores::{
build_provider_key_pool_score_upsert, provider_key_pool_score_id, provider_key_pool_score_scope,
};
pub(crate) use self::route::is_matching_stream_request as planner_is_matching_stream_request;
pub(crate) use self::runtime_miss::{
apply_local_runtime_candidate_terminal_reason, record_local_runtime_candidate_skip_reason,
};
pub(crate) use self::specialized::{
build_local_gemini_files_stream_attempt_source_for_kind,
build_local_gemini_files_stream_plan_and_reports_for_kind,

View File

@@ -9,7 +9,7 @@ use crate::{AiExecutionDecision, AppState, GatewayError};
use super::super::plans::{resolve_stream_spec, resolve_sync_spec};
use super::candidates::{
materialize_local_same_format_provider_candidate_attempts,
build_local_same_format_provider_candidate_attempt_source,
resolve_local_same_format_provider_decision_input,
};
use super::payload::maybe_build_local_same_format_provider_decision_payload_for_candidate;
@@ -55,7 +55,7 @@ pub(crate) async fn maybe_build_sync_local_same_format_provider_decision_payload
Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete",
);
let (attempts, candidate_count) = materialize_local_same_format_provider_candidate_attempts(
let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
state, trace_id, &input, body_json, spec,
)
.await?;
@@ -65,7 +65,7 @@ pub(crate) async fn maybe_build_sync_local_same_format_provider_decision_payload
candidate_count,
);
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
if let Some(payload) =
maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
@@ -122,7 +122,7 @@ pub(crate) async fn maybe_build_stream_local_same_format_provider_decision_paylo
Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete",
);
let (attempts, candidate_count) = materialize_local_same_format_provider_candidate_attempts(
let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
state, trace_id, &input, body_json, spec,
)
.await?;
@@ -132,7 +132,7 @@ pub(crate) async fn maybe_build_stream_local_same_format_provider_decision_paylo
candidate_count,
);
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
if let Some(payload) =
maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,

View File

@@ -20,7 +20,6 @@ pub(crate) use crate::ai_serving::{
use super::{
build_local_same_format_provider_candidate_attempt_source,
materialize_local_same_format_provider_candidate_attempts,
maybe_build_local_same_format_provider_decision_payload_for_candidate,
resolve_local_same_format_provider_decision_input, AiStreamAttempt, AiSyncAttempt, AppState,
GatewayControlDecision, GatewayError, LocalSameFormatProviderCandidateAttempt,
@@ -348,7 +347,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete",
);
let (attempts, candidate_count) = materialize_local_same_format_provider_candidate_attempts(
let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
state, trace_id, &input, body_json, spec,
)
.await?;
@@ -362,7 +361,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
}
let mut plans = Vec::new();
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -433,7 +432,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete",
);
let (attempts, candidate_count) = materialize_local_same_format_provider_candidate_attempts(
let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
state, trace_id, &input, body_json, spec,
)
.await?;
@@ -447,7 +446,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
}
let mut plans = Vec::new();
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,161 @@
use aether_ai_serving::{
score_pool_member_with_rules, PoolMemberScoreInput, PoolMemberScoreRules, POOL_SCORE_VERSION,
};
use aether_data_contracts::repository::pool_scores::{
PoolMemberIdentity, PoolMemberProbeStatus, PoolScoreScope, UpsertPoolMemberScore,
POOL_SCORE_CAPABILITY_ACCOUNT, POOL_SCORE_SCOPE_KIND_ACCOUNT,
};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use serde_json::Value;
use crate::handlers::shared::{provider_key_health_summary, provider_key_status_snapshot_payload};
pub(crate) fn build_provider_key_pool_score_upsert(
key: &StoredProviderCatalogKey,
provider_type: &str,
existing: Option<&aether_data_contracts::repository::pool_scores::StoredPoolMemberScore>,
now_unix_secs: u64,
score_rules: PoolMemberScoreRules,
) -> UpsertPoolMemberScore {
let identity = PoolMemberIdentity::provider_api_key(key.provider_id.clone(), key.id.clone());
let scope = provider_key_pool_score_scope();
let input = provider_key_score_input(
key,
provider_type,
identity.clone(),
scope.clone(),
existing,
now_unix_secs,
);
let output = score_pool_member_with_rules(&input, score_rules);
UpsertPoolMemberScore {
id: provider_key_pool_score_id(&identity, &scope),
identity,
scope,
score: output.score,
hard_state: output.hard_state,
score_version: POOL_SCORE_VERSION,
score_reason: output.score_reason,
last_ranked_at: Some(now_unix_secs),
last_scheduled_at: existing.and_then(|score| score.last_scheduled_at),
last_success_at: existing.and_then(|score| score.last_success_at),
last_failure_at: existing.and_then(|score| score.last_failure_at),
failure_count: existing.map(|score| score.failure_count).unwrap_or(0),
last_probe_attempt_at: existing.and_then(|score| score.last_probe_attempt_at),
last_probe_success_at: existing.and_then(|score| score.last_probe_success_at),
last_probe_failure_at: existing.and_then(|score| score.last_probe_failure_at),
probe_failure_count: existing.map(|score| score.probe_failure_count).unwrap_or(0),
probe_status: existing
.map(|score| score.probe_status)
.unwrap_or(PoolMemberProbeStatus::Never),
updated_at: now_unix_secs,
}
}
pub(crate) fn provider_key_pool_score_scope() -> PoolScoreScope {
PoolScoreScope {
capability: POOL_SCORE_CAPABILITY_ACCOUNT.to_string(),
scope_kind: POOL_SCORE_SCOPE_KIND_ACCOUNT.to_string(),
scope_id: None,
}
}
pub(crate) fn provider_key_pool_score_id(
identity: &PoolMemberIdentity,
scope: &PoolScoreScope,
) -> String {
let raw = format!(
"{}:{}:{}:{}:{}:{}:{}",
identity.pool_kind,
identity.pool_id,
identity.member_kind,
identity.member_id,
scope.capability,
scope.scope_kind,
scope.scope_id.as_deref().unwrap_or("*")
);
format!(
"pms-{:016x}-{:016x}",
stable_hash(raw.as_bytes()),
stable_hash(identity.member_id.as_bytes())
)
}
fn provider_key_score_input(
key: &StoredProviderCatalogKey,
provider_type: &str,
identity: PoolMemberIdentity,
scope: PoolScoreScope,
existing: Option<&aether_data_contracts::repository::pool_scores::StoredPoolMemberScore>,
now_unix_secs: u64,
) -> PoolMemberScoreInput {
let status_snapshot = provider_key_status_snapshot_payload(key, provider_type);
let quota_snapshot = status_snapshot
.as_object()
.and_then(|snapshot| snapshot.get("quota"))
.and_then(Value::as_object);
let account_snapshot = status_snapshot
.as_object()
.and_then(|snapshot| snapshot.get("account"))
.and_then(Value::as_object);
let (health_score, _, _, any_circuit_open, _) = provider_key_health_summary(key);
let health_score = key
.health_by_format
.as_ref()
.and_then(Value::as_object)
.filter(|payload| !payload.is_empty())
.map(|_| health_score);
PoolMemberScoreInput {
identity,
scope: scope.clone(),
internal_priority: key.internal_priority,
is_active: key.is_active,
health_score,
quota_usage_ratio: quota_snapshot
.and_then(|quota| quota.get("usage_ratio"))
.and_then(json_f64)
.map(|value| value.clamp(0.0, 1.0)),
quota_exhausted: quota_snapshot
.and_then(|quota| quota.get("exhausted"))
.and_then(Value::as_bool)
.unwrap_or(false),
account_blocked: account_snapshot
.and_then(|account| account.get("blocked"))
.and_then(Value::as_bool)
.unwrap_or(false),
oauth_invalid_reason: key.oauth_invalid_reason.clone(),
circuit_open: any_circuit_open,
success_count: key.success_count.unwrap_or(0).into(),
error_count: key.error_count.unwrap_or(0).into(),
total_response_time_ms: key.total_response_time_ms.unwrap_or(0).into(),
total_tokens: key.total_tokens,
total_cost_usd: key.total_cost_usd,
last_used_at: key.last_used_at_unix_secs,
last_probe_success_at: existing.and_then(|score| score.last_probe_success_at),
probe_failure_count: existing.map(|score| score.probe_failure_count).unwrap_or(0),
probe_status: existing
.map(|score| score.probe_status)
.unwrap_or(PoolMemberProbeStatus::Never),
now_unix_secs,
}
}
fn json_f64(value: &Value) -> Option<f64> {
value.as_f64().or_else(|| {
value
.as_str()
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(|value| value.parse::<f64>().ok())
})
}
fn stable_hash(bytes: &[u8]) -> u64 {
let mut hash = 0xcbf29ce484222325u64;
for byte in bytes {
hash ^= u64::from(*byte);
hash = hash.wrapping_mul(0x100000001b3);
}
hash
}

View File

@@ -318,10 +318,10 @@ pub(crate) async fn maybe_build_sync_local_gemini_files_decision_payload(
return Ok(None);
};
let attempts =
materialize_local_gemini_files_candidate_attempts(state, trace_id, &input).await?;
let (mut source, _) =
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
if let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
state,
parts,
@@ -359,11 +359,11 @@ pub(crate) async fn maybe_build_stream_local_gemini_files_decision_payload(
return Ok(None);
};
let attempts =
materialize_local_gemini_files_candidate_attempts(state, trace_id, &input).await?;
let (mut source, _) =
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
let empty_body_json = serde_json::Value::Null;
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
if let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
state,
parts,
@@ -407,11 +407,11 @@ async fn build_local_sync_plan_and_reports(
return Ok(Vec::new());
};
let attempts =
materialize_local_gemini_files_candidate_attempts(state, trace_id, &input).await?;
let (mut source, _) =
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
let mut plans = Vec::new();
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
state,
parts,
@@ -459,12 +459,12 @@ async fn build_local_stream_plan_and_reports(
return Ok(Vec::new());
};
let attempts =
materialize_local_gemini_files_candidate_attempts(state, trace_id, &input).await?;
let (mut source, _) =
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
let mut plans = Vec::new();
let empty_body_json = serde_json::Value::Null;
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
state,
parts,

View File

@@ -405,7 +405,7 @@ pub(crate) async fn maybe_build_sync_local_image_decision_payload(
return Ok(None);
};
let Some(attempts) = list_local_openai_image_candidate_attempts(
let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source(
state,
trace_id,
&input,
@@ -413,12 +413,12 @@ pub(crate) async fn maybe_build_sync_local_image_decision_payload(
spec_metadata.api_format,
spec_metadata.decision_kind,
)
.await
.await?
else {
return Ok(None);
};
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
if let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
state,
parts,
@@ -465,7 +465,7 @@ pub(crate) async fn maybe_build_stream_local_image_decision_payload(
return Ok(None);
};
let Some(attempts) = list_local_openai_image_candidate_attempts(
let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source(
state,
trace_id,
&input,
@@ -473,12 +473,12 @@ pub(crate) async fn maybe_build_stream_local_image_decision_payload(
spec_metadata.api_format,
spec_metadata.decision_kind,
)
.await
.await?
else {
return Ok(None);
};
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
if let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
state,
parts,
@@ -521,7 +521,7 @@ async fn build_local_sync_plan_and_reports(
return Ok(Vec::new());
};
let Some(attempts) = list_local_openai_image_candidate_attempts(
let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source(
state,
trace_id,
&input,
@@ -529,13 +529,13 @@ async fn build_local_sync_plan_and_reports(
spec_metadata.api_format,
spec_metadata.decision_kind,
)
.await
.await?
else {
return Ok(Vec::new());
};
let mut plans = Vec::new();
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
state,
parts,
@@ -597,7 +597,7 @@ async fn build_local_stream_plan_and_reports(
return Ok(Vec::new());
};
let Some(attempts) = list_local_openai_image_candidate_attempts(
let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source(
state,
trace_id,
&input,
@@ -605,13 +605,13 @@ async fn build_local_stream_plan_and_reports(
spec_metadata.api_format,
spec_metadata.decision_kind,
)
.await
.await?
else {
return Ok(Vec::new());
};
let mut plans = Vec::new();
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
state,
parts,

View File

@@ -180,7 +180,7 @@ pub(crate) async fn maybe_build_sync_local_video_decision_payload(
return Ok(None);
};
let Some(attempts) = list_local_video_create_candidate_attempts(
let Some((mut source, _)) = build_local_video_create_candidate_attempt_source(
state,
trace_id,
&input,
@@ -188,12 +188,12 @@ pub(crate) async fn maybe_build_sync_local_video_decision_payload(
spec_metadata.api_format,
spec_metadata.decision_kind,
)
.await
.await?
else {
return Ok(None);
};
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
if let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate(
state, parts, body_json, trace_id, &input, attempt, spec,
)
@@ -223,7 +223,7 @@ async fn build_local_sync_plan_and_reports(
return Ok(Vec::new());
};
let Some(attempts) = list_local_video_create_candidate_attempts(
let Some((mut source, _)) = build_local_video_create_candidate_attempt_source(
state,
trace_id,
&input,
@@ -231,13 +231,13 @@ async fn build_local_sync_plan_and_reports(
spec_metadata.api_format,
spec_metadata.decision_kind,
)
.await
.await?
else {
return Ok(Vec::new());
};
let mut plans = Vec::new();
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate(
state, parts, body_json, trace_id, &input, attempt, spec,
)

View File

@@ -20,8 +20,7 @@ use crate::ai_serving::GatewayControlDecision;
use crate::{AiExecutionDecision, AppState, GatewayError};
use super::candidates::{
build_local_standard_candidate_attempt_source, materialize_local_standard_candidate_attempts,
resolve_local_standard_decision_input,
build_local_standard_candidate_attempt_source, resolve_local_standard_decision_input,
};
use super::payload::maybe_build_local_standard_decision_payload_for_candidate;
use super::{LocalStandardDecisionInput, LocalStandardSpec};
@@ -323,12 +322,12 @@ pub(crate) async fn maybe_build_sync_via_standard_family_payload(
Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete",
);
let (attempts, candidate_count) =
materialize_local_standard_candidate_attempts(state, trace_id, &input, body_json, spec)
let (mut source, candidate_count) =
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
.await?;
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
if let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -372,12 +371,12 @@ pub(crate) async fn maybe_build_stream_via_standard_family_payload(
Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete",
);
let (attempts, candidate_count) =
materialize_local_standard_candidate_attempts(state, trace_id, &input, body_json, spec)
let (mut source, candidate_count) =
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
.await?;
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
if let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -427,15 +426,15 @@ pub(crate) async fn build_local_sync_plan_and_reports(
Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete",
);
let (attempts, candidate_count) =
materialize_local_standard_candidate_attempts(state, trace_id, &input, body_json, spec)
let (mut source, candidate_count) =
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
.await?;
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
if candidate_count == 0 {
return Ok(Vec::new());
}
let mut plans = Vec::new();
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -501,15 +500,15 @@ pub(crate) async fn build_local_stream_plan_and_reports(
Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete",
);
let (attempts, candidate_count) =
materialize_local_standard_candidate_attempts(state, trace_id, &input, body_json, spec)
let (mut source, candidate_count) =
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
.await?;
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
if candidate_count == 0 {
return Ok(Vec::new());
}
let mut plans = Vec::new();
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)

View File

@@ -1,28 +1,23 @@
use serde_json::Value;
use tracing::warn;
use crate::ai_serving::planner::common::{
OPENAI_CHAT_STREAM_PLAN_KIND, OPENAI_CHAT_SYNC_PLAN_KIND,
};
use crate::ai_serving::planner::runtime_miss::set_local_runtime_execution_exhausted_diagnostic;
use crate::ai_serving::GatewayControlDecision;
use crate::{AiExecutionDecision, AppState, GatewayError};
use tracing::warn;
mod decision;
mod plans;
use self::decision::{
build_lazy_local_openai_chat_candidate_attempt_source,
build_local_openai_chat_candidate_attempt_source,
materialize_local_openai_chat_candidate_attempts,
maybe_build_local_openai_chat_decision_payload_for_candidate, LocalOpenAiChatCandidateAttempt,
LocalOpenAiChatCandidateAttemptSource, LocalOpenAiChatDecisionInput,
};
use self::plans::{
build_local_openai_chat_stream_attempt_source, build_local_openai_chat_stream_plan_and_reports,
build_local_openai_chat_sync_attempt_source, build_local_openai_chat_sync_plan_and_reports,
list_local_openai_chat_candidates, resolve_local_openai_chat_decision_input,
set_local_openai_chat_miss_diagnostic,
resolve_local_openai_chat_decision_input,
};
pub(crate) async fn build_local_openai_chat_sync_plan_and_reports_for_kind(
@@ -146,32 +141,17 @@ pub(crate) async fn maybe_build_sync_local_decision_payload(
return Ok(None);
};
let (candidates, skipped_candidates) =
match list_local_openai_chat_candidates(state, &input, false).await {
Ok(value) => value,
Err(err) => {
warn!(
event_name = "local_openai_chat_scheduler_selection_failed",
log_type = "event",
trace_id = %trace_id,
error = ?err,
"gateway local openai chat sync decision scheduler selection failed"
);
return Ok(None);
}
};
let attempts = materialize_local_openai_chat_candidate_attempts(
state,
trace_id,
&input,
body_json,
candidates,
skipped_candidates,
let (mut source, _) = build_lazy_local_openai_chat_candidate_attempt_source(
state, trace_id, &input, body_json, false,
)
.await;
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
let upstream_is_stream = self::plans::openai_chat_upstream_is_stream_for_candidate(
&attempt.eligible.transport,
attempt.eligible.provider_api_format.as_str(),
false,
);
if let Some(payload) = maybe_build_local_openai_chat_decision_payload_for_candidate(
state,
parts,
@@ -181,7 +161,7 @@ pub(crate) async fn maybe_build_sync_local_decision_payload(
attempt,
OPENAI_CHAT_SYNC_PLAN_KIND,
"openai_chat_sync_success",
false,
upstream_is_stream,
)
.await
{
@@ -212,32 +192,17 @@ pub(crate) async fn maybe_build_stream_local_decision_payload(
return Ok(None);
};
let (candidates, skipped_candidates) =
match list_local_openai_chat_candidates(state, &input, true).await {
Ok(value) => value,
Err(err) => {
warn!(
event_name = "local_openai_chat_scheduler_selection_failed",
log_type = "event",
trace_id = %trace_id,
error = ?err,
"gateway local openai chat stream decision scheduler selection failed"
);
return Ok(None);
}
};
let attempts = materialize_local_openai_chat_candidate_attempts(
state,
trace_id,
&input,
body_json,
candidates,
skipped_candidates,
let (mut source, _) = build_lazy_local_openai_chat_candidate_attempt_source(
state, trace_id, &input, body_json, true,
)
.await;
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
let upstream_is_stream = self::plans::openai_chat_upstream_is_stream_for_candidate(
&attempt.eligible.transport,
attempt.eligible.provider_api_format.as_str(),
true,
);
if let Some(payload) = maybe_build_local_openai_chat_decision_payload_for_candidate(
state,
parts,
@@ -247,7 +212,7 @@ pub(crate) async fn maybe_build_stream_local_decision_payload(
attempt,
OPENAI_CHAT_STREAM_PLAN_KIND,
"openai_chat_stream_success",
true,
upstream_is_stream,
)
.await
{

View File

@@ -22,7 +22,7 @@ pub(super) use self::sync::{
build_local_openai_chat_sync_attempt_source, build_local_openai_chat_sync_plan_and_reports,
};
fn openai_chat_upstream_is_stream_for_candidate(
pub(super) fn openai_chat_upstream_is_stream_for_candidate(
transport: &GatewayProviderTransportSnapshot,
provider_api_format: &str,
client_is_stream: bool,

View File

@@ -3,13 +3,10 @@ use tracing::warn;
use super::super::{
build_lazy_local_openai_chat_candidate_attempt_source,
build_local_openai_chat_candidate_attempt_source,
materialize_local_openai_chat_candidate_attempts,
maybe_build_local_openai_chat_decision_payload_for_candidate, AppState, GatewayControlDecision,
GatewayError, LocalOpenAiChatCandidateAttempt, LocalOpenAiChatCandidateAttemptSource,
LocalOpenAiChatDecisionInput,
};
use super::candidates::list_local_openai_chat_candidates;
use super::diagnostic::{
set_local_openai_chat_candidate_evaluation_diagnostic, set_local_openai_chat_miss_diagnostic,
};
@@ -176,27 +173,12 @@ pub(crate) async fn build_local_openai_chat_stream_plan_and_reports(
return Ok(Vec::new());
};
let (candidates, skipped_candidates) =
match list_local_openai_chat_candidates(state, &input, true).await {
Ok(value) => value,
Err(err) => {
warn!(
trace_id = %trace_id,
error = ?err,
"gateway local openai chat stream decision scheduler selection failed"
);
set_local_openai_chat_miss_diagnostic(
state,
trace_id,
decision,
plan_kind,
Some(input.requested_model.as_str()),
"scheduler_selection_failed",
);
return Ok(Vec::new());
}
};
if candidates.is_empty() && skipped_candidates.is_empty() {
let Some((mut attempt_source, candidate_count)) =
build_local_openai_chat_stream_attempt_source(
state, parts, trace_id, decision, body_json, plan_kind,
)
.await?
else {
set_local_openai_chat_candidate_evaluation_diagnostic(
state,
trace_id,
@@ -206,59 +188,13 @@ pub(crate) async fn build_local_openai_chat_stream_plan_and_reports(
0,
);
return Ok(Vec::new());
}
set_local_openai_chat_candidate_evaluation_diagnostic(
state,
trace_id,
decision,
plan_kind,
Some(input.requested_model.as_str()),
candidates.len() + skipped_candidates.len(),
);
let attempts = materialize_local_openai_chat_candidate_attempts(
state,
trace_id,
&input,
body_json,
candidates,
skipped_candidates,
)
.await;
};
let mut plans = Vec::new();
for attempt in attempts {
let upstream_is_stream = openai_chat_upstream_is_stream_for_candidate(
&attempt.eligible.transport,
attempt.eligible.provider_api_format.as_str(),
true,
);
let Some(payload) = maybe_build_local_openai_chat_decision_payload_for_candidate(
state,
parts,
trace_id,
body_json,
&input,
attempt,
OPENAI_CHAT_STREAM_PLAN_KIND,
"openai_chat_stream_success",
upstream_is_stream,
)
.await
else {
continue;
};
match build_openai_chat_stream_plan_from_decision(parts, body_json, payload) {
Ok(Some(value)) => plans.push(value),
Ok(None) => {}
Err(err) => {
warn!(
trace_id = %trace_id,
error = ?err,
"gateway local openai chat stream decision plan build failed"
);
}
while let Some(attempt) = attempt_source.next_execution_attempt().await? {
plans.push(attempt);
if plans.len() >= candidate_count {
break;
}
}

View File

@@ -3,13 +3,10 @@ use tracing::warn;
use super::super::{
build_lazy_local_openai_chat_candidate_attempt_source,
build_local_openai_chat_candidate_attempt_source,
materialize_local_openai_chat_candidate_attempts,
maybe_build_local_openai_chat_decision_payload_for_candidate, AppState, GatewayControlDecision,
GatewayError, LocalOpenAiChatCandidateAttempt, LocalOpenAiChatCandidateAttemptSource,
LocalOpenAiChatDecisionInput,
};
use super::candidates::list_local_openai_chat_candidates;
use super::diagnostic::{
set_local_openai_chat_candidate_evaluation_diagnostic, set_local_openai_chat_miss_diagnostic,
};
@@ -176,27 +173,11 @@ pub(crate) async fn build_local_openai_chat_sync_plan_and_reports(
return Ok(Vec::new());
};
let (candidates, skipped_candidates) =
match list_local_openai_chat_candidates(state, &input, false).await {
Ok(value) => value,
Err(err) => {
warn!(
trace_id = %trace_id,
error = ?err,
"gateway local openai chat sync decision scheduler selection failed"
);
set_local_openai_chat_miss_diagnostic(
state,
trace_id,
decision,
plan_kind,
Some(input.requested_model.as_str()),
"scheduler_selection_failed",
);
return Ok(Vec::new());
}
};
if candidates.is_empty() && skipped_candidates.is_empty() {
let Some((mut attempt_source, candidate_count)) = build_local_openai_chat_sync_attempt_source(
state, parts, trace_id, decision, body_json, plan_kind,
)
.await?
else {
set_local_openai_chat_candidate_evaluation_diagnostic(
state,
trace_id,
@@ -206,59 +187,13 @@ pub(crate) async fn build_local_openai_chat_sync_plan_and_reports(
0,
);
return Ok(Vec::new());
}
set_local_openai_chat_candidate_evaluation_diagnostic(
state,
trace_id,
decision,
plan_kind,
Some(input.requested_model.as_str()),
candidates.len() + skipped_candidates.len(),
);
let attempts = materialize_local_openai_chat_candidate_attempts(
state,
trace_id,
&input,
body_json,
candidates,
skipped_candidates,
)
.await;
};
let mut plans = Vec::new();
for attempt in attempts {
let upstream_is_stream = openai_chat_upstream_is_stream_for_candidate(
&attempt.eligible.transport,
attempt.eligible.provider_api_format.as_str(),
false,
);
let Some(payload) = maybe_build_local_openai_chat_decision_payload_for_candidate(
state,
parts,
trace_id,
body_json,
&input,
attempt,
OPENAI_CHAT_SYNC_PLAN_KIND,
"openai_chat_sync_success",
upstream_is_stream,
)
.await
else {
continue;
};
match build_openai_chat_sync_plan_from_decision(parts, body_json, payload) {
Ok(Some(value)) => plans.push(value),
Ok(None) => {}
Err(err) => {
warn!(
trace_id = %trace_id,
error = ?err,
"gateway local openai chat sync decision plan build failed"
);
}
while let Some(attempt) = attempt_source.next_execution_attempt().await? {
plans.push(attempt);
if plans.len() >= candidate_count {
break;
}
}

View File

@@ -6,7 +6,7 @@ mod decision;
mod plans;
use self::decision::{
materialize_local_openai_responses_candidate_attempts,
build_local_openai_responses_candidate_attempt_source,
maybe_build_local_openai_responses_decision_payload_for_candidate,
resolve_local_openai_responses_decision_input,
};
@@ -108,12 +108,12 @@ pub(crate) async fn maybe_build_sync_local_openai_responses_decision_payload(
return Ok(None);
};
let (attempts, _) = materialize_local_openai_responses_candidate_attempts(
let (mut source, _) = build_local_openai_responses_candidate_attempt_source(
state, trace_id, &input, body_json, spec,
)
.await?;
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
if let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -146,12 +146,12 @@ pub(crate) async fn maybe_build_stream_local_openai_responses_decision_payload(
return Ok(None);
};
let (attempts, _) = materialize_local_openai_responses_candidate_attempts(
let (mut source, _) = build_local_openai_responses_candidate_attempt_source(
state, trace_id, &input, body_json, spec,
)
.await?;
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
if let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)

View File

@@ -3,7 +3,6 @@ use tracing::warn;
use super::decision::{
build_local_openai_responses_candidate_attempt_source,
materialize_local_openai_responses_candidate_attempts,
maybe_build_local_openai_responses_decision_payload_for_candidate,
resolve_local_openai_responses_decision_input, LocalOpenAiResponsesCandidateAttempt,
LocalOpenAiResponsesCandidateAttemptSource, LocalOpenAiResponsesDecisionInput,
@@ -312,7 +311,7 @@ pub(super) async fn build_local_sync_plan_and_reports(
"candidate_evaluation_incomplete",
);
let (attempts, candidate_count) = materialize_local_openai_responses_candidate_attempts(
let (mut source, candidate_count) = build_local_openai_responses_candidate_attempt_source(
state, trace_id, &input, body_json, spec,
)
.await?;
@@ -322,7 +321,7 @@ pub(super) async fn build_local_sync_plan_and_reports(
}
let mut plans = Vec::new();
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)
@@ -384,7 +383,7 @@ pub(super) async fn build_local_stream_plan_and_reports(
"candidate_evaluation_incomplete",
);
let (attempts, candidate_count) = materialize_local_openai_responses_candidate_attempts(
let (mut source, candidate_count) = build_local_openai_responses_candidate_attempt_source(
state, trace_id, &input, body_json, spec,
)
.await?;
@@ -394,7 +393,7 @@ pub(super) async fn build_local_stream_plan_and_reports(
}
let mut plans = Vec::new();
for attempt in attempts {
while let Some(attempt) = source.next_attempt().await {
let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec,
)

View File

@@ -2,9 +2,10 @@ pub(crate) use aether_ai_formats::api::{
aggregate_claude_stream_sync_response, aggregate_gemini_stream_sync_response,
aggregate_openai_chat_stream_sync_response, aggregate_openai_responses_stream_sync_response,
aggregate_standard_chat_stream_sync_response, aggregate_standard_cli_stream_sync_response,
api_format_alias_matches, apply_codex_openai_responses_special_body_edits,
apply_codex_openai_responses_special_headers, apply_model_directive_mapping_patch,
apply_model_directive_overrides_from_model, apply_model_directive_overrides_from_request,
api_format_alias_matches, api_format_storage_aliases,
apply_codex_openai_responses_special_body_edits, apply_codex_openai_responses_special_headers,
apply_model_directive_mapping_patch, apply_model_directive_overrides_from_model,
apply_model_directive_overrides_from_request,
apply_openai_responses_compact_special_body_edits, build_chatgpt_web_image_request_body,
build_core_error_body_for_client_format, build_cross_format_openai_chat_request_body,
build_cross_format_openai_chat_request_body_with_model_directives,

View File

@@ -28,4 +28,8 @@ impl AuthContextCache {
self.entries
.insert(cache_key, auth_context, ttl, max_entries);
}
pub(crate) fn clear(&self) {
self.entries.clear();
}
}

View File

@@ -339,9 +339,12 @@ pub(crate) fn validate_management_token_admin_route_permission(
.and_then(|value| value.rsplit_once(':').map(|(scope, _)| scope))
.unwrap_or_default();
let admin_permission = format!("admin:{scope}:admin");
let has_full_assignable_access =
management_token_permissions_cover_all_assignable_permissions(token_permissions);
if token_permissions
.iter()
.any(|permission| permission == &required_permission || permission == &admin_permission)
|| (scope == "management_tokens" && has_full_assignable_access)
{
Ok(())
} else {
@@ -468,6 +471,18 @@ fn is_assignable_management_token_permission(key: &str) -> bool {
.any(|item| item.key == key)
}
pub(crate) fn management_token_permissions_cover_all_assignable_permissions(
token_permissions: &[String],
) -> bool {
let permission_set = token_permissions
.iter()
.map(String::as_str)
.collect::<BTreeSet<_>>();
all_assignable_management_token_permissions()
.iter()
.all(|permission| permission_set.contains(permission.as_str()))
}
#[cfg(test)]
mod tests {
use super::*;
@@ -581,6 +596,25 @@ mod tests {
);
}
#[test]
fn full_assignable_token_permissions_can_cover_management_tokens_scope() {
let decision = GatewayControlDecision::synthetic(
"/api/admin/management-tokens".to_string(),
Some("admin_proxy".to_string()),
Some("management_tokens_manage".to_string()),
Some("list_tokens".to_string()),
Some("admin:management_tokens".to_string()),
);
let permissions = all_assignable_management_token_permissions();
assert!(validate_management_token_admin_route_permission(
&http::Method::GET,
&decision,
Some(&permissions),
)
.is_ok());
}
fn extract_admin_route_scopes(source: &'static str) -> BTreeSet<&'static str> {
let mut scopes = BTreeSet::new();
let mut remaining = source;

View File

@@ -16,6 +16,7 @@ pub(crate) use execute::{allows_control_execute_emergency, maybe_execute_via_con
pub(crate) use management_token_permissions::{
all_assignable_management_token_permissions, management_token_permission_catalog_payload,
management_token_permission_keys_from_value, management_token_permission_mode_and_summary,
management_token_permissions_cover_all_assignable_permissions,
management_token_required_permission, normalize_assignable_management_token_permissions,
validate_management_token_admin_route_permission,
};

View File

@@ -158,9 +158,21 @@ pub(super) fn classify_admin_observability_family_route(
"admin:api_keys",
false,
))
} else if method == http::Method::POST
&& normalized_path_no_trailing.starts_with("/api/admin/api-keys/")
&& normalized_path_no_trailing.ends_with("/install-sessions")
&& normalized_path_no_trailing.matches('/').count() == 5
{
Some(classified(
"admin_proxy",
"api_keys_manage",
"create_api_key_install_session",
"admin:api_keys",
false,
))
} else if method == http::Method::GET
&& normalized_path.starts_with("/api/admin/api-keys/")
&& normalized_path.matches('/').count() == 4
&& normalized_path_no_trailing.starts_with("/api/admin/api-keys/")
&& normalized_path_no_trailing.matches('/').count() == 4
{
Some(classified(
"admin_proxy",
@@ -170,8 +182,8 @@ pub(super) fn classify_admin_observability_family_route(
false,
))
} else if method == http::Method::PUT
&& normalized_path.starts_with("/api/admin/api-keys/")
&& normalized_path.matches('/').count() == 4
&& normalized_path_no_trailing.starts_with("/api/admin/api-keys/")
&& normalized_path_no_trailing.matches('/').count() == 4
{
Some(classified(
"admin_proxy",
@@ -240,6 +252,18 @@ pub(super) fn classify_admin_observability_family_route(
"admin:pool",
false,
))
} else if method == http::Method::GET
&& normalized_path_no_trailing.starts_with("/api/admin/pool/")
&& normalized_path_no_trailing.ends_with("/scores")
&& normalized_path_no_trailing.matches('/').count() == 5
{
Some(classified(
"admin_proxy",
"pool_manage",
"scores",
"admin:pool",
false,
))
} else if method == http::Method::POST
&& normalized_path_no_trailing.starts_with("/api/admin/pool/")
&& normalized_path_no_trailing.ends_with("/keys/batch-import")

View File

@@ -1,5 +1,8 @@
use http::Uri;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::shared::local_proxy_route_requires_buffered_body;
use super::{classify_control_route, headers};
#[test]
@@ -38,6 +41,50 @@ fn classifies_admin_api_keys_create_as_admin_proxy_route() {
assert!(!decision.is_execution_runtime_candidate());
}
#[test]
fn classifies_admin_api_key_install_session_create_as_admin_proxy_route() {
let headers = headers(&[]);
for path in [
"/api/admin/api-keys/key-123/install-sessions",
"/api/admin/api-keys/key-123/install-sessions/",
] {
let uri: Uri = path.parse().expect("uri should parse");
let decision = classify_control_route(&http::Method::POST, &uri, &headers)
.expect("route should classify");
assert_eq!(decision.route_class.as_deref(), Some("admin_proxy"));
assert_eq!(decision.route_family.as_deref(), Some("api_keys_manage"));
assert_eq!(
decision.route_kind.as_deref(),
Some("create_api_key_install_session")
);
assert_eq!(
decision.auth_endpoint_signature.as_deref(),
Some("admin:api_keys")
);
assert!(!decision.is_execution_runtime_candidate());
}
}
#[test]
fn admin_api_key_install_session_create_buffers_request_body() {
let headers = headers(&[]);
let uri: Uri = "/api/admin/api-keys/key-123/install-sessions"
.parse()
.expect("uri should parse");
let decision =
classify_control_route(&http::Method::POST, &uri, &headers).expect("route should classify");
let context = GatewayPublicRequestContext::from_request_parts(
"trace-admin-install-session",
&http::Method::POST,
&uri,
&headers,
Some(decision),
);
assert!(local_proxy_route_requires_buffered_body(&context));
}
#[test]
fn classifies_admin_api_keys_detail_as_admin_proxy_route() {
let headers = headers(&[]);

View File

@@ -52,6 +52,14 @@ fn classifies_admin_pool_provider_key_routes_as_admin_proxy_route() {
assert_eq!(list.route_family.as_deref(), Some("pool_manage"));
assert_eq!(list.route_kind.as_deref(), Some("list_keys"));
let scores_uri: Uri = "/api/admin/pool/provider-1/scores?api_format=openai:responses"
.parse()
.expect("uri should parse");
let scores = classify_control_route(&http::Method::GET, &scores_uri, &headers)
.expect("route should classify");
assert_eq!(scores.route_family.as_deref(), Some("pool_manage"));
assert_eq!(scores.route_kind.as_deref(), Some("scores"));
let batch_import_uri: Uri = "/api/admin/pool/provider-1/keys/batch-import"
.parse()
.expect("uri should parse");

View File

@@ -45,6 +45,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -86,6 +88,8 @@ impl GatewayDataState {
let gemini_file_mapping_writer = backends.write().gemini_file_mappings();
let provider_catalog_reader = backends.read().provider_catalog();
let provider_catalog_writer = backends.write().provider_catalog();
let pool_score_reader = backends.read().pool_scores();
let pool_score_writer = backends.write().pool_scores();
let provider_quota_reader = backends.read().provider_quotas();
let provider_quota_writer = backends.write().provider_quotas();
let usage_reader = backends.read().usage();
@@ -125,6 +129,8 @@ impl GatewayDataState {
request_candidate_writer,
provider_catalog_reader,
provider_catalog_writer,
pool_score_reader,
pool_score_writer,
provider_quota_reader,
provider_quota_writer,
usage_reader,
@@ -263,6 +269,14 @@ impl GatewayDataState {
self.provider_catalog_writer.is_some()
}
pub(crate) fn has_pool_score_reader(&self) -> bool {
self.pool_score_reader.is_some()
}
pub(crate) fn has_pool_score_writer(&self) -> bool {
self.pool_score_writer.is_some()
}
pub(crate) fn has_proxy_node_reader(&self) -> bool {
self.proxy_node_reader.is_some()
}

View File

@@ -95,7 +95,8 @@ use aether_data_contracts::repository::billing::{
};
use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredRequestedModelCandidateRowsQuery,
};
use aether_data_contracts::repository::candidates::{
PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateReadRepository,
@@ -109,6 +110,13 @@ use aether_data_contracts::repository::global_models::{
StoredProviderModelStats, StoredPublicCatalogModel, StoredPublicGlobalModel,
StoredPublicGlobalModelPage, UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
};
use aether_data_contracts::repository::pool_scores::{
GetPoolMemberScoresByIdsQuery, ListPoolMemberProbeCandidatesQuery, ListPoolMemberScoresQuery,
ListRankedPoolMembersQuery, PoolMemberHardState, PoolMemberIdentity, PoolMemberProbeAttempt,
PoolMemberProbeResult, PoolMemberProbeStatus, PoolMemberScheduleFeedback,
PoolMemberScoreWriteRepository, PoolScoreReadRepository, PoolScoreScope, StoredPoolMemberScore,
UpsertPoolMemberScore,
};
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyListQuery, ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyPage,
@@ -158,6 +166,8 @@ pub(crate) struct GatewayDataState {
request_candidate_writer: Option<Arc<dyn RequestCandidateWriteRepository>>,
provider_catalog_reader: Option<Arc<dyn ProviderCatalogReadRepository>>,
provider_catalog_writer: Option<Arc<dyn ProviderCatalogWriteRepository>>,
pool_score_reader: Option<Arc<dyn PoolScoreReadRepository>>,
pool_score_writer: Option<Arc<dyn PoolMemberScoreWriteRepository>>,
provider_quota_reader: Option<Arc<dyn ProviderQuotaReadRepository>>,
provider_quota_writer: Option<Arc<dyn ProviderQuotaWriteRepository>>,
usage_reader: Option<Arc<dyn UsageReadRepository>>,
@@ -259,6 +269,8 @@ impl fmt::Debug for GatewayDataState {
"has_provider_catalog_writer",
&self.provider_catalog_writer.is_some(),
)
.field("has_pool_score_reader", &self.pool_score_reader.is_some())
.field("has_pool_score_writer", &self.pool_score_writer.is_some())
.field(
"has_provider_quota_reader",
&self.provider_quota_reader.is_some(),
@@ -289,6 +301,7 @@ mod catalog;
mod core;
mod integrations;
mod models;
mod pool_scores;
mod runtime;
#[cfg(test)]
mod testing;

View File

@@ -2,7 +2,8 @@ use super::{
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
DataLayerError, GatewayDataState, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery,
PublicGlobalModelQuery, StoredAdminGlobalModel, StoredAdminGlobalModelPage,
StoredAdminProviderModel, StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
StoredAdminProviderModel, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredProviderActiveGlobalModel, StoredProviderModelStats, StoredPublicCatalogModel,
StoredPublicGlobalModel, StoredPublicGlobalModelPage, StoredRequestedModelCandidateRowsQuery,
UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
@@ -73,6 +74,16 @@ impl GatewayDataState {
}
}
pub(crate) async fn list_pool_key_candidate_rows_for_group_key_ids(
&self,
query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
match &self.minimal_candidate_selection_reader {
Some(repository) => repository.list_pool_key_rows_for_group_key_ids(query).await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn list_public_global_models(
&self,
query: &PublicGlobalModelQuery,

View File

@@ -0,0 +1,123 @@
use super::{
DataLayerError, GatewayDataState, GetPoolMemberScoresByIdsQuery,
ListPoolMemberProbeCandidatesQuery, ListPoolMemberScoresQuery, ListRankedPoolMembersQuery,
PoolMemberHardState, PoolMemberIdentity, PoolMemberProbeAttempt, PoolMemberProbeResult,
PoolMemberScheduleFeedback, PoolScoreScope, StoredPoolMemberScore, UpsertPoolMemberScore,
};
impl GatewayDataState {
pub(crate) async fn list_ranked_pool_members(
&self,
query: &ListRankedPoolMembersQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
match &self.pool_score_reader {
Some(repository) => repository.list_ranked_pool_members(query).await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn list_pool_member_probe_candidates(
&self,
query: &ListPoolMemberProbeCandidatesQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
match &self.pool_score_reader {
Some(repository) => repository.list_pool_member_probe_candidates(query).await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn list_pool_member_scores(
&self,
query: &ListPoolMemberScoresQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
match &self.pool_score_reader {
Some(repository) => repository.list_pool_member_scores(query).await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn get_pool_member_scores_by_ids(
&self,
query: &GetPoolMemberScoresByIdsQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
match &self.pool_score_reader {
Some(repository) => repository.get_pool_member_scores_by_ids(query).await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn upsert_pool_member_score(
&self,
score: UpsertPoolMemberScore,
) -> Result<Option<StoredPoolMemberScore>, DataLayerError> {
match &self.pool_score_writer {
Some(repository) => repository.upsert_pool_member_score(score).await.map(Some),
None => Ok(None),
}
}
pub(crate) async fn record_pool_member_probe_result(
&self,
result: PoolMemberProbeResult,
) -> Result<usize, DataLayerError> {
match &self.pool_score_writer {
Some(repository) => repository.record_pool_member_probe_result(result).await,
None => Ok(0),
}
}
pub(crate) async fn mark_pool_member_probe_in_progress(
&self,
attempt: PoolMemberProbeAttempt,
) -> Result<usize, DataLayerError> {
match &self.pool_score_writer {
Some(repository) => repository.mark_pool_member_probe_in_progress(attempt).await,
None => Ok(0),
}
}
pub(crate) async fn record_pool_member_schedule_feedback(
&self,
feedback: PoolMemberScheduleFeedback,
) -> Result<usize, DataLayerError> {
match &self.pool_score_writer {
Some(repository) => {
repository
.record_pool_member_schedule_feedback(feedback)
.await
}
None => Ok(0),
}
}
pub(crate) async fn mark_pool_member_hard_state(
&self,
identity: &PoolMemberIdentity,
scope: Option<&PoolScoreScope>,
hard_state: PoolMemberHardState,
updated_at: u64,
) -> Result<usize, DataLayerError> {
match &self.pool_score_writer {
Some(repository) => {
repository
.mark_pool_member_hard_state(identity, scope, hard_state, updated_at)
.await
}
None => Ok(0),
}
}
pub(crate) async fn delete_pool_member_scores_for_member(
&self,
identity: &PoolMemberIdentity,
) -> Result<usize, DataLayerError> {
match &self.pool_score_writer {
Some(repository) => {
repository
.delete_pool_member_scores_for_member(identity)
.await
}
None => Ok(0),
}
}
}

View File

@@ -34,6 +34,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -86,6 +88,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,

View File

@@ -2,6 +2,7 @@ use std::collections::BTreeMap;
use std::sync::{Arc, RwLock};
use aether_data_contracts::repository::candidates::RequestCandidateRepository;
use aether_data_contracts::repository::pool_scores::PoolMemberScoreRepository;
use aether_data_contracts::repository::quota::ProviderQuotaRepository;
use aether_data_contracts::repository::usage::UsageRepository;
@@ -12,12 +13,13 @@ use super::{
GeminiFileMappingWriteRepository, GlobalModelReadRepository, GlobalModelWriteRepository,
ManagementTokenReadRepository, ManagementTokenWriteRepository,
MinimalCandidateSelectionReadRepository, OAuthProviderReadRepository,
OAuthProviderWriteRepository, ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, ProxyNodeReadRepository,
ProxyNodeWriteRepository, RequestCandidateReadRepository, RequestCandidateWriteRepository,
SettlementWriteRepository, StoredSystemConfigEntry, StoredUserPreferenceRecord,
UsageReadRepository, UsageWriteRepository, UserReadRepository, VideoTaskReadRepository,
VideoTaskWriteRepository, WalletReadRepository, WalletWriteRepository,
OAuthProviderWriteRepository, PoolMemberScoreWriteRepository, PoolScoreReadRepository,
ProviderCatalogReadRepository, ProviderCatalogWriteRepository, ProviderQuotaReadRepository,
ProviderQuotaWriteRepository, ProxyNodeReadRepository, ProxyNodeWriteRepository,
RequestCandidateReadRepository, RequestCandidateWriteRepository, SettlementWriteRepository,
StoredSystemConfigEntry, StoredUserPreferenceRecord, UsageReadRepository, UsageWriteRepository,
UserReadRepository, VideoTaskReadRepository, VideoTaskWriteRepository, WalletReadRepository,
WalletWriteRepository,
};
mod announcements;
@@ -67,6 +69,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -118,6 +122,8 @@ impl GatewayDataState {
request_candidate_writer: Some(request_candidate_writer),
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -165,6 +171,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: Some(repository),
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -300,6 +308,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: Some(provider_catalog_reader),
provider_catalog_writer: Some(provider_catalog_writer),
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -333,6 +343,18 @@ impl GatewayDataState {
self
}
#[cfg(test)]
pub(crate) fn with_pool_score_repository_for_tests<T>(mut self, repository: Arc<T>) -> Self
where
T: PoolMemberScoreRepository + 'static,
{
let pool_score_reader: Arc<dyn PoolScoreReadRepository> = repository.clone();
let pool_score_writer: Arc<dyn PoolMemberScoreWriteRepository> = repository;
self.pool_score_reader = Some(pool_score_reader);
self.pool_score_writer = Some(pool_score_writer);
self
}
#[cfg(test)]
pub(crate) fn with_provider_catalog_and_request_candidate_reader_for_tests(
provider_catalog_repository: Arc<dyn ProviderCatalogReadRepository>,
@@ -363,6 +385,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: Some(provider_catalog_repository),
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -419,6 +443,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: Some(provider_catalog_repository),
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer),
usage_reader: None,
@@ -484,6 +510,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: Some(provider_catalog_reader),
provider_catalog_writer: Some(provider_catalog_writer),
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer),
usage_reader: None,
@@ -531,6 +559,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -579,6 +609,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: Some(provider_catalog_repository),
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -638,6 +670,8 @@ impl GatewayDataState {
request_candidate_writer: Some(request_candidate_writer),
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: Some(usage_reader),
@@ -699,6 +733,8 @@ impl GatewayDataState {
request_candidate_writer: Some(request_candidate_writer),
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -744,6 +780,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: Some(repository),
@@ -804,6 +842,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -857,6 +897,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -915,6 +957,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: Some(usage_reader),
@@ -974,6 +1018,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -1032,6 +1078,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: Some(provider_catalog_reader),
provider_catalog_writer: Some(provider_catalog_writer),
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: Some(usage_reader),
@@ -1079,6 +1127,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -1126,6 +1176,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -1185,6 +1237,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -1249,6 +1303,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -1296,6 +1352,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -1348,6 +1406,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -1417,6 +1477,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -1481,6 +1543,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -1529,6 +1593,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: Some(provider_catalog_repository),
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -1577,6 +1643,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: Some(repository),
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -1627,6 +1695,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: Some(provider_catalog_repository),
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: Some(usage_repository),
@@ -1675,6 +1745,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -1723,6 +1795,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -1779,6 +1853,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer),
usage_reader: None,
@@ -1836,6 +1912,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer),
usage_reader: None,
@@ -1896,6 +1974,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: Some(provider_catalog_repository),
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer),
usage_reader: None,
@@ -1962,6 +2042,8 @@ impl GatewayDataState {
request_candidate_writer: Some(request_candidate_writer),
provider_catalog_reader: Some(provider_catalog_reader),
provider_catalog_writer: Some(provider_catalog_writer),
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -2029,6 +2111,8 @@ impl GatewayDataState {
request_candidate_writer: Some(request_candidate_writer),
provider_catalog_reader: Some(provider_catalog_reader),
provider_catalog_writer: Some(provider_catalog_writer),
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -2100,6 +2184,8 @@ impl GatewayDataState {
request_candidate_writer: Some(request_candidate_writer),
provider_catalog_reader: Some(provider_catalog_reader),
provider_catalog_writer: Some(provider_catalog_writer),
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: Some(usage_reader),
@@ -2178,6 +2264,8 @@ impl GatewayDataState {
request_candidate_writer: Some(request_candidate_writer),
provider_catalog_reader: Some(provider_catalog_reader),
provider_catalog_writer: Some(provider_catalog_writer),
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: Some(usage_reader),
@@ -2238,6 +2326,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: Some(provider_catalog_repository),
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer),
usage_reader: None,
@@ -2289,6 +2379,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: Some(usage_reader),
@@ -2336,6 +2428,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -2389,6 +2483,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -2446,6 +2542,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: Some(usage_reader),
@@ -2504,6 +2602,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: Some(usage_reader),
@@ -2555,6 +2655,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer),
usage_reader: None,

View File

@@ -38,6 +38,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -92,6 +94,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -143,6 +147,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -198,6 +204,8 @@ impl GatewayDataState {
request_candidate_writer: None,
provider_catalog_reader: Some(provider_catalog_repository),
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -257,6 +265,8 @@ impl GatewayDataState {
request_candidate_writer: Some(request_candidate_writer),
provider_catalog_reader: None,
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,
@@ -325,6 +335,8 @@ impl GatewayDataState {
request_candidate_writer: Some(request_candidate_writer),
provider_catalog_reader: Some(provider_catalog_reader),
provider_catalog_writer: None,
pool_score_reader: None,
pool_score_writer: None,
provider_quota_reader: None,
provider_quota_writer: None,
usage_reader: None,

View File

@@ -0,0 +1,3 @@
pub(crate) mod pool;
pub(crate) mod pool_scheduler;
pub(crate) mod refs;

View File

@@ -0,0 +1,5 @@
use aether_dispatch_core::PoolWindowConfig;
pub(crate) fn default_pool_window_config() -> PoolWindowConfig {
PoolWindowConfig::default()
}

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,213 @@
use aether_dispatch_core::{
DispatchCandidateRef, DispatchRankFacts, KeyRef, PoolRef, ProviderEndpointRef,
};
use crate::ai_serving::{EligibleLocalExecutionCandidate, LocalExecutionCandidateKind};
pub(crate) fn dispatch_ref_for_local_candidate(
eligible: &EligibleLocalExecutionCandidate,
) -> DispatchCandidateRef {
let rank = DispatchRankFacts {
provider_priority: eligible.candidate.provider_priority,
key_priority: Some(eligible.candidate.key_internal_priority),
ranking_reason: eligible.ranking.as_ref().and_then(|ranking| {
ranking
.promoted_by
.or(ranking.demoted_by)
.map(str::to_string)
}),
};
match eligible.kind {
LocalExecutionCandidateKind::SingleKey => DispatchCandidateRef::SingleKey {
key: key_ref_for_candidate(eligible),
rank,
},
LocalExecutionCandidateKind::PoolGroup => DispatchCandidateRef::PoolRef {
pool: pool_ref_for_candidate(eligible),
rank,
},
}
}
pub(crate) fn key_ref_for_candidate(eligible: &EligibleLocalExecutionCandidate) -> KeyRef {
KeyRef {
provider_id: eligible.candidate.provider_id.clone(),
endpoint_id: eligible.candidate.endpoint_id.clone(),
key_id: eligible.candidate.key_id.clone(),
model_id: eligible.candidate.model_id.clone(),
selected_provider_model_name: eligible.candidate.selected_provider_model_name.clone(),
api_format: eligible.candidate.endpoint_api_format.clone(),
}
}
pub(crate) fn pool_ref_for_candidate(eligible: &EligibleLocalExecutionCandidate) -> PoolRef {
PoolRef {
provider_id: eligible.candidate.provider_id.clone(),
endpoint_id: eligible.candidate.endpoint_id.clone(),
model_id: eligible.candidate.model_id.clone(),
selected_provider_model_name: eligible.candidate.selected_provider_model_name.clone(),
api_format: eligible.candidate.endpoint_api_format.clone(),
pool_group_id: eligible
.orchestration
.candidate_group_id
.clone()
.unwrap_or_else(|| pool_group_id_for_provider_endpoint(eligible)),
}
}
pub(crate) fn provider_endpoint_ref_for_candidate(
eligible: &EligibleLocalExecutionCandidate,
) -> ProviderEndpointRef {
ProviderEndpointRef {
provider_id: eligible.candidate.provider_id.clone(),
endpoint_id: eligible.candidate.endpoint_id.clone(),
model_id: eligible.candidate.model_id.clone(),
selected_provider_model_name: eligible.candidate.selected_provider_model_name.clone(),
api_format: eligible.candidate.endpoint_api_format.clone(),
}
}
fn pool_group_id_for_provider_endpoint(eligible: &EligibleLocalExecutionCandidate) -> String {
format!(
"provider={}|endpoint={}|model={}|selected_model={}|api_format={}",
eligible.candidate.provider_id,
eligible.candidate.endpoint_id,
eligible.candidate.model_id,
eligible.candidate.selected_provider_model_name,
eligible.candidate.endpoint_api_format
)
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use aether_dispatch_core::DispatchCandidateRef;
use aether_provider_transport::snapshot::{
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider,
};
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
use super::dispatch_ref_for_local_candidate;
use crate::ai_serving::{EligibleLocalExecutionCandidate, LocalExecutionCandidateKind};
use crate::orchestration::LocalExecutionCandidateMetadata;
#[test]
fn pool_group_maps_to_pool_ref_without_exposing_internal_key() {
let eligible = sample_eligible(LocalExecutionCandidateKind::PoolGroup);
let dispatch_ref = dispatch_ref_for_local_candidate(&eligible);
match dispatch_ref {
DispatchCandidateRef::PoolRef { pool, rank } => {
assert_eq!(pool.provider_id, "provider-1");
assert_eq!(pool.endpoint_id, "endpoint-1");
assert_eq!(pool.pool_group_id, "group-1");
assert_eq!(rank.provider_priority, 10);
}
other => panic!("expected pool ref, got {other:?}"),
}
}
#[test]
fn single_key_maps_to_key_ref() {
let eligible = sample_eligible(LocalExecutionCandidateKind::SingleKey);
let dispatch_ref = dispatch_ref_for_local_candidate(&eligible);
match dispatch_ref {
DispatchCandidateRef::SingleKey { key, rank } => {
assert_eq!(key.key_id, "key-1");
assert_eq!(rank.key_priority, Some(7));
}
other => panic!("expected key ref, got {other:?}"),
}
}
fn sample_eligible(kind: LocalExecutionCandidateKind) -> EligibleLocalExecutionCandidate {
EligibleLocalExecutionCandidate {
kind,
candidate: SchedulerMinimalCandidateSelectionCandidate {
provider_id: "provider-1".to_string(),
provider_name: "Provider 1".to_string(),
provider_type: "openai".to_string(),
provider_priority: 10,
endpoint_id: "endpoint-1".to_string(),
endpoint_api_format: "openai:chat".to_string(),
key_id: "key-1".to_string(),
key_name: "Key 1".to_string(),
key_auth_type: "api_key".to_string(),
key_internal_priority: 7,
key_global_priority_for_format: None,
key_capabilities: None,
model_id: "model-1".to_string(),
global_model_id: "global-model-1".to_string(),
global_model_name: "gpt-5".to_string(),
selected_provider_model_name: "gpt-5".to_string(),
mapping_matched_model: None,
},
transport: Arc::new(crate::ai_serving::GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: "provider-1".to_string(),
name: "Provider 1".to_string(),
provider_type: "openai".to_string(),
website: None,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
},
endpoint: GatewayProviderTransportEndpoint {
id: "endpoint-1".to_string(),
provider_id: "provider-1".to_string(),
api_format: "openai:chat".to_string(),
api_family: Some("openai".to_string()),
endpoint_kind: Some("chat".to_string()),
is_active: true,
base_url: "https://example.com".to_string(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
},
key: GatewayProviderTransportKey {
id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
name: "Key 1".to_string(),
auth_type: "api_key".to_string(),
is_active: true,
api_formats: Some(vec!["openai:chat".to_string()]),
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
allowed_models: None,
capabilities: None,
rate_multipliers: None,
global_priority_by_format: None,
expires_at_unix_secs: None,
proxy: None,
fingerprint: None,
decrypted_api_key: "secret".to_string(),
decrypted_auth_config: None,
},
}),
provider_api_format: "openai:chat".to_string(),
orchestration: LocalExecutionCandidateMetadata {
candidate_group_id: Some("group-1".to_string()),
pool_key_index: None,
pool_key_lease: None,
scheduler_affinity_epoch: None,
},
ranking: None,
}
}
}

View File

@@ -0,0 +1,87 @@
use super::shared::{
admin_api_key_install_session_id_from_path, build_admin_api_keys_bad_request_response,
build_admin_api_keys_data_unavailable_response, build_admin_api_keys_not_found_response,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::handlers::public::{
build_api_key_install_session_response, CreateApiKeyInstallSessionRequest,
};
use crate::GatewayError;
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
};
pub(super) async fn build_admin_create_api_key_install_session_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_headers: &http::HeaderMap,
request_body: Option<&axum::body::Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_auth_api_key_data_reader() {
return Ok(build_admin_api_keys_data_unavailable_response());
}
let Some(api_key_id) = admin_api_key_install_session_id_from_path(request_context.path())
else {
return Ok(build_admin_api_keys_data_unavailable_response());
};
let Some(request_body) = request_body else {
return Ok(build_admin_api_keys_bad_request_response(
"请求数据验证失败",
));
};
let payload = match serde_json::from_slice::<CreateApiKeyInstallSessionRequest>(request_body) {
Ok(value) => value,
Err(_) => {
return Ok(build_admin_api_keys_bad_request_response(
"请求数据验证失败",
))
}
};
let Some(record) = state
.find_auth_api_key_export_standalone_record_by_id(&api_key_id)
.await?
else {
return Ok(build_admin_api_keys_not_found_response());
};
let Some(ciphertext) = record
.key_encrypted
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(build_admin_api_keys_bad_request_response(
"该密钥没有存储完整密钥信息",
));
};
let Some(api_key) = state.decrypt_catalog_secret_with_fallbacks(ciphertext) else {
return Ok((
http::StatusCode::INTERNAL_SERVER_ERROR,
axum::Json(serde_json::json!({ "detail": "解密密钥失败" })),
)
.into_response());
};
let response = build_api_key_install_session_response(
state.app(),
request_context.public(),
request_headers,
record.api_key_id.clone(),
record.name.unwrap_or_else(|| "API Key".to_string()),
api_key,
payload,
)
.await;
Ok(attach_admin_audit_response(
response,
"admin_standalone_api_key_install_session_created",
"create_standalone_api_key_install_session",
"api_key",
&api_key_id,
))
}

View File

@@ -18,11 +18,13 @@ use axum::{
};
use serde_json::json;
mod install_routes;
mod mutation_routes;
mod read_routes;
mod routes;
mod shared;
use self::install_routes::build_admin_create_api_key_install_session_response;
use self::mutation_routes::{
build_admin_create_api_key_response, build_admin_delete_api_key_response,
build_admin_toggle_api_key_response, build_admin_update_api_key_response,
@@ -39,8 +41,14 @@ use self::shared::{
pub(crate) async fn maybe_build_local_admin_api_keys_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_headers: &http::HeaderMap,
request_body: Option<&axum::body::Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
routes::maybe_build_local_admin_api_keys_routes_response(state, request_context, request_body)
.await
routes::maybe_build_local_admin_api_keys_routes_response(
state,
request_context,
request_headers,
request_body,
)
.await
}

View File

@@ -1,3 +1,4 @@
use super::install_routes::build_admin_create_api_key_install_session_response;
use super::mutation_routes::{
build_admin_create_api_key_response, build_admin_delete_api_key_response,
build_admin_toggle_api_key_response, build_admin_update_api_key_response,
@@ -11,6 +12,7 @@ use axum::{body::Body, http, response::Response};
pub(super) async fn maybe_build_local_admin_api_keys_routes_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_headers: &http::HeaderMap,
request_body: Option<&axum::body::Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.decision() else {
@@ -22,8 +24,10 @@ pub(super) async fn maybe_build_local_admin_api_keys_routes_response(
}
let path = request_context.path();
let path_no_trailing = path.trim_end_matches('/');
let is_api_keys_route = matches!(path, "/api/admin/api-keys" | "/api/admin/api-keys/")
|| (path.starts_with("/api/admin/api-keys/") && path.matches('/').count() == 4);
|| (path_no_trailing.starts_with("/api/admin/api-keys/")
&& matches!(path_no_trailing.matches('/').count(), 4 | 5));
if !is_api_keys_route {
return Ok(None);
@@ -54,6 +58,21 @@ pub(super) async fn maybe_build_local_admin_api_keys_routes_response(
build_admin_create_api_key_response(state, request_context, request_body).await?,
))
}
Some("create_api_key_install_session")
if request_context.method() == http::Method::POST
&& path_no_trailing.starts_with("/api/admin/api-keys/")
&& path_no_trailing.ends_with("/install-sessions") =>
{
Ok(Some(
build_admin_create_api_key_install_session_response(
state,
request_context,
request_headers,
request_body,
)
.await?,
))
}
Some("update_api_key")
if request_context.method() == http::Method::PUT
&& path.starts_with("/api/admin/api-keys/") =>

View File

@@ -91,6 +91,17 @@ pub(super) fn admin_api_keys_id_from_path(request_path: &str) -> Option<String>
}
}
pub(super) fn admin_api_key_install_session_id_from_path(request_path: &str) -> Option<String> {
let raw = request_path
.strip_prefix("/api/admin/api-keys/")?
.trim()
.trim_matches('/');
let mut segments = raw.split('/').map(str::trim);
let api_key_id = segments.next()?.to_string();
let suffix = segments.next()?;
(suffix == "install-sessions" && segments.next().is_none()).then_some(api_key_id)
}
pub(super) fn admin_api_keys_operator_id(
request_context: &AdminRequestContext<'_>,
) -> Option<String> {

View File

@@ -17,6 +17,7 @@ pub(crate) async fn maybe_build_local_admin_auth_response(
if let Some(response) = api_keys::maybe_build_local_admin_api_keys_response(
&request.state(),
&request.request_context(),
request.request_headers(),
request.request_body(),
)
.await?

View File

@@ -29,12 +29,14 @@ pub(crate) use self::provider::oauth::quota::antigravity::refresh_antigravity_pr
pub(crate) use self::provider::oauth::quota::chatgpt_web::refresh_chatgpt_web_provider_quota_locally;
pub(crate) use self::provider::oauth::quota::codex::refresh_codex_provider_quota_locally;
pub(crate) use self::provider::oauth::quota::kiro::refresh_kiro_provider_quota_locally;
pub(crate) use self::provider::oauth::quota::shared::provider_type_supports_quota_refresh;
pub(crate) use self::provider::oauth::runtime::{
provider_oauth_runtime_endpoint_for_provider, refresh_provider_oauth_account_state_after_update,
};
pub(crate) use self::provider::ops::providers::actions::admin_provider_ops_local_action_response;
pub(crate) use self::provider::pool::config::admin_provider_pool_config;
pub(crate) use self::provider::pool_admin::maybe_build_local_admin_pool_response;
pub(crate) use self::provider::write::provider::reconcile_admin_fixed_provider_template_endpoints;
pub(crate) use self::provider::{
maybe_build_local_admin_provider_oauth_response, maybe_build_local_admin_providers_response,
};

View File

@@ -265,8 +265,8 @@ async fn admin_monitoring_cache_stats_count_runtime_scheduler_affinities() {
"model-alpha",
)
.expect("scheduler affinity cache key should build");
state.scheduler_affinity_cache.insert(
affinity_cache_key,
state.remember_scheduler_affinity_target(
&affinity_cache_key,
crate::cache::SchedulerAffinityTarget {
provider_id: "provider-1".to_string(),
endpoint_id: "endpoint-1".to_string(),

View File

@@ -239,8 +239,8 @@ async fn admin_monitoring_cache_affinities_and_delete_use_runtime_scheduler_affi
"model-alpha",
)
.expect("scheduler affinity cache key should build");
state.scheduler_affinity_cache.insert(
affinity_cache_key.clone(),
state.remember_scheduler_affinity_target(
&affinity_cache_key,
crate::cache::SchedulerAffinityTarget {
provider_id: "provider-1".to_string(),
endpoint_id: "endpoint-1".to_string(),
@@ -393,8 +393,8 @@ async fn admin_monitoring_cache_affinities_parse_session_scoped_scheduler_affini
.next()
.expect("session hash should exist")
.to_string();
state.scheduler_affinity_cache.insert(
affinity_cache_key.clone(),
state.remember_scheduler_affinity_target(
&affinity_cache_key,
crate::cache::SchedulerAffinityTarget {
provider_id: "provider-1".to_string(),
endpoint_id: "endpoint-1".to_string(),
@@ -403,8 +403,8 @@ async fn admin_monitoring_cache_affinities_parse_session_scoped_scheduler_affini
crate::scheduler::affinity::SCHEDULER_AFFINITY_TTL,
128,
);
state.scheduler_affinity_cache.insert(
other_affinity_cache_key.clone(),
state.remember_scheduler_affinity_target(
&other_affinity_cache_key,
crate::cache::SchedulerAffinityTarget {
provider_id: "provider-1".to_string(),
endpoint_id: "endpoint-1".to_string(),

View File

@@ -1,6 +1,8 @@
use crate::handlers::admin::admin_provider_pool_config;
use crate::handlers::admin::provider::shared::paths::admin_provider_id_for_keys;
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyCreateRequest;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::maintenance::ensure_provider_key_pool_scores_for_keys;
use crate::provider_key_auth::provider_key_effective_api_formats;
use crate::{model_fetch::perform_model_fetch_for_key, GatewayError};
use axum::{
@@ -98,6 +100,27 @@ pub(super) async fn maybe_handle(
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
if let Some(pool_config) = admin_provider_pool_config(&provider) {
let score_ensure_budget = (pool_config.score_fallback_scan_limit as usize).clamp(1, 50_000);
if let Err(err) = ensure_provider_key_pool_scores_for_keys(
state.as_ref(),
&provider,
&pool_config,
&endpoints,
std::slice::from_ref(&created),
now_unix_secs,
score_ensure_budget,
)
.await
{
tracing::debug!(
provider_id = %provider.id,
key_id = %created.id,
error = ?err,
"gateway admin provider key create: failed to seed pool score rows"
);
}
}
let api_formats =
provider_key_effective_api_formats(&created, &provider.provider_type, &endpoints);

View File

@@ -1,6 +1,8 @@
use crate::handlers::admin::admin_provider_pool_config;
use crate::handlers::admin::provider::shared::paths::admin_update_key_id;
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdatePatch;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::maintenance::ensure_provider_key_pool_scores_for_keys;
use crate::provider_key_auth::provider_key_effective_api_formats;
use crate::{model_fetch::perform_model_fetch_for_key, GatewayError};
use axum::{
@@ -121,6 +123,27 @@ pub(super) async fn maybe_handle(
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
if let Some(pool_config) = admin_provider_pool_config(&provider) {
let score_ensure_budget = (pool_config.score_fallback_scan_limit as usize).clamp(1, 50_000);
if let Err(err) = ensure_provider_key_pool_scores_for_keys(
state.as_ref(),
&provider,
&pool_config,
&endpoints,
std::slice::from_ref(&updated),
now_unix_secs,
score_ensure_budget,
)
.await
{
tracing::debug!(
provider_id = %provider.id,
key_id = %updated.id,
error = ?err,
"gateway admin provider key update: failed to seed pool score rows"
);
}
}
let api_formats =
provider_key_effective_api_formats(&updated, &provider.provider_type, &endpoints);

View File

@@ -18,6 +18,24 @@ use super::super::oauth::quota::chatgpt_web::refresh_chatgpt_web_provider_quota_
use super::super::oauth::quota::codex::refresh_codex_provider_quota_locally;
use super::super::oauth::quota::kiro::refresh_kiro_provider_quota_locally;
use super::super::oauth::quota::shared::normalize_string_id_list;
use super::super::oauth::quota::shared::{
provider_type_supports_quota_refresh, unsupported_provider_quota_refresh_message,
};
use super::super::oauth::runtime::provider_oauth_runtime_endpoint_for_provider;
use super::super::write::provider::reconcile_admin_fixed_provider_template_endpoints;
fn unsupported_provider_quota_refresh_response(provider_type: &str) -> Response<Body> {
let message = unsupported_provider_quota_refresh_message(provider_type);
Json(json!({
"success": 0,
"failed": 0,
"total": 0,
"results": [],
"message": message,
"auto_removed": 0,
}))
.into_response()
}
pub(super) async fn maybe_handle(
state: &AdminAppState<'_>,
@@ -85,41 +103,48 @@ pub(super) async fn maybe_handle(
let raw_key_ids = payload.key_ids;
let selected_key_ids = normalize_string_id_list(raw_key_ids.clone());
let explicit_key_ids_requested = raw_key_ids.is_some();
let endpoints = state
let is_fixed_provider = state
.fixed_provider_template(&provider.provider_type)
.is_some();
if !is_fixed_provider && !provider_type_supports_quota_refresh(&normalized_provider_type) {
return Ok(None);
}
let mut endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let endpoint = match normalized_provider_type.as_str() {
"codex" => endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& crate::ai_serving::is_openai_responses_format(&endpoint.api_format)
}),
"antigravity" => endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:generate_content")
}),
"kiro" => endpoints
.iter()
.find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("claude:messages")
})
.cloned()
.or_else(|| endpoints.into_iter().find(|endpoint| endpoint.is_active)),
"chatgpt_web" => endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("openai:image")
}),
_ => return Ok(None),
};
let mut endpoint =
provider_oauth_runtime_endpoint_for_provider(&normalized_provider_type, &endpoints);
if endpoint.is_none() && is_fixed_provider {
if !state.has_provider_catalog_data_writer() {
if !provider_type_supports_quota_refresh(&normalized_provider_type) {
return Ok(Some(unsupported_provider_quota_refresh_response(
&normalized_provider_type,
)));
}
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "固定 Provider 端点缺失,且 provider catalog writer 不可用,无法自动补全端点" })),
)
.into_response(),
));
}
reconcile_admin_fixed_provider_template_endpoints(state, &provider).await?;
endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
endpoint =
provider_oauth_runtime_endpoint_for_provider(&normalized_provider_type, &endpoints);
}
if !provider_type_supports_quota_refresh(&normalized_provider_type) {
return Ok(Some(unsupported_provider_quota_refresh_response(
&normalized_provider_type,
)));
}
let Some(endpoint) = endpoint else {
let detail = match normalized_provider_type.as_str() {
@@ -127,6 +152,8 @@ pub(super) async fn maybe_handle(
"antigravity" => "找不到有效的 gemini:generate_content 端点",
"kiro" => "找不到有效的 Kiro 端点",
"chatgpt_web" => "找不到有效的 openai:image 端点",
"claude_code" => "找不到有效的 claude:messages 端点",
"gemini_cli" | "vertex_ai" => "找不到有效的 gemini:generate_content 端点",
_ => "找不到有效端点",
};
return Ok(Some(

View File

@@ -18,7 +18,7 @@ use crate::handlers::admin::provider::oauth::provisioning::{
provider_oauth_key_proxy_value, update_existing_provider_oauth_catalog_key,
};
use crate::handlers::admin::provider::oauth::runtime::{
provider_oauth_runtime_endpoint_for_provider,
resolve_provider_oauth_runtime_endpoints,
spawn_provider_oauth_account_state_refresh_after_update,
};
use crate::handlers::admin::provider::oauth::state::{
@@ -224,11 +224,11 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
});
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(&[provider_id.to_string()])
.await?;
let endpoint_resolution =
resolve_provider_oauth_runtime_endpoints(state, &provider, provider_type).await?;
let endpoints = endpoint_resolution.endpoints;
let api_formats = provider_oauth_active_api_formats(&endpoints);
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider(provider_type, &endpoints);
let runtime_endpoint = endpoint_resolution.runtime_endpoint;
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
proxy_node_id,

View File

@@ -13,7 +13,7 @@ use crate::handlers::admin::provider::oauth::provisioning::{
provider_oauth_key_proxy_value, update_existing_provider_oauth_catalog_key,
};
use crate::handlers::admin::provider::oauth::runtime::{
provider_oauth_runtime_endpoint_for_provider,
resolve_provider_oauth_runtime_endpoints,
spawn_provider_oauth_account_state_refresh_after_update,
};
use crate::handlers::admin::provider::oauth::state::decode_jwt_claims;
@@ -59,11 +59,11 @@ pub(super) async fn execute_admin_provider_oauth_kiro_batch_import(
});
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(&[provider_id.to_string()])
.await?;
let endpoint_resolution =
resolve_provider_oauth_runtime_endpoints(state, &provider, "kiro").await?;
let endpoints = endpoint_resolution.endpoints;
let api_formats = provider_oauth_active_api_formats(&endpoints);
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider("kiro", &endpoints);
let runtime_endpoint = endpoint_resolution.runtime_endpoint;
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
proxy_node_id,

View File

@@ -1,7 +1,7 @@
use super::super::super::errors::build_internal_control_error_response;
use super::super::super::provisioning::provider_oauth_token_payload_expires_at_unix_secs;
use super::super::super::quota::codex::refresh_codex_provider_quota_locally;
use super::super::super::runtime::provider_oauth_runtime_endpoint_for_provider;
use super::super::super::runtime::resolve_provider_oauth_runtime_endpoints;
use super::super::super::state::{
admin_provider_oauth_template, enrich_admin_provider_oauth_auth_config,
is_fixed_provider_type_for_provider_oauth, json_non_empty_string,
@@ -126,10 +126,10 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
"该 Provider 不支持 OAuth 授权",
));
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider(&provider_type, &endpoints);
let endpoint_resolution =
resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?;
let endpoints = endpoint_resolution.endpoints;
let runtime_endpoint = endpoint_resolution.runtime_endpoint;
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
payload.proxy_node_id.as_deref(),
@@ -223,9 +223,6 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
let mut account_state_recheck_attempted = false;
let mut account_state_recheck_error = None::<String>;
if provider_type == "codex" {
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
if let Some(endpoint) = endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& crate::ai_serving::is_openai_responses_format(&endpoint.api_format)

View File

@@ -6,7 +6,7 @@ use super::super::super::provisioning::{
update_existing_provider_oauth_catalog_key,
};
use super::super::super::runtime::{
provider_oauth_runtime_endpoint_for_provider,
resolve_provider_oauth_runtime_endpoints,
spawn_provider_oauth_account_state_refresh_after_update,
};
use super::super::super::state::{
@@ -117,10 +117,10 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
"该 Provider 不支持 OAuth 授权",
));
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider(&provider_type, &endpoints);
let endpoint_resolution =
resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?;
let endpoints = endpoint_resolution.endpoints;
let runtime_endpoint = endpoint_resolution.runtime_endpoint;
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
payload.proxy_node_id.as_deref(),

View File

@@ -1,6 +1,6 @@
use super::session::AdminProviderOAuthDeviceAuthorizePayload;
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
use crate::handlers::admin::provider::oauth::runtime::provider_oauth_runtime_endpoint_for_provider;
use crate::handlers::admin::provider::oauth::runtime::resolve_provider_oauth_runtime_endpoints;
use crate::handlers::admin::provider::oauth::state::{
build_admin_provider_oauth_backend_unavailable_response, current_unix_secs,
default_kiro_device_start_url, generate_provider_oauth_nonce, json_non_empty_string,
@@ -169,10 +169,9 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
"设备授权仅支持 Kiro provider",
));
}
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider("kiro", &endpoints);
let endpoint_resolution =
resolve_provider_oauth_runtime_endpoints(state, &provider, "kiro").await?;
let runtime_endpoint = endpoint_resolution.runtime_endpoint;
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
payload.proxy_node_id.as_deref(),

View File

@@ -10,7 +10,7 @@ use crate::handlers::admin::provider::oauth::provisioning::{
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
};
use crate::handlers::admin::provider::oauth::runtime::{
provider_oauth_runtime_endpoint_for_provider,
resolve_provider_oauth_runtime_endpoints,
spawn_provider_oauth_account_state_refresh_after_update,
};
use crate::handlers::admin::provider::oauth::state::{
@@ -322,10 +322,10 @@ pub(super) async fn handle_admin_provider_oauth_device_poll(
"Provider 不存在",
));
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider("kiro", &endpoints);
let endpoint_resolution =
resolve_provider_oauth_runtime_endpoints(state, &provider, "kiro").await?;
let endpoints = endpoint_resolution.endpoints;
let runtime_endpoint = endpoint_resolution.runtime_endpoint;
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
session.proxy_node_id.as_deref(),

View File

@@ -6,7 +6,7 @@ use super::super::provisioning::{
update_existing_provider_oauth_catalog_key,
};
use super::super::runtime::{
provider_oauth_runtime_endpoint_for_provider,
resolve_provider_oauth_runtime_endpoints,
spawn_provider_oauth_account_state_refresh_after_update,
};
use super::super::state::{
@@ -300,10 +300,10 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let runtime_endpoint = provider_oauth_runtime_endpoint_for_provider(&provider_type, &endpoints);
let endpoint_resolution =
resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?;
let endpoints = endpoint_resolution.endpoints;
let runtime_endpoint = endpoint_resolution.runtime_endpoint;
let request_proxy = state
.resolve_admin_provider_oauth_operation_proxy_snapshot(
proxy_node_id.as_deref(),

View File

@@ -1,4 +1,4 @@
use super::super::super::runtime::provider_oauth_runtime_endpoint_for_provider;
use super::super::super::runtime::resolve_provider_oauth_runtime_endpoints;
use super::super::super::state::is_fixed_provider_type_for_provider_oauth;
use super::helpers::{self, RefreshDispatch, RefreshRequestContext};
use super::response;
@@ -76,11 +76,19 @@ pub(super) async fn parse_admin_provider_oauth_refresh_request(
)));
}
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let Some(endpoint) = provider_oauth_runtime_endpoint_for_provider(&provider_type, &endpoints)
else {
let endpoint_resolution =
resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?;
let Some(endpoint) = endpoint_resolution.runtime_endpoint else {
if state
.fixed_provider_template(&provider.provider_type)
.is_some()
&& !state.has_provider_catalog_data_writer()
{
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::BAD_REQUEST,
"固定 Provider 端点缺失,且 provider catalog writer 不可用,无法自动补全端点",
)));
}
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::BAD_REQUEST,
"找不到有效端点,无法 refresh",

View File

@@ -50,6 +50,25 @@ pub(crate) fn normalize_string_id_list(values: Option<Vec<String>>) -> Option<Ve
admin_provider_quota_pure::normalize_string_id_list(values)
}
pub(crate) fn provider_type_supports_quota_refresh(provider_type: &str) -> bool {
matches!(
provider_type.trim().to_ascii_lowercase().as_str(),
"codex" | "kiro" | "antigravity" | "chatgpt_web"
)
}
pub(crate) fn unsupported_provider_quota_refresh_message(provider_type: &str) -> String {
match provider_type.trim().to_ascii_lowercase().as_str() {
"claude_code" => "Claude Code 暂不支持自动刷新额度:上游没有稳定可用的账号额度查询接口",
"gemini_cli" => {
"Gemini CLI 暂不支持自动刷新额度:当前只能通过模型同步/缓存快照展示已知配额信息"
}
"vertex_ai" => "Vertex AI 暂不支持自动刷新额度:额度属于 Google Cloud 项目/区域配额",
_ => "该 Provider 暂不支持自动刷新额度",
}
.to_string()
}
pub(super) fn coerce_json_u64(value: &serde_json::Value) -> Option<u64> {
admin_provider_quota_pure::coerce_json_u64(value)
}

View File

@@ -2,6 +2,7 @@ use super::quota::antigravity::refresh_antigravity_provider_quota_locally;
use super::quota::chatgpt_web::refresh_chatgpt_web_provider_quota_locally;
use super::quota::codex::refresh_codex_provider_quota_locally;
use super::quota::kiro::refresh_kiro_provider_quota_locally;
use crate::handlers::admin::provider::write::provider::reconcile_admin_fixed_provider_template_endpoints;
use crate::handlers::admin::request::AdminAppState;
use crate::provider_key_auth::provider_key_is_oauth_managed;
use crate::task_runtime::{spawn_fire_and_forget, TASK_KEY_PROVIDER_OAUTH_ACCOUNT_REFRESH};
@@ -60,6 +61,48 @@ pub(crate) fn provider_oauth_runtime_endpoint_for_provider(
.find(|endpoint| endpoint.is_active)
.cloned()
}),
"claude_code" => endpoints
.iter()
.find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("claude:messages")
})
.cloned(),
"gemini_cli" => endpoints
.iter()
.find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:generate_content")
})
.cloned(),
"vertex_ai" => endpoints
.iter()
.find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:generate_content")
})
.cloned()
.or_else(|| {
endpoints
.iter()
.find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("claude:messages")
})
.cloned()
}),
_ => endpoints
.iter()
.find(|endpoint| endpoint.is_active)
@@ -67,6 +110,41 @@ pub(crate) fn provider_oauth_runtime_endpoint_for_provider(
}
}
#[derive(Debug, Clone)]
pub(crate) struct ProviderOAuthRuntimeEndpoints {
pub(crate) endpoints: Vec<StoredProviderCatalogEndpoint>,
pub(crate) runtime_endpoint: Option<StoredProviderCatalogEndpoint>,
}
pub(crate) async fn resolve_provider_oauth_runtime_endpoints(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
provider_type: &str,
) -> Result<ProviderOAuthRuntimeEndpoints, GatewayError> {
let mut endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let mut runtime_endpoint =
provider_oauth_runtime_endpoint_for_provider(provider_type, &endpoints);
if runtime_endpoint.is_none()
&& state
.fixed_provider_template(&provider.provider_type)
.is_some()
&& state.has_provider_catalog_data_writer()
{
reconcile_admin_fixed_provider_template_endpoints(state, provider).await?;
endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
runtime_endpoint = provider_oauth_runtime_endpoint_for_provider(provider_type, &endpoints);
}
Ok(ProviderOAuthRuntimeEndpoints {
endpoints,
runtime_endpoint,
})
}
pub(crate) async fn refresh_provider_oauth_account_state_after_update(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
@@ -81,11 +159,10 @@ pub(crate) async fn refresh_provider_oauth_account_state_after_update(
return Ok((false, None));
}
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let Some(endpoint) = provider_oauth_runtime_endpoint_for_provider(&provider_type, &endpoints)
else {
let ProviderOAuthRuntimeEndpoints {
runtime_endpoint, ..
} = resolve_provider_oauth_runtime_endpoints(state, provider, &provider_type).await?;
let Some(endpoint) = runtime_endpoint else {
return Ok((false, None));
};
let Some(key) = state

View File

@@ -1,6 +1,7 @@
use crate::handlers::admin::provider::shared::support::{
AdminProviderPoolConfig, AdminProviderPoolSchedulingPreset, AdminProviderPoolUnschedulableRule,
};
use aether_ai_serving::{PoolMemberScoreRules, PoolMemberScoreWeights};
use serde_json::{Map, Value};
const POOL_ALLOWED_SCHEDULING_PRESETS: &[&str] = &[
@@ -26,6 +27,123 @@ fn json_u64(value: &Value) -> Option<u64> {
.or_else(|| value.as_i64().and_then(|raw| u64::try_from(raw).ok()))
}
fn json_f64(value: &Value) -> Option<f64> {
value.as_f64().or_else(|| {
value
.as_str()
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(|value| value.parse::<f64>().ok())
})
}
fn pool_score_weight(object: &Map<String, Value>, names: &[&str], current: f64) -> f64 {
names
.iter()
.find_map(|name| {
object
.get(*name)
.and_then(json_f64)
.filter(|value| value.is_finite() && *value >= 0.0)
})
.unwrap_or(current)
}
fn parse_pool_score_weights(
raw_weights: Option<&Map<String, Value>>,
current: PoolMemberScoreWeights,
) -> PoolMemberScoreWeights {
let Some(raw_weights) = raw_weights else {
return current;
};
PoolMemberScoreWeights {
manual_priority: pool_score_weight(
raw_weights,
&["manual_priority", "priority", "internal_priority"],
current.manual_priority,
),
health: pool_score_weight(raw_weights, &["health"], current.health),
probe_freshness: pool_score_weight(
raw_weights,
&["probe_freshness", "freshness", "probe"],
current.probe_freshness,
),
quota_remaining: pool_score_weight(
raw_weights,
&["quota_remaining", "quota", "quota_available"],
current.quota_remaining,
),
latency: pool_score_weight(raw_weights, &["latency"], current.latency),
cost_lru: pool_score_weight(
raw_weights,
&["cost_lru", "cost_remaining", "cost", "lru"],
current.cost_lru,
),
}
}
fn parse_pool_score_rules(pool_advanced: &Map<String, Value>) -> PoolMemberScoreRules {
let mut rules = PoolMemberScoreRules::default();
for key in ["score_weights", "pool_score_weights", "scoring_weights"] {
rules.weights = parse_pool_score_weights(
pool_advanced.get(key).and_then(Value::as_object),
rules.weights,
);
}
if let Some(score_rules) = pool_advanced
.get("score_rules")
.or_else(|| pool_advanced.get("pool_score_rules"))
.and_then(Value::as_object)
{
rules.weights = parse_pool_score_weights(
score_rules.get("weights").and_then(Value::as_object),
rules.weights,
);
if let Some(ttl_seconds) = score_rules
.get("probe_freshness_ttl_seconds")
.or_else(|| score_rules.get("score_probe_freshness_ttl_seconds"))
.and_then(json_u64)
.filter(|value| *value > 0)
{
rules.probe_freshness_ttl_seconds = ttl_seconds.min(7 * 24 * 3600);
}
if let Some(cap) = score_rules
.get("unschedulable_score_cap")
.or_else(|| score_rules.get("hard_state_score_cap"))
.and_then(json_f64)
.filter(|value| value.is_finite())
{
rules.unschedulable_score_cap = cap.clamp(0.0, 1.0);
}
if let Some(penalty) = score_rules
.get("probe_failure_penalty")
.and_then(json_f64)
.filter(|value| value.is_finite())
{
rules.probe_failure_penalty = penalty.clamp(0.0, 1.0);
}
if let Some(penalty) = score_rules
.get("request_failure_penalty")
.or_else(|| score_rules.get("runtime_failure_penalty"))
.and_then(json_f64)
.filter(|value| value.is_finite())
{
rules.request_failure_penalty = penalty.clamp(0.0, 1.0);
}
if let Some(threshold) = score_rules
.get("probe_failure_cooldown_threshold")
.or_else(|| score_rules.get("probe_failure_hard_state_threshold"))
.and_then(json_u64)
{
rules.probe_failure_cooldown_threshold = threshold.min(100);
}
}
rules.effective()
}
fn normalize_pool_preset_mode(preset: &str, raw_mode: Option<&Value>) -> Option<String> {
match preset {
"free_first" | "team_first" | "plus_first" | "pro_first" => {
@@ -270,6 +388,10 @@ pub(crate) fn admin_provider_pool_config_from_config_value(
health_policy_enabled: true,
probing_enabled: false,
probing_interval_minutes: 10,
probe_concurrency: 4,
score_top_n: 128,
score_fallback_scan_limit: 1024,
score_rules: PoolMemberScoreRules::default(),
stream_timeout_threshold: 3,
stream_timeout_window_seconds: 1800,
stream_timeout_cooldown_seconds: 300,
@@ -278,6 +400,7 @@ pub(crate) fn admin_provider_pool_config_from_config_value(
let scheduling_presets = parse_pool_scheduling_presets(pool_advanced);
let unschedulable_rules = parse_pool_unschedulable_rules(pool_advanced);
let score_rules = parse_pool_score_rules(pool_advanced);
Some(AdminProviderPoolConfig {
lru_enabled: admin_provider_pool_lru_enabled(&scheduling_presets),
@@ -333,6 +456,25 @@ pub(crate) fn admin_provider_pool_config_from_config_value(
.filter(|value| *value > 0)
.map(|value| value.min(1440))
.unwrap_or(10),
probe_concurrency: pool_advanced
.get("probe_concurrency")
.and_then(json_u64)
.filter(|value| *value > 0)
.map(|value| value.min(64))
.unwrap_or(4),
score_top_n: pool_advanced
.get("score_top_n")
.and_then(json_u64)
.filter(|value| *value > 0)
.map(|value| value.min(4096))
.unwrap_or(128),
score_fallback_scan_limit: pool_advanced
.get("score_fallback_scan_limit")
.and_then(json_u64)
.filter(|value| *value > 0)
.map(|value| value.min(50_000))
.unwrap_or(1024),
score_rules,
stream_timeout_threshold: pool_advanced
.get("stream_timeout_threshold")
.and_then(json_u64)
@@ -405,6 +547,24 @@ mod tests {
"health_policy_enabled": false,
"probing_enabled": true,
"probing_interval_minutes": 20,
"probe_concurrency": 6,
"score_top_n": 256,
"score_fallback_scan_limit": 2048,
"score_rules": {
"weights": {
"manual_priority": 0.4,
"health": 0.2,
"probe_freshness": 0.2,
"quota_remaining": 0.1,
"latency": 0.05,
"cost_lru": 0.05
},
"probe_freshness_ttl_seconds": 1200,
"unschedulable_score_cap": 0.03,
"probe_failure_penalty": 0.08,
"request_failure_penalty": 0.01,
"probe_failure_cooldown_threshold": 2
},
"stream_timeout_threshold": 4,
"stream_timeout_window_seconds": 900,
"stream_timeout_cooldown_seconds": 180
@@ -424,6 +584,16 @@ mod tests {
assert!(!config.health_policy_enabled);
assert!(config.probing_enabled);
assert_eq!(config.probing_interval_minutes, 20);
assert_eq!(config.probe_concurrency, 6);
assert_eq!(config.score_top_n, 256);
assert_eq!(config.score_fallback_scan_limit, 2048);
assert_eq!(config.score_rules.weights.manual_priority, 0.4);
assert_eq!(config.score_rules.weights.health, 0.2);
assert_eq!(config.score_rules.probe_freshness_ttl_seconds, 1200);
assert_eq!(config.score_rules.unschedulable_score_cap, 0.03);
assert_eq!(config.score_rules.probe_failure_penalty, 0.08);
assert_eq!(config.score_rules.request_failure_penalty, 0.01);
assert_eq!(config.score_rules.probe_failure_cooldown_threshold, 2);
assert_eq!(config.stream_timeout_threshold, 4);
assert_eq!(config.stream_timeout_window_seconds, 900);
assert_eq!(config.stream_timeout_cooldown_seconds, 180);
@@ -462,6 +632,27 @@ mod tests {
assert_eq!(config.sticky_session_ttl_seconds, 0);
}
#[test]
fn parses_legacy_pool_score_weights_from_pool_advanced() {
let config = admin_provider_pool_config_from_config_value(Some(&json!({
"pool_advanced": {
"scoring_weights": {
"manual_priority": 0,
"health": 2,
"probe": 1,
"quota_remaining": 0,
"latency": 0,
"cost_remaining": 1
}
}
})))
.expect("pool config should parse");
assert_eq!(config.score_rules.weights.health, 0.5);
assert_eq!(config.score_rules.weights.probe_freshness, 0.25);
assert_eq!(config.score_rules.weights.cost_lru, 0.25);
}
#[test]
fn parses_pool_config_from_generic_config_value() {
let config = admin_provider_pool_config_from_config_value(Some(&json!({

View File

@@ -14,10 +14,6 @@ pub(super) fn pool_cooldown_key(provider_id: &str, key_id: &str) -> String {
format!("ap:{provider_id}:cooldown:{key_id}")
}
pub(super) fn pool_lease_key(provider_id: &str, key_id: &str) -> String {
format!("ap:{provider_id}:lease:{key_id}")
}
pub(super) fn pool_cooldown_index_key(provider_id: &str) -> String {
format!("ap:{provider_id}:cooldown_idx")
}

View File

@@ -1,23 +1,4 @@
use super::keys::pool_lease_key;
use aether_runtime_state::{DataLayerError, RuntimeLockLease, RuntimeState};
use std::time::Duration;
pub(crate) const ADMIN_PROVIDER_POOL_KEY_LEASE_TTL_MS: u64 = 15 * 60 * 1000;
pub(crate) async fn try_claim_admin_provider_pool_key(
runtime: &RuntimeState,
provider_id: &str,
key_id: &str,
owner: &str,
) -> Result<Option<RuntimeLockLease>, DataLayerError> {
runtime
.lock_try_acquire(
&pool_lease_key(provider_id, key_id),
owner,
Duration::from_millis(ADMIN_PROVIDER_POOL_KEY_LEASE_TTL_MS),
)
.await
}
pub(crate) async fn release_admin_provider_pool_key_lease(
runtime: &RuntimeState,

View File

@@ -5,16 +5,14 @@ mod reads;
mod status;
mod writes;
pub(crate) use self::leases::{
release_admin_provider_pool_key_lease, try_claim_admin_provider_pool_key,
ADMIN_PROVIDER_POOL_KEY_LEASE_TTL_MS,
};
pub(crate) use self::leases::release_admin_provider_pool_key_lease;
pub(crate) use self::mutations::{
clear_admin_provider_pool_cooldown, reset_admin_provider_pool_cost,
};
pub(crate) use self::reads::{
read_admin_provider_pool_cooldown_count, read_admin_provider_pool_cooldown_counts,
read_admin_provider_pool_cooldown_key_ids, read_admin_provider_pool_runtime_state,
read_admin_provider_pool_cooldown_key_ids, read_admin_provider_pool_key_cooldown_reason,
read_admin_provider_pool_runtime_state,
};
pub(crate) use self::status::build_admin_provider_pool_status_payload;
pub(crate) use self::writes::{

View File

@@ -7,7 +7,7 @@ use crate::handlers::admin::provider::pool::config::admin_provider_pool_cache_af
use crate::handlers::admin::provider::shared::support::{
AdminProviderPoolConfig, AdminProviderPoolRuntimeState,
};
use aether_runtime_state::RuntimeState;
use aether_runtime_state::{DataLayerError, RuntimeState};
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
use tracing::warn;
@@ -202,3 +202,13 @@ pub(crate) async fn read_admin_provider_pool_cooldown_key_ids(
.await
.unwrap_or_default()
}
pub(crate) async fn read_admin_provider_pool_key_cooldown_reason(
runtime: &RuntimeState,
provider_id: &str,
key_id: &str,
) -> Result<Option<String>, DataLayerError> {
runtime
.kv_get(&pool_cooldown_key(provider_id, key_id))
.await
}

View File

@@ -615,6 +615,10 @@ mod tests {
health_policy_enabled: true,
probing_enabled: false,
probing_interval_minutes: 10,
probe_concurrency: 4,
score_top_n: 128,
score_fallback_scan_limit: 1024,
score_rules: aether_ai_serving::PoolMemberScoreRules::default(),
stream_timeout_threshold: 3,
stream_timeout_window_seconds: 1800,
stream_timeout_cooldown_seconds: 300,

View File

@@ -24,6 +24,8 @@ mod read_overview;
mod read_presets;
#[path = "read_routes/resolve_selection.rs"]
mod read_resolve_selection;
#[path = "read_routes/scores.rs"]
mod read_scores;
pub(crate) mod selection;
mod support;
@@ -34,11 +36,11 @@ pub(crate) use self::batch_shared::{
AdminPoolBatchImportRequest,
};
pub(crate) use self::support::{
admin_pool_provider_id_from_path, parse_admin_pool_key_sort, parse_admin_pool_page,
parse_admin_pool_page_size, parse_admin_pool_quick_selectors, parse_admin_pool_search,
parse_admin_pool_status_filter, AdminPoolKeySort, AdminPoolKeySortDirection,
AdminPoolKeySortField, AdminPoolResolveSelectionRequest,
ADMIN_POOL_BANNED_KEY_CLEANUP_EMPTY_MESSAGE,
admin_pool_provider_id_from_path, admin_pool_provider_id_from_scores_path,
parse_admin_pool_key_sort, parse_admin_pool_page, parse_admin_pool_page_size,
parse_admin_pool_quick_selectors, parse_admin_pool_search, parse_admin_pool_status_filter,
AdminPoolKeySort, AdminPoolKeySortDirection, AdminPoolKeySortField,
AdminPoolResolveSelectionRequest, ADMIN_POOL_BANNED_KEY_CLEANUP_EMPTY_MESSAGE,
ADMIN_POOL_PROVIDER_CATALOG_READER_UNAVAILABLE_DETAIL,
ADMIN_POOL_PROVIDER_CATALOG_WRITER_UNAVAILABLE_DETAIL,
};
@@ -103,6 +105,11 @@ pub(crate) async fn maybe_build_local_admin_pool_response(
read_keys::build_admin_pool_list_keys_response(state, request_context).await?,
));
}
Some("scores") => {
return Ok(Some(
read_scores::build_admin_pool_scores_response(state, request_context).await?,
));
}
Some("resolve_selection") => {
return Ok(Some(
read_resolve_selection::build_admin_pool_resolve_selection_response(

View File

@@ -6,6 +6,7 @@ use crate::handlers::admin::shared::{provider_key_status_snapshot_payload, unix_
use crate::provider_key_auth::{provider_key_auth_semantics, provider_key_effective_api_formats};
use aether_admin::provider::pool as admin_provider_pool_pure;
use aether_admin::provider::quota as admin_provider_quota_pure;
use aether_data_contracts::repository::pool_scores::StoredPoolMemberScore;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
};
@@ -906,6 +907,7 @@ pub(super) fn build_admin_pool_key_payload(
key: &StoredProviderCatalogKey,
runtime: &AdminProviderPoolRuntimeState,
pool_config: Option<AdminProviderPoolConfig>,
pool_score: Option<&StoredPoolMemberScore>,
codex_cycle_usage_by_code: Option<&BTreeMap<String, StoredProviderApiKeyWindowUsageSummary>>,
now_unix_secs: u64,
) -> serde_json::Value {
@@ -1087,6 +1089,34 @@ pub(super) fn build_admin_pool_key_payload(
payload.insert("status_snapshot".to_string(), status_snapshot);
payload.insert("quota_updated_at".to_string(), json!(quota_updated_at));
payload.insert("health_score".to_string(), json!(health_score));
payload.insert(
"pool_score".to_string(),
pool_score
.map(|score| {
json!({
"id": score.id.clone(),
"capability": score.capability.clone(),
"scope_kind": score.scope_kind.clone(),
"scope_id": score.scope_id.clone(),
"score": score.score,
"hard_state": score.hard_state.as_database(),
"score_version": score.score_version,
"score_reason": score.score_reason.clone(),
"last_ranked_at": score.last_ranked_at,
"last_scheduled_at": score.last_scheduled_at,
"last_success_at": score.last_success_at,
"last_failure_at": score.last_failure_at,
"failure_count": score.failure_count,
"last_probe_attempt_at": score.last_probe_attempt_at,
"last_probe_success_at": score.last_probe_success_at,
"last_probe_failure_at": score.last_probe_failure_at,
"probe_failure_count": score.probe_failure_count,
"probe_status": score.probe_status.as_database(),
"updated_at": score.updated_at,
})
})
.unwrap_or(serde_json::Value::Null),
);
payload.insert(
"circuit_breaker_open".to_string(),
json!(circuit_breaker_open),

View File

@@ -7,9 +7,13 @@ use super::{
AdminPoolKeySortField, AdminProviderPoolRuntimeState, ProviderCatalogKeyListOrder,
ProviderCatalogKeyListQuery, ADMIN_POOL_PROVIDER_CATALOG_READER_UNAVAILABLE_DETAIL,
};
use crate::ai_serving::{provider_key_pool_score_id, provider_key_pool_score_scope};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError;
use aether_admin::provider::pool as admin_provider_pool_pure;
use aether_data_contracts::repository::pool_scores::{
GetPoolMemberScoresByIdsQuery, PoolMemberIdentity, StoredPoolMemberScore,
};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use aether_data_contracts::repository::usage::{
ProviderApiKeyWindowUsageRequest, StoredProviderApiKeyWindowUsageSummary,
@@ -56,6 +60,36 @@ fn admin_pool_current_unix_secs() -> u64 {
.unwrap_or(0)
}
async fn read_admin_pool_scores_by_key_id(
state: &AdminAppState<'_>,
provider_id: &str,
key_ids: &[String],
) -> Result<BTreeMap<String, StoredPoolMemberScore>, GatewayError> {
if key_ids.is_empty() {
return Ok(BTreeMap::new());
}
let score_scope = provider_key_pool_score_scope();
let score_ids = key_ids
.iter()
.map(|key_id| {
let identity =
PoolMemberIdentity::provider_api_key(provider_id.to_string(), key_id.clone());
provider_key_pool_score_id(&identity, &score_scope)
})
.collect::<Vec<_>>();
let scores = state
.app()
.data
.get_pool_member_scores_by_ids(&GetPoolMemberScoresByIdsQuery { ids: score_ids })
.await
.map_err(|err| GatewayError::Internal(format!("{err:?}")))?;
Ok(scores
.into_iter()
.map(|score| (score.member_id.clone(), score))
.collect::<BTreeMap<_, _>>())
}
fn admin_pool_codex_cycle_usage_request(
key: &StoredProviderCatalogKey,
window: &serde_json::Map<String, serde_json::Value>,
@@ -379,6 +413,9 @@ pub(super) async fn build_admin_pool_list_keys_response(
};
let key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>();
let pool_scores_by_key_id = read_admin_pool_scores_by_key_id(state, &provider.id, &key_ids)
.await
.unwrap_or_default();
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
@@ -414,6 +451,7 @@ pub(super) async fn build_admin_pool_list_keys_response(
&key,
&runtime,
pool_config.clone(),
pool_scores_by_key_id.get(&key.id),
codex_cycle_usage_by_key.get(&key.id),
now_unix_secs,
)

View File

@@ -0,0 +1,193 @@
use super::{
admin_pool_provider_id_from_scores_path, build_admin_pool_error_response,
parse_admin_pool_page, parse_admin_pool_page_size,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::shared::query_param_value;
use crate::GatewayError;
use aether_data_contracts::repository::pool_scores::{
ListPoolMemberScoresQuery, PoolMemberHardState, PoolMemberProbeStatus,
POOL_KIND_PROVIDER_KEY_POOL, POOL_SCORE_CAPABILITY_ACCOUNT, POOL_SCORE_SCOPE_KIND_ACCOUNT,
};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::collections::BTreeMap;
pub(super) async fn build_admin_pool_scores_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Response<Body>, GatewayError> {
let Some(provider_id) = admin_pool_provider_id_from_scores_path(request_context.path()) else {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"provider_id 无效",
));
};
let query = request_context.query_string();
let page = match parse_admin_pool_page(query) {
Ok(value) => value,
Err(message) => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
message,
));
}
};
let page_size = match parse_admin_pool_page_size(query) {
Ok(value) => value.min(500),
Err(message) => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
message,
));
}
};
let offset = page.saturating_sub(1).saturating_mul(page_size);
let hard_states = match parse_hard_state_filter(query) {
Ok(value) => value,
Err(message) => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
message,
));
}
};
let probe_statuses = match parse_probe_status_filter(query) {
Ok(value) => value,
Err(message) => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
message,
));
}
};
let scores = state
.app()
.data
.list_pool_member_scores(&ListPoolMemberScoresQuery {
pool_kind: POOL_KIND_PROVIDER_KEY_POOL.to_string(),
pool_id: provider_id.clone(),
capability: Some(POOL_SCORE_CAPABILITY_ACCOUNT.to_string()),
scope_kind: Some(POOL_SCORE_SCOPE_KIND_ACCOUNT.to_string()),
scope_id: None,
hard_states,
probe_statuses,
offset,
limit: page_size,
})
.await
.map_err(|err| GatewayError::Internal(format!("{err:?}")))?;
let key_ids = scores
.iter()
.map(|score| score.member_id.clone())
.collect::<Vec<_>>();
let keys = state
.app()
.read_provider_catalog_keys_by_ids(&key_ids)
.await
.unwrap_or_default()
.into_iter()
.map(|key| (key.id.clone(), key))
.collect::<BTreeMap<_, _>>();
let items = scores
.into_iter()
.map(|score| {
let key = keys.get(&score.member_id);
json!({
"id": score.id,
"pool_kind": score.pool_kind,
"pool_id": score.pool_id,
"member_kind": score.member_kind,
"member_id": score.member_id,
"capability": score.capability,
"scope_kind": score.scope_kind,
"scope_id": score.scope_id,
"score": score.score,
"hard_state": score.hard_state.as_database(),
"score_version": score.score_version,
"score_reason": score.score_reason,
"last_ranked_at": score.last_ranked_at,
"last_scheduled_at": score.last_scheduled_at,
"last_success_at": score.last_success_at,
"last_failure_at": score.last_failure_at,
"failure_count": score.failure_count,
"last_probe_attempt_at": score.last_probe_attempt_at,
"last_probe_success_at": score.last_probe_success_at,
"last_probe_failure_at": score.last_probe_failure_at,
"probe_failure_count": score.probe_failure_count,
"probe_status": score.probe_status.as_database(),
"updated_at": score.updated_at,
"key": key.map(|key| json!({
"id": key.id,
"name": key.name,
"auth_type": key.auth_type,
"is_active": key.is_active,
"internal_priority": key.internal_priority,
"last_used_at": key.last_used_at_unix_secs,
}))
})
})
.collect::<Vec<_>>();
Ok(Json(json!({
"provider_id": provider_id,
"page": page,
"page_size": page_size,
"filters": {
"api_format": serde_json::Value::Null,
"model_id": serde_json::Value::Null,
"hard_state": query_param_value(query, "hard_state"),
"probe_status": query_param_value(query, "probe_status")
},
"items": items
}))
.into_response())
}
fn parse_hard_state_filter(query: Option<&str>) -> Result<Vec<PoolMemberHardState>, String> {
let Some(raw) = query_param_value(query, "hard_state") else {
return Ok(Vec::new());
};
raw.split(',')
.map(|value| match value.trim() {
"available" => Ok(PoolMemberHardState::Available),
"unknown" => Ok(PoolMemberHardState::Unknown),
"cooldown" => Ok(PoolMemberHardState::Cooldown),
"quota_exhausted" => Ok(PoolMemberHardState::QuotaExhausted),
"auth_invalid" => Ok(PoolMemberHardState::AuthInvalid),
"banned" => Ok(PoolMemberHardState::Banned),
"inactive" => Ok(PoolMemberHardState::Inactive),
_ => Err("hard_state must be one of: available, unknown, cooldown, quota_exhausted, auth_invalid, banned, inactive".to_string()),
})
.collect()
}
fn parse_probe_status_filter(
query: Option<&str>,
) -> Result<Option<Vec<PoolMemberProbeStatus>>, String> {
let Some(raw) = query_param_value(query, "probe_status") else {
return Ok(None);
};
raw.split(',')
.map(|value| match value.trim() {
"never" => Ok(PoolMemberProbeStatus::Never),
"ok" => Ok(PoolMemberProbeStatus::Ok),
"failed" => Ok(PoolMemberProbeStatus::Failed),
"stale" => Ok(PoolMemberProbeStatus::Stale),
"in_progress" => Ok(PoolMemberProbeStatus::InProgress),
_ => Err(
"probe_status must be one of: never, ok, failed, stale, in_progress".to_string(),
),
})
.collect::<Result<Vec<_>, _>>()
.map(Some)
}

View File

@@ -149,6 +149,18 @@ pub(crate) fn admin_pool_provider_id_from_path(request_path: &str) -> Option<Str
}
}
pub(crate) fn admin_pool_provider_id_from_scores_path(request_path: &str) -> Option<String> {
let raw = request_path.strip_prefix("/api/admin/pool/")?;
let mut segments = raw.split('/');
let provider_id = segments.next()?.trim();
let scores_segment = segments.next()?.trim_end_matches('/').trim();
if provider_id.is_empty() || scores_segment != "scores" {
None
} else {
Some(provider_id.to_string())
}
}
pub(crate) fn is_admin_pool_route(request_context: &AdminRequestContext<'_>) -> bool {
let normalized_path = request_context.path().trim_end_matches('/');
let path = if normalized_path.is_empty() {
@@ -164,6 +176,10 @@ pub(crate) fn is_admin_pool_route(request_context: &AdminRequestContext<'_>) ->
&& path.starts_with("/api/admin/pool/")
&& path.ends_with("/keys")
&& path.matches('/').count() == 5)
|| (request_context.method() == http::Method::GET
&& path.starts_with("/api/admin/pool/")
&& path.ends_with("/scores")
&& path.matches('/').count() == 5)
|| (request_context.method() == http::Method::POST
&& path.starts_with("/api/admin/pool/")
&& path.ends_with("/keys/batch-import")

View File

@@ -1,5 +1,6 @@
use crate::handlers::admin::request::AdminAppState;
use crate::LocalProviderDeleteTaskState;
use aether_ai_serving::PoolMemberScoreRules;
use serde_json::json;
use std::collections::BTreeMap;
@@ -39,6 +40,10 @@ pub(crate) struct AdminProviderPoolConfig {
pub(crate) health_policy_enabled: bool,
pub(crate) probing_enabled: bool,
pub(crate) probing_interval_minutes: u64,
pub(crate) probe_concurrency: u64,
pub(crate) score_top_n: u64,
pub(crate) score_fallback_scan_limit: u64,
pub(crate) score_rules: PoolMemberScoreRules,
pub(crate) stream_timeout_threshold: u64,
pub(crate) stream_timeout_window_seconds: u64,
pub(crate) stream_timeout_cooldown_seconds: u64,

View File

@@ -1,7 +1,7 @@
use super::{AdminAppState, AdminRequestContext};
use crate::{AppState, GatewayError};
use axum::body::{Body, Bytes};
use axum::http::Response;
use axum::http::{HeaderMap, Response};
pub(crate) enum AdminCancelVideoTaskError {
NotFound,
@@ -14,6 +14,7 @@ pub(crate) enum AdminCancelVideoTaskError {
pub(crate) struct AdminRouteRequest<'a> {
state: AdminAppState<'a>,
request_context: AdminRequestContext<'a>,
request_headers: &'a HeaderMap,
request_body: Option<&'a Bytes>,
}
@@ -21,11 +22,13 @@ impl<'a> AdminRouteRequest<'a> {
pub(crate) fn new(
state: &'a AppState,
request_context: &'a crate::control::GatewayPublicRequestContext,
request_headers: &'a HeaderMap,
request_body: Option<&'a Bytes>,
) -> Self {
Self {
state: AdminAppState::new(state),
request_context: AdminRequestContext::new(request_context),
request_headers,
request_body,
}
}
@@ -38,6 +41,10 @@ impl<'a> AdminRouteRequest<'a> {
self.request_context
}
pub(crate) fn request_headers(self) -> &'a HeaderMap {
self.request_headers
}
pub(crate) fn request_body(self) -> Option<&'a Bytes> {
self.request_body
}

View File

@@ -326,14 +326,14 @@ async fn refresh_imported_oauth_key_after_persist(
provider: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider,
key_id: &str,
) -> Result<(), GatewayError> {
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let Some(endpoint) =
crate::handlers::admin::provider::oauth::runtime::provider_oauth_runtime_endpoint_for_provider(
crate::handlers::admin::provider::oauth::runtime::resolve_provider_oauth_runtime_endpoints(
state,
provider,
provider.provider_type.as_str(),
&endpoints,
)
.await?
.runtime_endpoint
else {
return Ok(());
};

View File

@@ -1,5 +1,7 @@
use crate::control::{
management_token_permission_catalog_payload, normalize_assignable_management_token_permissions,
management_token_permission_catalog_payload,
management_token_permissions_cover_all_assignable_permissions,
normalize_assignable_management_token_permissions,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{query_param_optional_bool, query_param_value};
@@ -317,12 +319,17 @@ pub(crate) async fn maybe_build_local_admin_management_tokens_response(
return Ok(None);
}
if decision
let is_management_token = decision
.admin_principal
.as_ref()
.and_then(|principal| principal.management_token_id.as_deref())
.is_some()
{
.is_some();
let management_token_is_full = decision
.admin_principal
.as_ref()
.and_then(|principal| principal.management_token_permissions.as_deref())
.is_none_or(management_token_permissions_cover_all_assignable_permissions);
if is_management_token && !management_token_is_full {
return Ok(Some(
(
http::StatusCode::FORBIDDEN,

View File

@@ -30,6 +30,7 @@ pub(super) async fn maybe_build_local_internal_proxy_response(
pub(super) async fn maybe_build_local_admin_proxy_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_headers: &http::HeaderMap,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.control_decision.as_ref() else {
@@ -49,6 +50,7 @@ pub(super) async fn maybe_build_local_admin_proxy_response(
admin_api::maybe_build_local_admin_response(admin_api::AdminRouteRequest::new(
state,
request_context,
request_headers,
request_body,
))
.await

View File

@@ -874,9 +874,13 @@ pub(crate) async fn proxy_request(
request_permit.take(),
));
}
if let Some(response) =
maybe_build_local_admin_proxy_response(&state, &request_context, local_proxy_body.as_ref())
.await?
if let Some(response) = maybe_build_local_admin_proxy_response(
&state,
&request_context,
&parts.headers,
local_proxy_body.as_ref(),
)
.await?
{
let execution_path =
resolve_local_proxy_execution_path(&response, EXECUTION_PATH_PUBLIC_PROXY_PASSTHROUGH);

View File

@@ -20,6 +20,7 @@ pub(crate) use self::system_modules_helpers::{
};
pub(crate) use self::support::{
build_unhandled_public_support_response, matches_model_mapping_for_models,
maybe_build_local_admin_announcements_response, maybe_build_local_public_support_response,
build_api_key_install_session_response, build_unhandled_public_support_response,
matches_model_mapping_for_models, maybe_build_local_admin_announcements_response,
maybe_build_local_public_support_response, CreateApiKeyInstallSessionRequest,
};

View File

@@ -66,6 +66,9 @@ use self::support_auth::{
};
use self::support_billing::maybe_build_local_billing_response;
use self::support_dashboard::maybe_build_local_dashboard_response;
pub(crate) use self::support_install::{
build_api_key_install_session_response, CreateApiKeyInstallSessionRequest,
};
use self::support_install::{
handle_users_me_api_key_install_session_create, maybe_build_local_install_response,
users_me_api_key_install_sessions_path_matches,

View File

@@ -17,7 +17,7 @@ const INSTALL_SESSION_KEY_PREFIX: &str = "install:session:";
#[derive(Debug, Clone, Copy, Deserialize, Serialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
enum InstallTargetCli {
pub(crate) enum InstallTargetCli {
ClaudeCode,
CodexCli,
GeminiCli,
@@ -25,7 +25,7 @@ enum InstallTargetCli {
#[derive(Debug, Clone, Copy, Deserialize, Serialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
enum InstallTargetSystem {
pub(crate) enum InstallTargetSystem {
Macos,
Linux,
Windows,
@@ -33,9 +33,9 @@ enum InstallTargetSystem {
}
#[derive(Debug, Deserialize)]
struct UsersMeCreateInstallSessionRequest {
target_cli: InstallTargetCli,
target_system: InstallTargetSystem,
pub(crate) struct CreateApiKeyInstallSessionRequest {
pub(crate) target_cli: InstallTargetCli,
pub(crate) target_system: InstallTargetSystem,
}
#[derive(Debug, Serialize, Deserialize)]
@@ -240,20 +240,60 @@ PY
;;
codex_cli)
mkdir -p "$HOME/.codex"
cat > "$HOME/.codex/auth.json" <<EOF
{{"OPENAI_API_KEY":"$AETHER_API_KEY"}}
EOF
cat > "$HOME/.codex/config.toml" <<EOF
# Managed by Aether
model_provider = "aether"
python3 - "$HOME/.codex/config.toml" "$AETHER_BASE_URL" "$AETHER_API_KEY" <<'PY'
import pathlib, re, sys
[model_providers.aether]
name = "Aether"
base_url = "$AETHER_BASE_URL/v1"
env_key = "OPENAI_API_KEY"
wire_api = "chat"
EOF
chmod 600 "$HOME/.codex/auth.json" "$HOME/.codex/config.toml" 2>/dev/null || true
path = pathlib.Path(sys.argv[1])
base_url = sys.argv[2].rstrip('/') + '/v1'
api_key = sys.argv[3]
text = path.read_text() if path.exists() else ''
lines = text.splitlines()
def quote_toml(value: str) -> str:
return '"' + value.replace('\\', '\\\\').replace('"', '\\"') + '"'
result = []
in_aether = False
top_model_provider_set = False
seen_section = False
for line in lines:
stripped = line.strip()
if re.match(r'^\[.*\]$', stripped):
seen_section = True
in_aether = stripped == '[model_providers.aether]'
if in_aether:
continue
if in_aether:
continue
if not seen_section and re.match(r'^model_provider\s*=', stripped):
if not top_model_provider_set:
result.append('model_provider = "aether"')
top_model_provider_set = True
continue
result.append(line)
if not top_model_provider_set:
insert_at = next((idx for idx, line in enumerate(result) if line.strip().startswith('[')), len(result))
while insert_at > 0 and result[insert_at - 1].strip() == '':
insert_at -= 1
result[insert_at:insert_at] = ['model_provider = "aether"', '']
while result and result[-1].strip() == '':
result.pop()
if result:
result.append('')
result.extend([
'# Managed by Aether',
'[model_providers.aether]',
'name = "Aether"',
f'base_url = {{quote_toml(base_url)}}',
'wire_api = "responses"',
'requires_openai_auth = false',
f'experimental_bearer_token = {{quote_toml(api_key)}}',
])
path.write_text('\n'.join(result) + '\n')
PY
chmod 600 "$HOME/.codex/config.toml" 2>/dev/null || true
;;
gemini_cli)
mkdir -p "$HOME/.gemini"
@@ -329,8 +369,51 @@ if ($TargetCli -eq 'claude_code') {{
$Data | ConvertTo-Json -Depth 8 | Set-Content $Path -Encoding UTF8
}} elseif ($TargetCli -eq 'codex_cli') {{
$Dir = Join-Path $HomeDir '.codex'; New-Item -ItemType Directory -Force -Path $Dir | Out-Null
Set-Content (Join-Path $Dir 'auth.json') -Value (@{{ OPENAI_API_KEY = $AetherApiKey }} | ConvertTo-Json) -Encoding UTF8
Set-Content (Join-Path $Dir 'config.toml') -Value "# Managed by Aether`nmodel_provider = \"aether\"`n`n[model_providers.aether]`nname = \"Aether\"`nbase_url = \"$AetherBaseUrl/v1\"`nenv_key = \"OPENAI_API_KEY\"`nwire_api = \"chat\"`n" -Encoding UTF8
$Path = Join-Path $Dir 'config.toml'
$Text = if (Test-Path $Path) {{ Get-Content $Path -Raw }} else {{ '' }}
$Lines = if ($Text.Length -gt 0) {{ $Text -split "`r?`n" }} else {{ @() }}
$Result = New-Object System.Collections.Generic.List[string]
$InAether = $false
$TopModelProviderSet = $false
$SeenSection = $false
foreach ($Line in $Lines) {{
$Stripped = $Line.Trim()
if ($Stripped -match '^\[.*\]$') {{
$SeenSection = $true
$InAether = $Stripped -eq '[model_providers.aether]'
if ($InAether) {{ continue }}
}}
if ($InAether) {{ continue }}
if (-not $SeenSection -and $Stripped -match '^model_provider\s*=') {{
if (-not $TopModelProviderSet) {{
$Result.Add('model_provider = "aether"')
$TopModelProviderSet = $true
}}
continue
}}
$Result.Add($Line)
}}
if (-not $TopModelProviderSet) {{
$InsertAt = $Result.Count
for ($Index = 0; $Index -lt $Result.Count; $Index++) {{
if ($Result[$Index].Trim().StartsWith('[')) {{ $InsertAt = $Index; break }}
}}
while ($InsertAt -gt 0 -and $Result[$InsertAt - 1].Trim() -eq '') {{ $InsertAt-- }}
$Result.Insert($InsertAt, '')
$Result.Insert($InsertAt, 'model_provider = "aether"')
}}
while ($Result.Count -gt 0 -and $Result[$Result.Count - 1].Trim() -eq '') {{ $Result.RemoveAt($Result.Count - 1) }}
if ($Result.Count -gt 0) {{ $Result.Add('') }}
$EscapedBaseUrl = ($AetherBaseUrl.TrimEnd('/') + '/v1').Replace('\', '\\').Replace('"', '\"')
$EscapedApiKey = $AetherApiKey.Replace('\', '\\').Replace('"', '\"')
$Result.Add('# Managed by Aether')
$Result.Add('[model_providers.aether]')
$Result.Add('name = "Aether"')
$Result.Add("base_url = `"$EscapedBaseUrl`"")
$Result.Add('wire_api = "responses"')
$Result.Add('requires_openai_auth = false')
$Result.Add("experimental_bearer_token = `"$EscapedApiKey`"")
Set-Content -Path $Path -Value (($Result -join "`n") + "`n") -Encoding UTF8
}} elseif ($TargetCli -eq 'gemini_cli') {{
$Dir = Join-Path $HomeDir '.gemini'; New-Item -ItemType Directory -Force -Path $Dir | Out-Null
Set-Content (Join-Path $Dir '.env') -Value "GEMINI_API_KEY=$AetherApiKey`nGOOGLE_API_KEY=$AetherApiKey`nGOOGLE_GEMINI_BASE_URL=$AetherBaseUrl`nAETHER_BASE_URL=$AetherBaseUrl`n" -Encoding UTF8
@@ -375,7 +458,7 @@ pub(super) async fn handle_users_me_api_key_install_session_create(
let Some(request_body) = request_body else {
return build_auth_error_response(http::StatusCode::BAD_REQUEST, "请求数据验证失败", false);
};
let payload = match serde_json::from_slice::<UsersMeCreateInstallSessionRequest>(request_body) {
let payload = match serde_json::from_slice::<CreateApiKeyInstallSessionRequest>(request_body) {
Ok(value) => value,
Err(_) => {
return build_auth_error_response(
@@ -426,11 +509,32 @@ pub(super) async fn handle_users_me_api_key_install_session_create(
);
};
build_api_key_install_session_response(
state,
request_context,
headers,
record.api_key_id.clone(),
record.name.unwrap_or_else(|| "API Key".to_string()),
api_key,
payload,
)
.await
}
pub(crate) async fn build_api_key_install_session_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
headers: &http::HeaderMap,
api_key_id: String,
api_key_name: String,
api_key: String,
payload: CreateApiKeyInstallSessionRequest,
) -> Response<Body> {
let code = generate_install_code();
let expires_at_unix_secs = unix_secs_now().saturating_add(INSTALL_SESSION_TTL_SECS);
let session = StoredInstallSession {
api_key_id: record.api_key_id.clone(),
api_key_name: record.name.unwrap_or_else(|| "API Key".to_string()),
api_key_id,
api_key_name,
api_key,
base_url: base_url_from_request(headers, request_context),
target_cli: payload.target_cli,
@@ -559,3 +663,49 @@ pub(super) async fn maybe_build_local_install_response(
);
Some(response)
}
#[cfg(test)]
mod tests {
use super::*;
fn test_session(target_cli: InstallTargetCli) -> StoredInstallSession {
StoredInstallSession {
api_key_id: "key-1".to_string(),
api_key_name: "Key 1".to_string(),
api_key: "sk-test".to_string(),
base_url: "http://localhost:8084".to_string(),
target_cli,
target_system: InstallTargetSystem::Linux,
expires_at_unix_secs: u64::MAX,
}
}
#[test]
fn codex_unix_script_preserves_config_and_uses_responses_bearer_token() {
let script = build_unix_script(&test_session(InstallTargetCli::CodexCli));
assert!(script.contains("path.read_text() if path.exists() else ''"));
assert!(script.contains("stripped == '[model_providers.aether]'"));
assert!(script.contains("model_provider = \"aether\""));
assert!(script.contains("wire_api = \"responses\""));
assert!(script.contains("requires_openai_auth = false"));
assert!(script.contains("experimental_bearer_token ="));
assert!(!script.contains("wire_api = \"chat\""));
assert!(!script.contains("cat > \"$HOME/.codex/config.toml\""));
assert!(!script.contains("auth.json"));
}
#[test]
fn codex_powershell_script_preserves_config_and_uses_responses_bearer_token() {
let script = build_powershell_script(&test_session(InstallTargetCli::CodexCli));
assert!(script.contains("Get-Content $Path -Raw"));
assert!(script.contains("$Stripped -eq '[model_providers.aether]'"));
assert!(script.contains("model_provider = \"aether\""));
assert!(script.contains("wire_api = \"responses\""));
assert!(script.contains("requires_openai_auth = false"));
assert!(script.contains("experimental_bearer_token ="));
assert!(!script.contains("wire_api = \"chat\""));
assert!(!script.contains("auth.json"));
}
}

View File

@@ -2,10 +2,10 @@ pub(crate) use super::super::admin::provider::pool::config::{
admin_provider_pool_cache_affinity_enabled, admin_provider_pool_config_from_config_value,
};
pub(crate) use super::super::admin::provider::pool::runtime::{
admin_provider_pool_key_circuit_breaker_reason, read_admin_provider_pool_runtime_state,
record_admin_provider_pool_error, record_admin_provider_pool_stream_timeout,
record_admin_provider_pool_success, release_admin_provider_pool_key_lease,
try_claim_admin_provider_pool_key, ADMIN_PROVIDER_POOL_KEY_LEASE_TTL_MS,
admin_provider_pool_key_circuit_breaker_reason, read_admin_provider_pool_key_cooldown_reason,
read_admin_provider_pool_runtime_state, record_admin_provider_pool_error,
record_admin_provider_pool_stream_timeout, record_admin_provider_pool_success,
release_admin_provider_pool_key_lease,
};
pub(crate) use super::super::admin::provider::shared::support::{
AdminProviderPoolConfig, AdminProviderPoolRuntimeState, AdminProviderPoolSchedulingPreset,

View File

@@ -311,7 +311,11 @@ pub(crate) fn admin_proxy_local_requires_buffered_body(
| (Some("payments_manage"), http::Method::POST, Some("credit_order"))
| (Some("payments_manage"), http::Method::POST, Some("create_redeem_code_batch"))
| (Some("payments_manage"), http::Method::POST, Some("delete_redeem_code_batch"))
| (Some("api_keys_manage"), http::Method::POST, Some("create_api_key"))
| (
Some("api_keys_manage"),
http::Method::POST,
Some("create_api_key" | "create_api_key_install_session"),
)
| (Some("api_keys_manage"), http::Method::PUT, Some("update_api_key"))
| (Some("api_keys_manage"), http::Method::PATCH, Some("toggle_api_key"))
| (Some("adaptive_manage"), http::Method::PATCH, Some("toggle_mode"))

View File

@@ -36,6 +36,7 @@ mod clock;
mod constants;
mod control;
mod data;
mod dispatch;
mod error;
mod execution_runtime;
mod executor;

View File

@@ -4,14 +4,15 @@ mod tests;
pub(crate) use runtime::{
cancel_proxy_upgrade_rollout, clear_proxy_upgrade_rollout_conflicts,
inspect_proxy_upgrade_rollout, list_admin_cleanup_run_records,
perform_oauth_token_refresh_once, perform_pool_quota_probe_once, perform_provider_checkin_once,
rebuild_admin_stats_once, record_completed_cleanup_run, record_proxy_upgrade_traffic_success,
ensure_provider_key_pool_scores_for_keys, inspect_proxy_upgrade_rollout,
list_admin_cleanup_run_records, perform_oauth_token_refresh_once,
perform_pool_quota_probe_once, perform_provider_checkin_once, rebuild_admin_stats_once,
record_completed_cleanup_run, record_proxy_upgrade_traffic_success,
restore_proxy_upgrade_rollout_skipped_nodes, retry_proxy_upgrade_rollout_node,
run_admin_system_cleanup_once, skip_proxy_upgrade_rollout_node, spawn_audit_cleanup_worker,
spawn_db_maintenance_worker, spawn_gemini_file_mapping_cleanup_worker,
spawn_oauth_token_refresh_worker, spawn_pending_cleanup_worker, spawn_pool_monitor_worker,
spawn_pool_quota_probe_worker, spawn_provider_checkin_worker,
spawn_pool_quota_probe_worker, spawn_pool_score_rebuild_worker, spawn_provider_checkin_worker,
spawn_proxy_node_metrics_cleanup_worker, spawn_proxy_node_stale_cleanup_worker,
spawn_proxy_upgrade_rollout_worker, spawn_request_candidate_cleanup_worker,
spawn_stats_aggregation_worker, spawn_stats_hourly_aggregation_worker,

View File

@@ -20,6 +20,8 @@ mod oauth_token_refresh;
mod pending_cleanup;
#[path = "runtime/pool_quota_probe.rs"]
mod pool_quota_probe;
#[path = "runtime/pool_score_rebuild.rs"]
mod pool_score_rebuild;
#[path = "runtime/provider_checkin.rs"]
mod provider_checkin;
#[path = "runtime/proxy_node_metrics_cleanup.rs"]
@@ -67,6 +69,11 @@ pub(crate) use pool_quota_probe::{
select_pool_quota_probe_key_ids, spawn_pool_quota_probe_worker, PoolQuotaProbeRunSummary,
PoolQuotaProbeWorkerConfig,
};
pub(crate) use pool_score_rebuild::{
ensure_provider_key_pool_scores_for_keys, perform_pool_score_rebuild_once,
perform_pool_score_rebuild_once_with_config, spawn_pool_score_rebuild_worker,
PoolScoreRebuildRunSummary, PoolScoreRebuildWorkerConfig,
};
pub(crate) use provider_checkin::{perform_provider_checkin_once, ProviderCheckinRunSummary};
use proxy_node_metrics_cleanup::*;
use proxy_node_staleness::*;

View File

@@ -1,24 +1,34 @@
use std::collections::BTreeMap;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_data_contracts::repository::pool_scores::{
ListPoolMemberProbeCandidatesQuery, PoolMemberHardState, PoolMemberIdentity,
PoolMemberProbeAttempt, PoolMemberProbeResult, PoolMemberProbeStatus,
POOL_KIND_PROVIDER_KEY_POOL,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_runtime_state::{RuntimeLockLease, RuntimeState};
use futures_util::{stream, StreamExt};
use serde_json::Value;
use tracing::{debug, info, warn};
use crate::admin_api::{
admin_provider_pool_config, provider_oauth_runtime_endpoint_for_provider,
provider_type_supports_quota_refresh, reconcile_admin_fixed_provider_template_endpoints,
refresh_antigravity_provider_quota_locally, refresh_chatgpt_web_provider_quota_locally,
refresh_codex_provider_quota_locally, refresh_kiro_provider_quota_locally, AdminAppState,
};
use crate::{AppState, GatewayError};
use super::pool_score_rebuild::ensure_provider_key_pool_scores_for_keys;
const POOL_QUOTA_PROBE_REDIS_PREFIX: &str = "ap:quota_probe:last";
const POOL_QUOTA_PROBE_DEFAULT_SCAN_INTERVAL_SECONDS: u64 = 60;
const POOL_QUOTA_PROBE_MIN_SCAN_INTERVAL_SECONDS: u64 = 15;
const POOL_QUOTA_PROBE_DEFAULT_MAX_KEYS_PER_PROVIDER: usize = 50;
const POOL_QUOTA_PROBE_DEFAULT_GLOBAL_CONCURRENCY: usize = 16;
const POOL_QUOTA_PROBE_PROVIDER_LOCK_TTL_MS: u64 = 30_000;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -50,6 +60,7 @@ impl PoolQuotaProbeRunSummary {
pub(crate) struct PoolQuotaProbeWorkerConfig {
pub(crate) scan_interval: Duration,
pub(crate) max_keys_per_provider: usize,
pub(crate) global_concurrency: usize,
}
impl PoolQuotaProbeWorkerConfig {
@@ -63,9 +74,15 @@ impl PoolQuotaProbeWorkerConfig {
"POOL_QUOTA_PROBE_MAX_KEYS_PER_PROVIDER",
POOL_QUOTA_PROBE_DEFAULT_MAX_KEYS_PER_PROVIDER,
);
let global_concurrency = env_usize(
"POOL_QUOTA_PROBE_GLOBAL_CONCURRENCY",
POOL_QUOTA_PROBE_DEFAULT_GLOBAL_CONCURRENCY,
)
.clamp(1, 256);
Self {
scan_interval: Duration::from_secs(scan_interval_seconds),
max_keys_per_provider,
global_concurrency,
}
}
}
@@ -92,10 +109,7 @@ fn now_unix_secs() -> u64 {
}
fn provider_supports_quota_probe(provider_type: &str) -> bool {
matches!(
provider_type.trim().to_ascii_lowercase().as_str(),
"codex" | "kiro" | "antigravity" | "chatgpt_web"
)
provider_type_supports_quota_refresh(provider_type)
}
fn json_number(value: Option<&Value>) -> Option<f64> {
@@ -172,6 +186,49 @@ pub(crate) fn select_pool_quota_probe_key_ids(
stale.into_iter().map(|(_, key_id)| key_id).collect()
}
async fn select_score_probe_key_ids(
state: &AppState,
provider_id: &str,
now_ts: u64,
interval_seconds: u64,
limit: usize,
) -> Vec<String> {
if limit == 0 {
return Vec::new();
}
let stale_before_unix_secs = now_ts.saturating_sub(interval_seconds);
let query = ListPoolMemberProbeCandidatesQuery {
pool_kind: POOL_KIND_PROVIDER_KEY_POOL.to_string(),
pool_id: provider_id.to_string(),
capability: None,
stale_before_unix_secs,
limit: limit.saturating_mul(4).max(limit),
};
let scores = match state.data.list_pool_member_probe_candidates(&query).await {
Ok(scores) => scores,
Err(err) => {
debug!(
provider_id,
error = ?err,
"gateway pool quota probe: failed to read score probe candidates"
);
return Vec::new();
}
};
let mut selected = Vec::new();
let mut seen = std::collections::BTreeSet::new();
for score in scores {
if !seen.insert(score.member_id.clone()) {
continue;
}
selected.push(score.member_id);
if selected.len() >= limit {
break;
}
}
selected
}
fn probe_stamp_key(provider_id: &str, key_id: &str) -> String {
format!("{POOL_QUOTA_PROBE_REDIS_PREFIX}:{provider_id}:{key_id}")
}
@@ -295,7 +352,28 @@ async fn select_keys_for_provider(
let key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>();
let probe_stamps = load_probe_timestamps(runtime, &provider.id, &key_ids).await;
let selected_ids = select_pool_quota_probe_key_ids(
let mut selected_ids = select_score_probe_key_ids(
state,
&provider.id,
now_ts,
interval_seconds,
max_keys_per_provider,
)
.await
.into_iter()
.filter(|key_id| key_ids.iter().any(|known_id| known_id == key_id))
.filter(|key_id| {
probe_stamps.get(key_id).is_none_or(|last_probe_ts| {
now_ts.saturating_sub(*last_probe_ts) >= interval_seconds
})
})
.collect::<Vec<_>>();
let mut selected_seen = selected_ids
.iter()
.cloned()
.collect::<std::collections::BTreeSet<_>>();
let remaining = max_keys_per_provider.saturating_sub(selected_ids.len());
let fallback_selected_ids = select_pool_quota_probe_key_ids(
&keys,
provider_type,
now_ts,
@@ -303,6 +381,14 @@ async fn select_keys_for_provider(
&probe_stamps,
max_keys_per_provider,
);
for key_id in fallback_selected_ids {
if selected_ids.len() >= max_keys_per_provider || remaining == 0 {
break;
}
if selected_seen.insert(key_id.clone()) {
selected_ids.push(key_id);
}
}
if selected_ids.is_empty() {
return Ok(Vec::new());
}
@@ -338,6 +424,37 @@ fn endpoint_for_probe(
provider_oauth_runtime_endpoint_for_provider(provider_type, endpoints)
}
async fn endpoint_for_probe_with_reconcile(
state: &AppState,
admin_state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
provider_type: &str,
endpoints_by_provider: &mut BTreeMap<String, Vec<StoredProviderCatalogEndpoint>>,
) -> Result<Option<StoredProviderCatalogEndpoint>, GatewayError> {
let endpoints = endpoints_by_provider
.get(&provider.id)
.map(Vec::as_slice)
.unwrap_or(&[]);
if let Some(endpoint) = endpoint_for_probe(provider_type, endpoints) {
return Ok(Some(endpoint));
}
if admin_state
.fixed_provider_template(&provider.provider_type)
.is_none()
{
return Ok(None);
}
reconcile_admin_fixed_provider_template_endpoints(admin_state, provider).await?;
let refreshed = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let endpoint = endpoint_for_probe(provider_type, &refreshed);
endpoints_by_provider.insert(provider.id.clone(), refreshed);
Ok(endpoint)
}
async fn refresh_provider_probe_keys(
admin_state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
@@ -381,6 +498,168 @@ fn update_summary_from_payload(
.unwrap_or(0) as usize;
}
async fn record_score_probe_results_from_payload(
state: &AppState,
provider_id: &str,
selected_key_ids: &[String],
payload: Option<&Value>,
attempted_at: u64,
) {
let mut recorded = std::collections::BTreeSet::new();
if let Some(results) = payload
.and_then(|value| value.get("results"))
.and_then(Value::as_array)
{
for item in results {
let Some(key_id) = item
.get("key_id")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
continue;
};
recorded.insert(key_id.to_string());
record_score_probe_result_for_key(
state,
provider_id,
key_id,
attempted_at,
probe_result_succeeded(item),
probe_result_hard_state(item).or_else(|| {
(!probe_result_succeeded(item)).then_some(PoolMemberHardState::Cooldown)
}),
serde_json::json!({
"last_probe": {
"source": "pool_quota_probe",
"status": item.get("status").cloned().unwrap_or(Value::Null),
"status_code": item.get("status_code").cloned().unwrap_or(Value::Null),
"message": item.get("message").cloned().unwrap_or(Value::Null),
"auto_removed": item.get("auto_removed").cloned().unwrap_or(Value::Null)
}
}),
)
.await;
}
}
for key_id in selected_key_ids {
if recorded.contains(key_id) {
continue;
}
record_score_probe_result_for_key(
state,
provider_id,
key_id,
attempted_at,
false,
Some(PoolMemberHardState::Cooldown),
serde_json::json!({
"last_probe": {
"source": "pool_quota_probe",
"status": "missing_result"
}
}),
)
.await;
}
}
async fn record_score_probe_result_for_key(
state: &AppState,
provider_id: &str,
key_id: &str,
attempted_at: u64,
succeeded: bool,
hard_state: Option<PoolMemberHardState>,
score_reason_patch: Value,
) {
let result = PoolMemberProbeResult {
identity: PoolMemberIdentity::provider_api_key(provider_id.to_string(), key_id.to_string()),
scope: None,
attempted_at,
succeeded,
hard_state,
probe_status: if succeeded {
PoolMemberProbeStatus::Ok
} else {
PoolMemberProbeStatus::Failed
},
score_reason_patch: Some(score_reason_patch),
};
if let Err(err) = state.data.record_pool_member_probe_result(result).await {
debug!(
provider_id,
key_id,
error = ?err,
"gateway pool quota probe: failed to record score probe result"
);
}
}
async fn record_score_probe_in_progress_for_key(
state: &AppState,
provider_id: &str,
key_id: &str,
attempted_at: u64,
) {
let attempt = PoolMemberProbeAttempt {
identity: PoolMemberIdentity::provider_api_key(provider_id.to_string(), key_id.to_string()),
scope: None,
attempted_at,
score_reason_patch: Some(serde_json::json!({
"last_probe": {
"source": "pool_quota_probe",
"status": "in_progress"
}
})),
};
if let Err(err) = state.data.mark_pool_member_probe_in_progress(attempt).await {
debug!(
provider_id,
key_id,
error = ?err,
"gateway pool quota probe: failed to mark score probe in progress"
);
}
}
fn probe_result_succeeded(item: &Value) -> bool {
item.get("status")
.and_then(Value::as_str)
.is_some_and(|status| status == "success")
}
fn probe_result_hard_state(item: &Value) -> Option<PoolMemberHardState> {
if probe_result_succeeded(item) {
return Some(PoolMemberHardState::Available);
}
if item
.get("auto_removed")
.and_then(Value::as_bool)
.unwrap_or(false)
{
return Some(PoolMemberHardState::Banned);
}
let status = item
.get("status")
.and_then(Value::as_str)
.unwrap_or_default()
.trim()
.to_ascii_lowercase();
match status.as_str() {
"auth_invalid" | "forbidden" => Some(PoolMemberHardState::AuthInvalid),
"workspace_deactivated" => Some(PoolMemberHardState::Banned),
"quota_exhausted" => Some(PoolMemberHardState::QuotaExhausted),
_ => match item.get("status_code").and_then(Value::as_u64) {
Some(401 | 403) => Some(PoolMemberHardState::AuthInvalid),
Some(402) => Some(PoolMemberHardState::QuotaExhausted),
Some(429 | 500..=599) => Some(PoolMemberHardState::Cooldown),
_ => None,
},
}
}
pub(crate) async fn perform_pool_quota_probe_once_with_config(
state: &AppState,
config: PoolQuotaProbeWorkerConfig,
@@ -404,11 +683,7 @@ pub(crate) async fn perform_pool_quota_probe_once_with_config(
.filter_map(|(provider, provider_type)| {
let pool_config = admin_provider_pool_config(&provider)?;
if pool_config.probing_enabled {
Some((
provider,
provider_type,
pool_config.probing_interval_minutes,
))
Some((provider, provider_type, pool_config))
} else {
None
}
@@ -441,11 +716,16 @@ pub(crate) async fn perform_pool_quota_probe_once_with_config(
..PoolQuotaProbeRunSummary::empty()
};
for (provider, provider_type, interval_minutes) in providers {
let endpoints = endpoints_by_provider
.remove(&provider.id)
.unwrap_or_default();
let Some(endpoint) = endpoint_for_probe(&provider_type, &endpoints) else {
for (provider, provider_type, pool_config) in providers {
let Some(endpoint) = endpoint_for_probe_with_reconcile(
state,
&admin_state,
&provider,
&provider_type,
&mut endpoints_by_provider,
)
.await?
else {
summary.providers_skipped += 1;
debug!(
provider_id = %provider.id,
@@ -454,7 +734,11 @@ pub(crate) async fn perform_pool_quota_probe_once_with_config(
);
continue;
};
let endpoints = endpoints_by_provider
.remove(&provider.id)
.unwrap_or_else(|| vec![endpoint.clone()]);
let interval_minutes = pool_config.probing_interval_minutes;
let interval_seconds = interval_minutes.clamp(1, 1440).saturating_mul(60);
let keys = select_keys_for_provider(
state,
@@ -474,42 +758,131 @@ pub(crate) async fn perform_pool_quota_probe_once_with_config(
summary.providers_probed += 1;
summary.selected_keys += selected_count;
let provider_short_id = provider.id.chars().take(8).collect::<String>();
match refresh_provider_probe_keys(&admin_state, &provider, &endpoint, &provider_type, keys)
.await
let selected_key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>();
let score_ensure_budget = (pool_config.score_fallback_scan_limit as usize)
.min(50_000)
.max(selected_count.min(50_000));
match ensure_provider_key_pool_scores_for_keys(
state,
&provider,
&pool_config,
&endpoints,
&keys,
now_ts,
score_ensure_budget,
)
.await
{
Ok(payload) => {
update_summary_from_payload(&mut summary, selected_count, payload.as_ref());
let probe_success = payload
.as_ref()
.and_then(|value| value.get("success"))
.and_then(Value::as_u64)
.unwrap_or(0);
let probe_failed = payload
.as_ref()
.and_then(|value| value.get("failed"))
.and_then(Value::as_u64)
.unwrap_or(0);
info!(
provider_id = %provider_short_id,
provider_type,
selected = selected_count,
success = probe_success,
failed = probe_failed,
"gateway pool quota probe completed"
Ok(upserted) if upserted > 0 => {
debug!(
provider_id = %provider.id,
key_count = selected_count,
scores_upserted = upserted,
"gateway pool quota probe: ensured score rows for selected probe keys"
);
}
Ok(_) => {}
Err(err) => {
summary.failed += selected_count;
warn!(
provider_id = %provider_short_id,
provider_type,
selected = selected_count,
provider_id = %provider.id,
key_count = selected_count,
error = ?err,
"gateway pool quota probe failed"
"gateway pool quota probe: failed to ensure score rows for selected probe keys"
);
}
}
for key_id in &selected_key_ids {
record_score_probe_in_progress_for_key(state, &provider.id, key_id, now_ts).await;
}
let provider_short_id = provider.id.chars().take(8).collect::<String>();
let probe_concurrency = pool_config.probe_concurrency.clamp(1, 64) as usize;
let probe_concurrency = probe_concurrency.min(config.global_concurrency).max(1);
let probe_results = stream::iter(keys.into_iter().map(|key| {
let key_id = key.id.clone();
let admin_state = &admin_state;
let provider = &provider;
let endpoint = &endpoint;
let provider_type = provider_type.as_str();
async move {
let result = refresh_provider_probe_keys(
admin_state,
provider,
endpoint,
provider_type,
vec![key],
)
.await;
(key_id, result)
}
}))
.buffer_unordered(probe_concurrency)
.collect::<Vec<_>>()
.await;
let mut probe_success = 0usize;
let mut probe_failed = 0usize;
for (key_id, result) in probe_results {
match result {
Ok(payload) => {
update_summary_from_payload(&mut summary, 1, payload.as_ref());
probe_success += payload
.as_ref()
.and_then(|value| value.get("success"))
.and_then(Value::as_u64)
.unwrap_or(0) as usize;
probe_failed += payload
.as_ref()
.and_then(|value| value.get("failed"))
.and_then(Value::as_u64)
.unwrap_or(0) as usize;
record_score_probe_results_from_payload(
state,
&provider.id,
std::slice::from_ref(&key_id),
payload.as_ref(),
now_ts,
)
.await;
}
Err(err) => {
summary.failed += 1;
probe_failed += 1;
record_score_probe_result_for_key(
state,
&provider.id,
&key_id,
now_ts,
false,
Some(PoolMemberHardState::Cooldown),
serde_json::json!({
"last_probe": {
"source": "pool_quota_probe",
"status": "worker_error",
"message": format!("{err:?}")
}
}),
)
.await;
warn!(
provider_id = %provider_short_id,
provider_type,
key_id,
error = ?err,
"gateway pool quota probe failed"
);
}
}
}
info!(
provider_id = %provider_short_id,
provider_type,
selected = selected_count,
success = probe_success,
failed = probe_failed,
concurrency = probe_concurrency,
"gateway pool quota probe completed"
);
}
Ok(summary)

View File

@@ -0,0 +1,418 @@
use std::collections::BTreeMap;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_data_contracts::repository::pool_scores::GetPoolMemberScoresByIdsQuery;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use tracing::{debug, info, warn};
use crate::admin_api::admin_provider_pool_config;
use crate::ai_serving::build_provider_key_pool_score_upsert;
use crate::handlers::shared::provider_pool::AdminProviderPoolConfig;
use crate::{AppState, GatewayError};
const POOL_SCORE_REBUILD_DEFAULT_INTERVAL_SECONDS: u64 = 300;
const POOL_SCORE_REBUILD_MIN_INTERVAL_SECONDS: u64 = 30;
const POOL_SCORE_REBUILD_DEFAULT_MAX_UPSERTS_PER_TICK: usize = 20_000;
const POOL_SCORE_REBUILD_PROVIDER_CURSOR_KEY: &str = "ap:pool_score_rebuild:provider_cursor";
const POOL_SCORE_REBUILD_PROVIDER_OFFSET_PREFIX: &str = "ap:pool_score_rebuild:provider_offset";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct PoolScoreRebuildRunSummary {
pub(crate) providers_checked: usize,
pub(crate) providers_scored: usize,
pub(crate) keys_seen: usize,
pub(crate) scores_upserted: usize,
}
impl PoolScoreRebuildRunSummary {
const fn empty() -> Self {
Self {
providers_checked: 0,
providers_scored: 0,
keys_seen: 0,
scores_upserted: 0,
}
}
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct PoolScoreRebuildWorkerConfig {
pub(crate) interval: Duration,
pub(crate) max_upserts_per_tick: usize,
}
impl PoolScoreRebuildWorkerConfig {
fn from_env() -> Self {
let interval_seconds = env_u64(
"POOL_SCORE_REBUILD_INTERVAL_SECONDS",
POOL_SCORE_REBUILD_DEFAULT_INTERVAL_SECONDS,
)
.max(POOL_SCORE_REBUILD_MIN_INTERVAL_SECONDS);
let max_upserts_per_tick = env_usize(
"POOL_SCORE_REBUILD_MAX_UPSERTS_PER_TICK",
POOL_SCORE_REBUILD_DEFAULT_MAX_UPSERTS_PER_TICK,
)
.max(1);
Self {
interval: Duration::from_secs(interval_seconds),
max_upserts_per_tick,
}
}
}
fn env_u64(name: &str, default_value: u64) -> u64 {
std::env::var(name)
.ok()
.and_then(|value| value.trim().parse::<u64>().ok())
.unwrap_or(default_value)
}
fn env_usize(name: &str, default_value: usize) -> usize {
std::env::var(name)
.ok()
.and_then(|value| value.trim().parse::<usize>().ok())
.unwrap_or(default_value)
}
fn now_unix_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
}
async fn load_runtime_usize(state: &AppState, key: &str) -> usize {
state
.runtime_state
.kv_get(key)
.await
.ok()
.flatten()
.and_then(|value| value.trim().parse::<usize>().ok())
.unwrap_or(0)
}
async fn store_runtime_usize(state: &AppState, key: &str, value: usize) {
if let Err(err) = state
.runtime_state
.kv_set(key, value.to_string(), None)
.await
{
debug!(
key,
error = ?err,
"gateway pool score rebuild: failed to store cursor"
);
}
}
fn provider_offset_cursor_key(provider_id: &str) -> String {
format!("{POOL_SCORE_REBUILD_PROVIDER_OFFSET_PREFIX}:{provider_id}")
}
pub(crate) async fn ensure_provider_key_pool_scores_for_keys(
state: &AppState,
provider: &StoredProviderCatalogProvider,
pool_config: &AdminProviderPoolConfig,
_endpoints: &[StoredProviderCatalogEndpoint],
keys: &[StoredProviderCatalogKey],
now_unix_secs: u64,
max_upserts: usize,
) -> Result<usize, GatewayError> {
if max_upserts == 0
|| keys.is_empty()
|| !state.data.has_pool_score_reader()
|| !state.data.has_pool_score_writer()
{
return Ok(0);
}
let keys = keys
.iter()
.filter(|key| key.is_active && key.provider_id == provider.id)
.collect::<Vec<_>>();
if keys.is_empty() {
return Ok(0);
}
let build_items = keys
.into_iter()
.take(max_upserts)
.map(|key| {
let draft = build_provider_key_pool_score_upsert(
key,
provider.provider_type.as_str(),
None,
now_unix_secs,
pool_config.score_rules,
);
(key, draft.id)
})
.collect::<Vec<_>>();
if build_items.is_empty() {
return Ok(0);
}
let existing_score_ids = state
.data
.get_pool_member_scores_by_ids(&GetPoolMemberScoresByIdsQuery {
ids: build_items
.iter()
.map(|(_, score_id)| score_id.clone())
.collect(),
})
.await
.unwrap_or_else(|err| {
debug!(
provider_id = %provider.id,
error = ?err,
"gateway pool score ensure: failed to read existing scores by id"
);
Vec::new()
})
.into_iter()
.map(|score| score.id)
.collect::<std::collections::BTreeSet<_>>();
let mut upserted = 0usize;
for (key, score_id) in &build_items {
if existing_score_ids.contains(score_id) {
continue;
}
let upsert = build_provider_key_pool_score_upsert(
key,
provider.provider_type.as_str(),
None,
now_unix_secs,
pool_config.score_rules,
);
if state
.data
.upsert_pool_member_score(upsert)
.await
.map_err(|err| GatewayError::Internal(format!("{err:?}")))?
.is_some()
{
upserted = upserted.saturating_add(1);
}
}
Ok(upserted)
}
pub(crate) async fn perform_pool_score_rebuild_once_with_config(
state: &AppState,
config: PoolScoreRebuildWorkerConfig,
) -> Result<PoolScoreRebuildRunSummary, GatewayError> {
if !state.has_provider_catalog_data_reader()
|| !state.data.has_pool_score_reader()
|| !state.data.has_pool_score_writer()
{
return Ok(PoolScoreRebuildRunSummary::empty());
}
let mut providers = state
.list_provider_catalog_providers(true)
.await?
.into_iter()
.filter_map(|provider| {
admin_provider_pool_config(&provider).map(|config| (provider, config))
})
.collect::<Vec<_>>();
providers.sort_by(|left, right| left.0.id.cmp(&right.0.id));
if providers.is_empty() {
return Ok(PoolScoreRebuildRunSummary::empty());
}
let provider_ids = providers
.iter()
.map(|(provider, _)| provider.id.clone())
.collect::<Vec<_>>();
let mut keys_by_provider = BTreeMap::new();
for key in state
.list_provider_catalog_keys_by_provider_ids(&provider_ids)
.await?
{
keys_by_provider
.entry(key.provider_id.clone())
.or_insert_with(Vec::new)
.push(key);
}
for keys in keys_by_provider.values_mut() {
keys.sort_by(|left, right| left.id.cmp(&right.id));
}
let now = now_unix_secs();
let mut summary = PoolScoreRebuildRunSummary {
providers_checked: providers.len(),
..PoolScoreRebuildRunSummary::empty()
};
let start_provider_index =
load_runtime_usize(state, POOL_SCORE_REBUILD_PROVIDER_CURSOR_KEY).await % providers.len();
let mut last_provider_index = None;
for provider_index in
(0..providers.len()).map(|offset| (start_provider_index + offset) % providers.len())
{
if summary.scores_upserted >= config.max_upserts_per_tick {
break;
}
last_provider_index = Some(provider_index);
let (provider, pool_config) = providers[provider_index].clone();
let keys = keys_by_provider.remove(&provider.id).unwrap_or_default();
let keys = keys
.into_iter()
.filter(|key| key.is_active)
.collect::<Vec<_>>();
if keys.is_empty() {
continue;
}
let total_keys = keys.len();
let provider_cursor_key = provider_offset_cursor_key(&provider.id);
let provider_cursor = load_runtime_usize(state, &provider_cursor_key).await % total_keys;
let remaining_budget = config
.max_upserts_per_tick
.saturating_sub(summary.scores_upserted);
let provider_budget = remaining_budget.min(total_keys);
let mut build_items = Vec::with_capacity(provider_budget);
for offset in 0..provider_budget {
let key_index = (provider_cursor + offset) % total_keys;
let key = &keys[key_index];
let draft = build_provider_key_pool_score_upsert(
key,
provider.provider_type.as_str(),
None,
now,
pool_config.score_rules,
);
build_items.push((key_index, draft.id));
}
if build_items.is_empty() {
store_runtime_usize(
state,
&provider_cursor_key,
(provider_cursor + provider_budget) % total_keys,
)
.await;
continue;
}
let existing_scores = state
.data
.get_pool_member_scores_by_ids(&GetPoolMemberScoresByIdsQuery {
ids: build_items
.iter()
.map(|(_, score_id)| score_id.clone())
.collect(),
})
.await
.unwrap_or_else(|err| {
debug!(
provider_id = %provider.id,
error = ?err,
"gateway pool score rebuild: failed to read existing scores by id"
);
Vec::new()
})
.into_iter()
.map(|score| (score.id.clone(), score))
.collect::<BTreeMap<_, _>>();
let mut provider_upserts = 0usize;
summary.keys_seen = summary.keys_seen.saturating_add(keys.len());
for (key_index, score_id) in &build_items {
if summary.scores_upserted >= config.max_upserts_per_tick {
break;
}
let key = &keys[*key_index];
let existing = existing_scores.get(score_id);
let upsert = build_provider_key_pool_score_upsert(
key,
provider.provider_type.as_str(),
existing,
now,
pool_config.score_rules,
);
if state
.data
.upsert_pool_member_score(upsert)
.await
.map_err(|err| GatewayError::Internal(format!("{err:?}")))?
.is_some()
{
summary.scores_upserted = summary.scores_upserted.saturating_add(1);
provider_upserts = provider_upserts.saturating_add(1);
}
}
store_runtime_usize(
state,
&provider_cursor_key,
(provider_cursor + provider_budget) % total_keys,
)
.await;
if provider_upserts > 0 {
summary.providers_scored = summary.providers_scored.saturating_add(1);
}
}
if let Some(last_provider_index) = last_provider_index {
store_runtime_usize(
state,
POOL_SCORE_REBUILD_PROVIDER_CURSOR_KEY,
(last_provider_index + 1) % providers.len(),
)
.await;
}
Ok(summary)
}
pub(crate) async fn perform_pool_score_rebuild_once(
state: &AppState,
) -> Result<PoolScoreRebuildRunSummary, GatewayError> {
perform_pool_score_rebuild_once_with_config(state, PoolScoreRebuildWorkerConfig::from_env())
.await
}
pub(crate) fn spawn_pool_score_rebuild_worker(
state: AppState,
) -> Option<tokio::task::JoinHandle<()>> {
if !state.has_provider_catalog_data_reader()
|| !state.data.has_pool_score_reader()
|| !state.data.has_pool_score_writer()
{
return None;
}
let config = PoolScoreRebuildWorkerConfig::from_env();
Some(tokio::spawn(async move {
if let Err(err) = perform_pool_score_rebuild_once_with_config(&state, config).await {
warn!(
error = ?err,
"gateway pool score rebuild initial tick failed"
);
}
let mut interval = tokio::time::interval(config.interval);
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
interval.tick().await;
match perform_pool_score_rebuild_once_with_config(&state, config).await {
Ok(summary) if summary.scores_upserted > 0 => {
info!(
providers_checked = summary.providers_checked,
providers_scored = summary.providers_scored,
keys_seen = summary.keys_seen,
scores_upserted = summary.scores_upserted,
"gateway pool score rebuild completed"
);
}
Ok(_) => {}
Err(err) => {
warn!(
error = ?err,
"gateway pool score rebuild worker tick failed"
);
}
}
}
}))
}

View File

@@ -2,6 +2,9 @@ use std::collections::BTreeMap;
use aether_admin::provider::quota as admin_provider_quota_pure;
use aether_contracts::{ExecutionPlan, ExecutionTelemetry};
use aether_data_contracts::repository::pool_scores::{
PoolMemberHardState, PoolMemberIdentity, PoolMemberScheduleFeedback,
};
use aether_scheduler_core::{
build_scheduler_affinity_cache_key_for_api_key_id_with_client_session,
count_recent_rpm_requests_for_provider_key, ClientSessionAffinity, SchedulerAffinityTarget,
@@ -381,6 +384,19 @@ async fn record_sync_pool_success_effect(
resolve_ttfb_ms(payload.telemetry.as_ref()),
)
.await;
record_pool_score_schedule_feedback(
state,
context,
Some(true),
Some(PoolMemberHardState::Available),
Some(50),
serde_json::json!({
"last_request_feedback": {
"source": "sync_success"
}
}),
)
.await;
}
async fn record_adaptive_rate_limit_effect(
@@ -600,6 +616,19 @@ async fn record_stream_pool_success_effect(
resolve_ttfb_ms(payload.telemetry.as_ref()),
)
.await;
record_pool_score_schedule_feedback(
state,
context,
Some(true),
Some(PoolMemberHardState::Available),
Some(50),
serde_json::json!({
"last_request_feedback": {
"source": "stream_success"
}
}),
)
.await;
}
async fn record_pool_error_effect(
@@ -636,6 +665,21 @@ async fn record_pool_error_effect(
Some(effect.headers),
)
.await;
record_pool_score_schedule_feedback(
state,
context,
Some(false),
pool_score_hard_state_for_status(effect.status_code, effect.error_body),
Some(pool_score_delta_for_status(effect.status_code)),
serde_json::json!({
"last_request_feedback": {
"source": "pool_error",
"status_code": effect.status_code,
"classification": format!("{:?}", effect.classification)
}
}),
)
.await;
}
async fn open_pool_key_circuit_breaker(
@@ -730,6 +774,21 @@ async fn record_oauth_invalidation_effect(
plan.provider_id, plan.endpoint_id, plan.key_id, err
);
}
record_pool_score_schedule_feedback(
state,
context,
Some(false),
Some(PoolMemberHardState::AuthInvalid),
Some(-2_000),
serde_json::json!({
"last_request_feedback": {
"source": "oauth_invalidation",
"status_code": effect.status_code,
"reason": invalid_reason
}
}),
)
.await;
}
fn resolve_local_oauth_invalid_reason(
@@ -792,6 +851,92 @@ async fn record_pool_stream_timeout_effect(
&pool_context.pool_config,
)
.await;
record_pool_score_schedule_feedback(
state,
context,
Some(false),
Some(PoolMemberHardState::Cooldown),
Some(-250),
serde_json::json!({
"last_request_feedback": {
"source": "stream_timeout"
}
}),
)
.await;
}
async fn record_pool_score_schedule_feedback(
state: &AppState,
context: LocalExecutionEffectContext<'_>,
succeeded: Option<bool>,
hard_state: Option<PoolMemberHardState>,
score_delta: Option<i32>,
score_reason_patch: Value,
) {
if context.plan.provider_id.trim().is_empty() || context.plan.key_id.trim().is_empty() {
return;
}
let feedback = PoolMemberScheduleFeedback {
identity: PoolMemberIdentity::provider_api_key(
context.plan.provider_id.clone(),
context.plan.key_id.clone(),
),
scope: None,
scheduled_at: current_unix_secs(),
succeeded,
hard_state,
score_delta,
score_reason_patch: Some(score_reason_patch),
};
if let Err(err) = state
.data
.record_pool_member_schedule_feedback(feedback)
.await
{
warn!(
provider_id = %context.plan.provider_id,
key_id = %context.plan.key_id,
error = ?err,
"gateway orchestration effects: failed to record pool score schedule feedback"
);
}
}
fn pool_score_hard_state_for_status(
status_code: u16,
error_body: Option<&str>,
) -> Option<PoolMemberHardState> {
match status_code {
401 | 403 => Some(PoolMemberHardState::AuthInvalid),
402 => Some(PoolMemberHardState::QuotaExhausted),
429 | 500..=599 => Some(PoolMemberHardState::Cooldown),
_ => {
let body = error_body.unwrap_or_default().to_ascii_lowercase();
if body.contains("quota") && body.contains("exceed") {
Some(PoolMemberHardState::QuotaExhausted)
} else if body.contains("invalid") && body.contains("token") {
Some(PoolMemberHardState::AuthInvalid)
} else if body.contains("banned")
|| body.contains("suspended")
|| body.contains("blocked")
{
Some(PoolMemberHardState::Banned)
} else {
None
}
}
}
}
fn pool_score_delta_for_status(status_code: u16) -> i32 {
match status_code {
401 | 403 => -2_000,
402 => -1_000,
429 => -500,
500..=599 => -300,
_ => -100,
}
}
#[cfg(test)]
@@ -1188,8 +1333,8 @@ mod tests {
build_scheduler_affinity_cache_key_for_api_key_id("api-key-1", "openai:chat", "gpt-5")
.expect("scheduler affinity cache key should build");
state.scheduler_affinity_cache.insert(
cache_key.clone(),
state.remember_scheduler_affinity_target(
&cache_key,
SchedulerAffinityTarget {
provider_id: "prov-1".to_string(),
endpoint_id: "ep-1".to_string(),
@@ -1231,8 +1376,8 @@ mod tests {
.expect("legacy scheduler affinity cache key should build");
for cache_key in [&session_cache_key, &legacy_cache_key] {
state.scheduler_affinity_cache.insert(
cache_key.to_string(),
state.remember_scheduler_affinity_target(
cache_key.as_str(),
SchedulerAffinityTarget {
provider_id: "prov-1".to_string(),
endpoint_id: "ep-1".to_string(),
@@ -1277,8 +1422,8 @@ mod tests {
build_scheduler_affinity_cache_key_for_api_key_id("api-key-1", "openai:chat", "gpt-5")
.expect("scheduler affinity cache key should build");
state.scheduler_affinity_cache.insert(
cache_key.clone(),
state.remember_scheduler_affinity_target(
&cache_key,
SchedulerAffinityTarget {
provider_id: "prov-1".to_string(),
endpoint_id: "ep-1".to_string(),
@@ -1319,8 +1464,8 @@ mod tests {
build_scheduler_affinity_cache_key_for_api_key_id("api-key-1", "openai:chat", "gpt-5")
.expect("scheduler affinity cache key should build");
state.scheduler_affinity_cache.insert(
cache_key.clone(),
state.remember_scheduler_affinity_target(
&cache_key,
SchedulerAffinityTarget {
provider_id: "prov-1".to_string(),
endpoint_id: "ep-1".to_string(),
@@ -1467,8 +1612,8 @@ mod tests {
build_scheduler_affinity_cache_key_for_api_key_id("api-key-1", "openai:chat", "gpt-5")
.expect("scheduler affinity cache key should build");
state.scheduler_affinity_cache.insert(
cache_key.clone(),
state.remember_scheduler_affinity_target(
&cache_key,
SchedulerAffinityTarget {
provider_id: "prov-1".to_string(),
endpoint_id: "ep-1".to_string(),

View File

@@ -120,6 +120,10 @@ impl FrontdoorUserRpmLimiter {
self.resolve_system_default_limit(state).await
}
pub(crate) fn clear_system_default_cache(&self) {
self.system_default_cache.clear();
}
pub(crate) fn current_bucket(&self, now_ts: u64) -> u64 {
self.config.current_bucket(now_ts)
}

View File

@@ -88,6 +88,8 @@ fn frontend_path_bypasses_static(path: &str) -> bool {
|| path.starts_with("/upload/")
|| path.starts_with("/_gateway/")
|| path.starts_with("/.well-known/")
|| path.starts_with("/install/")
|| path.starts_with("/i/")
}
fn frontend_path_targets_static_asset(path: &str) -> bool {

View File

@@ -31,7 +31,9 @@ use regex::Regex;
use sha2::{Digest, Sha256};
use std::collections::BTreeMap;
pub(crate) use self::selection::SchedulerSkippedCandidate;
pub(crate) use self::selection::{
SchedulerSkippedCandidate, API_KEY_CONCURRENCY_LIMIT_SKIP_REASON,
};
use crate::data::auth::GatewayAuthApiKeySnapshot;
use crate::data::candidate_selection::{

View File

@@ -20,7 +20,7 @@ pub(crate) struct SchedulerSkippedCandidate {
pub(crate) skip_reason: &'static str,
}
pub(super) const API_KEY_CONCURRENCY_LIMIT_SKIP_REASON: &str = "api_key_concurrency_limit_reached";
pub(crate) const API_KEY_CONCURRENCY_LIMIT_SKIP_REASON: &str = "api_key_concurrency_limit_reached";
pub(super) fn is_exact_all_skipped_by_auth_limit(
selected: &[SchedulerMinimalCandidateSelectionCandidate],

View File

@@ -210,8 +210,8 @@ async fn reuses_cached_scheduler_affinity_candidate_before_sorted_fallback() {
let cache_key =
build_scheduler_affinity_cache_key(Some(&auth_snapshot), "openai:chat", "gpt-4.1", None)
.expect("cache key should build");
state.scheduler_affinity_cache.insert(
cache_key,
state.remember_scheduler_affinity_target(
&cache_key,
SchedulerAffinityTarget {
provider_id: "provider-b".to_string(),
endpoint_id: "endpoint-b".to_string(),
@@ -315,8 +315,8 @@ async fn cached_affinity_candidate_cannot_use_reserved_provider_key_rpm_capacity
let cache_key =
build_scheduler_affinity_cache_key(Some(&auth_snapshot), "openai:chat", "gpt-4.1", None)
.expect("cache key should build");
state.scheduler_affinity_cache.insert(
cache_key,
state.remember_scheduler_affinity_target(
&cache_key,
SchedulerAffinityTarget {
provider_id: "provider-a".to_string(),
endpoint_id: "endpoint-a".to_string(),

View File

@@ -176,8 +176,8 @@ async fn required_capability_without_model_uses_session_scoped_affinity() {
Some(&client_session_affinity),
)
.expect("session affinity cache key should build");
state.scheduler_affinity_cache.insert(
cache_key,
state.remember_scheduler_affinity_target(
&cache_key,
SchedulerAffinityTarget {
provider_id: "provider-b".to_string(),
endpoint_id: "endpoint-b".to_string(),

View File

@@ -459,8 +459,8 @@ async fn fixed_order_ignores_cached_scheduler_affinity_promotion() {
);
let auth_snapshot = sample_auth_snapshot("affinity-key-1");
state.scheduler_affinity_cache.insert(
"scheduler_affinity:affinity-key-1:openai:chat:gpt-4.1".to_string(),
state.remember_scheduler_affinity_target(
"scheduler_affinity:affinity-key-1:openai:chat:gpt-4.1",
SchedulerAffinityTarget {
provider_id: "provider-b".to_string(),
endpoint_id: "endpoint-b".to_string(),
@@ -578,8 +578,8 @@ async fn cache_affinity_promotes_cached_scheduler_affinity_candidate_when_enable
);
let auth_snapshot = sample_auth_snapshot("affinity-key-1");
state.scheduler_affinity_cache.insert(
"scheduler_affinity:affinity-key-1:openai:chat:gpt-4.1".to_string(),
state.remember_scheduler_affinity_target(
"scheduler_affinity:affinity-key-1:openai:chat:gpt-4.1",
SchedulerAffinityTarget {
provider_id: "provider-b".to_string(),
endpoint_id: "endpoint-b".to_string(),
@@ -683,8 +683,8 @@ async fn load_balance_ignores_provider_priority_and_cached_affinity() {
);
let auth_snapshot = sample_auth_snapshot("affinity-key-1");
state.scheduler_affinity_cache.insert(
"scheduler_affinity:affinity-key-1:openai:chat:gpt-4.1".to_string(),
state.remember_scheduler_affinity_target(
"scheduler_affinity:affinity-key-1:openai:chat:gpt-4.1",
SchedulerAffinityTarget {
provider_id: "provider-b".to_string(),
endpoint_id: "endpoint-b".to_string(),

View File

@@ -1,7 +1,8 @@
use super::{AppState, GatewayError, LocalMutationOutcome, LocalProviderDeleteTaskState};
use crate::handlers::shared::sync_provider_key_oauth_status_snapshot;
use aether_data_contracts::repository::{candidates, global_models, provider_catalog};
use aether_data_contracts::repository::{candidates, global_models, pool_scores, provider_catalog};
use std::time::{SystemTime, UNIX_EPOCH};
use tracing::warn;
impl AppState {
pub fn has_provider_catalog_data_reader(&self) -> bool {
@@ -468,7 +469,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if created.is_some() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(created)
}
@@ -484,7 +485,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if created.is_some() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(created)
}
@@ -499,7 +500,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated.is_some() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(updated)
}
@@ -514,7 +515,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if deleted {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(deleted)
}
@@ -529,8 +530,27 @@ impl AppState {
.cleanup_deleted_provider_catalog_refs(provider_id, endpoint_ids, key_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
for key_id in key_ids {
if let Err(err) = self
.data
.delete_pool_member_scores_for_member(
&pool_scores::PoolMemberIdentity::provider_api_key(
provider_id.to_string(),
key_id.to_string(),
),
)
.await
{
warn!(
provider_id,
key_id,
error = ?err,
"gateway provider catalog cleanup: failed to delete pool member scores"
);
}
}
if !endpoint_ids.is_empty() || !key_ids.is_empty() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(())
}
@@ -545,7 +565,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if created.is_some() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(created)
}
@@ -560,7 +580,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated.is_some() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(updated)
}
@@ -575,7 +595,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if deleted {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(deleted)
}
@@ -590,7 +610,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated.is_some() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(updated)
}
@@ -611,7 +631,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(updated)
}
@@ -620,13 +640,39 @@ impl AppState {
&self,
key_id: &str,
) -> Result<bool, GatewayError> {
let existing_key = self
.data
.list_provider_catalog_keys_by_ids(&[key_id.to_string()])
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
.into_iter()
.next();
let deleted = self
.data
.delete_provider_catalog_key(key_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if deleted {
self.clear_provider_transport_snapshot_cache();
if let Some(key) = existing_key {
if let Err(err) = self
.data
.delete_pool_member_scores_for_member(
&pool_scores::PoolMemberIdentity::provider_api_key(
key.provider_id.clone(),
key.id.clone(),
),
)
.await
{
warn!(
provider_id = %key.provider_id,
key_id = %key.id,
error = ?err,
"gateway provider catalog key delete: failed to delete pool member scores"
);
}
}
self.invalidate_provider_routing_caches();
}
Ok(deleted)
}
@@ -760,8 +806,128 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(updated)
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use crate::cache::SchedulerAffinityTarget;
use crate::data::GatewayDataState;
use crate::AppState;
fn sample_provider() -> StoredProviderCatalogProvider {
StoredProviderCatalogProvider::new(
"provider-1".to_string(),
"Provider 1".to_string(),
Some("https://example.com".to_string()),
"openai".to_string(),
)
.expect("provider should build")
}
fn sample_endpoint() -> StoredProviderCatalogEndpoint {
StoredProviderCatalogEndpoint::new(
"endpoint-1".to_string(),
"provider-1".to_string(),
"openai:chat".to_string(),
Some("openai".to_string()),
Some("chat".to_string()),
true,
)
.expect("endpoint should build")
.with_transport_fields(
"https://api.example.com/v1".to_string(),
None,
None,
None,
None,
None,
None,
None,
)
.expect("endpoint transport should build")
}
fn sample_key() -> StoredProviderCatalogKey {
StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"Key 1".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
}
#[tokio::test]
async fn provider_catalog_update_invalidates_scheduler_affinity_and_transport_snapshot_cache() {
let provider = sample_provider();
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider.clone()],
vec![sample_endpoint()],
vec![sample_key()],
));
let state = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(repository)
.with_encryption_key_for_tests("test-encryption-key"),
);
let snapshot = state
.read_provider_transport_snapshot("provider-1", "endpoint-1", "key-1")
.await
.expect("provider transport should read")
.expect("provider transport should exist");
assert!(!snapshot.provider.keep_priority_on_conversion);
let cache_key = "scheduler_affinity:api-key-1:openai:chat:gpt-5";
let ttl = Duration::from_secs(300);
state.remember_scheduler_affinity_target(
cache_key,
SchedulerAffinityTarget {
provider_id: "provider-1".to_string(),
endpoint_id: "endpoint-1".to_string(),
key_id: "key-1".to_string(),
},
ttl,
128,
);
assert!(state
.read_scheduler_affinity_target(cache_key, ttl)
.is_some());
let initial_epoch = state.scheduler_affinity_epoch();
let mut updated_provider = provider;
updated_provider.keep_priority_on_conversion = true;
updated_provider.provider_priority = -10;
state
.update_provider_catalog_provider(&updated_provider)
.await
.expect("provider update should succeed")
.expect("provider should update");
assert!(state.scheduler_affinity_epoch() > initial_epoch);
assert!(state
.read_scheduler_affinity_target(cache_key, ttl)
.is_none());
let snapshot = state
.read_provider_transport_snapshot("provider-1", "endpoint-1", "key-1")
.await
.expect("provider transport should read after update")
.expect("provider transport should exist after update");
assert!(snapshot.provider.keep_priority_on_conversion);
}
}

View File

@@ -47,6 +47,7 @@ use crate::maintenance::spawn_oauth_token_refresh_worker;
use crate::maintenance::spawn_pending_cleanup_worker;
use crate::maintenance::spawn_pool_monitor_worker;
use crate::maintenance::spawn_pool_quota_probe_worker;
use crate::maintenance::spawn_pool_score_rebuild_worker;
use crate::maintenance::spawn_provider_checkin_worker;
use crate::maintenance::spawn_proxy_node_metrics_cleanup_worker;
use crate::maintenance::spawn_proxy_node_stale_cleanup_worker;
@@ -58,6 +59,30 @@ use crate::maintenance::spawn_usage_cleanup_worker;
use crate::maintenance::spawn_wallet_daily_usage_aggregation_worker;
const SYSTEM_CONFIG_CACHE_TTL: Duration = Duration::from_secs(3);
const SCHEDULER_AFFECTING_SYSTEM_CONFIG_KEYS: &[&str] = &[
"enable_format_conversion",
"keep_priority_on_conversion",
"provider_priority_mode",
"scheduling_mode",
];
const AUTH_AFFECTING_SYSTEM_CONFIG_KEYS: &[&str] =
&[crate::constants::DEFAULT_USER_GROUP_CONFIG_KEY];
const FRONTDOOR_RPM_AFFECTING_SYSTEM_CONFIG_KEYS: &[&str] = &["rate_limit_per_minute"];
fn system_config_key_affects_scheduler(key: &str) -> bool {
let key = key.trim();
SCHEDULER_AFFECTING_SYSTEM_CONFIG_KEYS.contains(&key)
}
fn system_config_key_affects_auth(key: &str) -> bool {
let key = key.trim();
AUTH_AFFECTING_SYSTEM_CONFIG_KEYS.contains(&key)
}
fn system_config_key_affects_frontdoor_rpm(key: &str) -> bool {
let key = key.trim();
FRONTDOOR_RPM_AFFECTING_SYSTEM_CONFIG_KEYS.contains(&key)
}
impl AppState {
fn usage_worker_queue_for(
@@ -142,7 +167,10 @@ impl AppState {
pub(crate) fn replace_data_state(&mut self, data: Arc<GatewayDataState>) {
self.clear_provider_transport_snapshot_cache();
self.invalidate_scheduler_affinity_cache();
self.invalidate_auth_context_cache();
self.system_config_cache.clear();
self.frontdoor_user_rpm.clear_system_default_cache();
let data = Arc::new(
(*data)
.clone()
@@ -491,11 +519,7 @@ impl AppState {
.upsert_system_config_value(key, value, description)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.system_config_cache.insert(
key.to_string(),
Some(value.clone()),
SYSTEM_CONFIG_CACHE_TTL,
);
self.remember_system_config_write(key, Some(value.clone()));
Ok(value)
}
@@ -514,10 +538,13 @@ impl AppState {
value: &serde_json::Value,
description: Option<&str>,
) -> Result<crate::data::state::StoredSystemConfigEntry, GatewayError> {
self.data
let entry = self
.data
.upsert_system_config_entry(key, value, description)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.remember_system_config_write(entry.key.as_str(), Some(entry.value.clone()));
Ok(entry)
}
pub(crate) async fn delete_system_config_value(&self, key: &str) -> Result<bool, GatewayError> {
@@ -528,9 +555,41 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.system_config_cache
.insert(key.to_string(), None, SYSTEM_CONFIG_CACHE_TTL);
if deleted && system_config_key_affects_scheduler(key) {
self.invalidate_scheduler_affinity_cache();
}
if deleted && system_config_key_affects_auth(key) {
self.invalidate_auth_context_cache();
}
if deleted && system_config_key_affects_frontdoor_rpm(key) {
self.frontdoor_user_rpm.clear_system_default_cache();
}
Ok(deleted)
}
pub(crate) fn invalidate_provider_routing_caches(&self) {
self.clear_provider_transport_snapshot_cache();
self.invalidate_scheduler_affinity_cache();
}
pub(crate) fn invalidate_auth_context_cache(&self) {
self.auth_context_cache.clear();
}
fn remember_system_config_write(&self, key: &str, value: Option<serde_json::Value>) {
self.system_config_cache
.insert(key.to_string(), value, SYSTEM_CONFIG_CACHE_TTL);
if system_config_key_affects_scheduler(key) {
self.invalidate_scheduler_affinity_cache();
}
if system_config_key_affects_auth(key) {
self.invalidate_auth_context_cache();
}
if system_config_key_affects_frontdoor_rpm(key) {
self.frontdoor_user_rpm.clear_system_default_cache();
}
}
pub(crate) async fn read_admin_system_stats(
&self,
) -> Result<aether_data::repository::system::AdminSystemStats, GatewayError> {
@@ -557,7 +616,7 @@ impl AppState {
| aether_data::repository::system::AdminSystemPurgeTarget::Stats
) {
self.system_config_cache.clear();
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(summary)
}
@@ -1113,6 +1172,10 @@ impl AppState {
crate::task_runtime::TASK_KEY_POOL_QUOTA_PROBE,
spawn_pool_quota_probe_worker(self.clone()),
);
supervise_worker(
crate::task_runtime::TASK_KEY_POOL_SCORE_REBUILD,
spawn_pool_score_rebuild_worker(self.clone()),
);
supervise_worker(
crate::task_runtime::TASK_KEY_STATS_HOURLY_AGG,
spawn_stats_hourly_aggregation_worker(self.data.clone()),
@@ -1185,6 +1248,7 @@ mod tests {
use serde_json::json;
use super::AppState;
use crate::cache::SchedulerAffinityTarget;
use crate::data::GatewayDataState;
#[tokio::test]
@@ -1232,6 +1296,91 @@ mod tests {
);
}
#[tokio::test]
async fn system_config_entry_write_refreshes_cache_and_scheduler_affinity_for_routing_keys() {
let state = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(
GatewayDataState::disabled().with_system_config_values_for_tests([(
"keep_priority_on_conversion".to_string(),
json!(false),
)]),
);
let cache_key = "scheduler_affinity:api-key-1:openai:chat:gpt-5";
let ttl = std::time::Duration::from_secs(300);
assert_eq!(
state
.read_system_config_json_value("keep_priority_on_conversion")
.await
.expect("system config read should succeed"),
Some(json!(false))
);
state.remember_scheduler_affinity_target(
cache_key,
SchedulerAffinityTarget {
provider_id: "provider-old".to_string(),
endpoint_id: "endpoint-old".to_string(),
key_id: "key-old".to_string(),
},
ttl,
128,
);
assert!(state
.read_scheduler_affinity_target(cache_key, ttl)
.is_some());
let initial_epoch = state.scheduler_affinity_epoch();
state
.upsert_system_config_entry("keep_priority_on_conversion", &json!(true), None)
.await
.expect("admin config write should succeed");
assert_eq!(
state
.read_system_config_json_value("keep_priority_on_conversion")
.await
.expect("system config read should use refreshed cache"),
Some(json!(true))
);
assert!(state.scheduler_affinity_epoch() > initial_epoch);
assert_eq!(state.read_scheduler_affinity_target(cache_key, ttl), None);
}
#[tokio::test]
async fn system_config_write_refreshes_frontdoor_rpm_default_cache() {
let state = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(
GatewayDataState::disabled().with_system_config_values_for_tests([(
"rate_limit_per_minute".to_string(),
json!(1),
)]),
);
assert_eq!(
state
.frontdoor_user_rpm()
.current_system_default_limit(&state)
.await
.expect("default rpm limit should read"),
1
);
state
.upsert_system_config_entry("rate_limit_per_minute", &json!(0), None)
.await
.expect("rpm system config should update");
assert_eq!(
state
.frontdoor_user_rpm()
.current_system_default_limit(&state)
.await
.expect("default rpm limit should use refreshed value"),
0
);
}
#[tokio::test]
async fn replacing_data_state_clears_system_config_cache() {
let mut state = AppState::new()

View File

@@ -182,10 +182,15 @@ impl AppState {
record: aether_data::repository::auth::CreateUserApiKeyRecord,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.create_user_api_key(record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn create_standalone_api_key(
@@ -193,10 +198,15 @@ impl AppState {
record: aether_data::repository::auth::CreateStandaloneApiKeyRecord,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.create_standalone_api_key(record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn update_user_api_key_basic(
@@ -204,10 +214,15 @@ impl AppState {
record: aether_data::repository::auth::UpdateUserApiKeyBasicRecord,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.update_user_api_key_basic(record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn update_standalone_api_key_basic(
@@ -215,10 +230,15 @@ impl AppState {
record: aether_data::repository::auth::UpdateStandaloneApiKeyBasicRecord,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.update_standalone_api_key_basic(record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn set_user_api_key_active(
@@ -228,10 +248,15 @@ impl AppState {
is_active: bool,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.set_user_api_key_active(user_id, api_key_id, is_active)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn set_standalone_api_key_active(
@@ -240,10 +265,15 @@ impl AppState {
is_active: bool,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.set_standalone_api_key_active(api_key_id, is_active)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn set_user_api_key_locked(
@@ -252,10 +282,15 @@ impl AppState {
api_key_id: &str,
is_locked: bool,
) -> Result<bool, GatewayError> {
self.data
let updated = self
.data
.set_user_api_key_locked(user_id, api_key_id, is_locked)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated {
self.invalidate_auth_context_cache();
}
Ok(updated)
}
pub(crate) async fn set_user_api_key_allowed_providers(
@@ -265,10 +300,15 @@ impl AppState {
allowed_providers: Option<Vec<String>>,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.set_user_api_key_allowed_providers(user_id, api_key_id, allowed_providers)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn set_user_api_key_force_capabilities(
@@ -278,10 +318,15 @@ impl AppState {
force_capabilities: Option<serde_json::Value>,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.set_user_api_key_force_capabilities(user_id, api_key_id, force_capabilities)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn delete_user_api_key(
@@ -289,19 +334,29 @@ impl AppState {
user_id: &str,
api_key_id: &str,
) -> Result<bool, GatewayError> {
self.data
let deleted = self
.data
.delete_user_api_key(user_id, api_key_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if deleted {
self.invalidate_auth_context_cache();
}
Ok(deleted)
}
pub(crate) async fn delete_standalone_api_key(
&self,
api_key_id: &str,
) -> Result<bool, GatewayError> {
self.data
let deleted = self
.data
.delete_standalone_api_key(api_key_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if deleted {
self.invalidate_auth_context_cache();
}
Ok(deleted)
}
}

View File

@@ -248,10 +248,15 @@ impl AppState {
&self,
record: aether_data::repository::users::UpsertUserGroupRecord,
) -> Result<Option<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.data
let group = self
.data
.create_user_group(record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if group.is_some() {
self.invalidate_auth_context_cache();
}
Ok(group)
}
pub(crate) async fn update_user_group(
@@ -259,17 +264,27 @@ impl AppState {
group_id: &str,
record: aether_data::repository::users::UpsertUserGroupRecord,
) -> Result<Option<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.data
let group = self
.data
.update_user_group(group_id, record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if group.is_some() {
self.invalidate_auth_context_cache();
}
Ok(group)
}
pub(crate) async fn delete_user_group(&self, group_id: &str) -> Result<bool, GatewayError> {
self.data
let deleted = self
.data
.delete_user_group(group_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if deleted {
self.invalidate_auth_context_cache();
}
Ok(deleted)
}
pub(crate) async fn list_user_group_members(
@@ -287,10 +302,13 @@ impl AppState {
group_id: &str,
user_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroupMember>, GatewayError> {
self.data
let members = self
.data
.replace_user_group_members(group_id, user_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.invalidate_auth_context_cache();
Ok(members)
}
pub(crate) async fn list_user_groups_for_user(
@@ -318,10 +336,13 @@ impl AppState {
user_id: &str,
group_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.data
let groups = self
.data
.replace_user_groups_for_user(user_id, group_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.invalidate_auth_context_cache();
Ok(groups)
}
pub(crate) async fn add_user_to_group(
@@ -329,10 +350,15 @@ impl AppState {
group_id: &str,
user_id: &str,
) -> Result<bool, GatewayError> {
self.data
let added = self
.data
.add_user_to_group(group_id, user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if added {
self.invalidate_auth_context_cache();
}
Ok(added)
}
pub(crate) async fn is_other_user_auth_email_taken(
@@ -417,13 +443,19 @@ impl AppState {
.lock()
.expect("auth user store should lock")
.insert(user.id.clone(), user.clone());
self.invalidate_auth_context_cache();
return Ok(Some(user));
}
self.data
let user = self
.data
.update_local_auth_user_profile(user_id, email, username)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if user.is_some() {
self.invalidate_auth_context_cache();
}
Ok(user)
}
pub(crate) async fn update_local_auth_user_password_hash(
@@ -606,10 +638,14 @@ impl AppState {
user.is_active = is_active;
}
let _ = (rate_limit_present, rate_limit);
return Ok(Some(user.clone()));
let user = user.clone();
drop(guard);
self.invalidate_auth_context_cache();
return Ok(Some(user));
}
self.data
let user = self
.data
.update_local_auth_user_admin_fields(
user_id,
role,
@@ -624,7 +660,11 @@ impl AppState {
is_active,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if user.is_some() {
self.invalidate_auth_context_cache();
}
Ok(user)
}
pub(crate) async fn update_local_auth_user_policy_modes(
@@ -651,10 +691,14 @@ impl AppState {
user.allowed_models_mode = mode;
}
let _ = rate_limit_mode;
return Ok(Some(user.clone()));
let user = user.clone();
drop(guard);
self.invalidate_auth_context_cache();
return Ok(Some(user));
}
self.data
let user = self
.data
.update_local_auth_user_policy_modes(
user_id,
allowed_providers_mode,
@@ -663,7 +707,11 @@ impl AppState {
rate_limit_mode,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if user.is_some() {
self.invalidate_auth_context_cache();
}
Ok(user)
}
pub(crate) async fn touch_auth_user_last_login(
@@ -821,3 +869,90 @@ fn normalized_user_group_ids(group_ids: &[String]) -> BTreeSet<String> {
.map(ToOwned::to_owned)
.collect()
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use aether_data::repository::users::{InMemoryUserReadRepository, UpsertUserGroupRecord};
use crate::control::GatewayControlAuthContext;
use crate::data::GatewayDataState;
use crate::AppState;
fn user_group_record(
allowed_models: Option<Vec<&str>>,
allowed_models_mode: &str,
) -> UpsertUserGroupRecord {
UpsertUserGroupRecord {
name: "Team".to_string(),
description: None,
priority: 0,
allowed_providers: None,
allowed_providers_mode: "unrestricted".to_string(),
allowed_api_formats: None,
allowed_api_formats_mode: "unrestricted".to_string(),
allowed_models: allowed_models.map(|values| {
values
.into_iter()
.map(ToOwned::to_owned)
.collect::<Vec<_>>()
}),
allowed_models_mode: allowed_models_mode.to_string(),
rate_limit: None,
rate_limit_mode: "inherit".to_string(),
}
}
fn cached_auth_context() -> GatewayControlAuthContext {
GatewayControlAuthContext {
user_id: "user-1".to_string(),
api_key_id: "key-1".to_string(),
username: Some("alice".to_string()),
api_key_name: Some("default".to_string()),
balance_remaining: None,
access_allowed: true,
user_rate_limit: None,
api_key_rate_limit: None,
api_key_is_standalone: false,
admin_bypass_limits: false,
local_rejection: None,
allowed_models: Some(vec!["gpt-4.1".to_string()]),
}
}
#[tokio::test]
async fn updating_user_group_invalidates_cached_auth_context() {
let repository = Arc::new(InMemoryUserReadRepository::default());
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(GatewayDataState::with_user_reader_for_tests(repository));
let group = state
.create_user_group(user_group_record(Some(vec!["gpt-4.1"]), "specific"))
.await
.expect("group should create")
.expect("group should exist");
let cache_key = "auth-context-cache-key".to_string();
let ttl = Duration::from_secs(60);
state
.auth_context_cache
.insert(cache_key.clone(), cached_auth_context(), ttl, 10);
assert!(state
.auth_context_cache
.get_fresh(&cache_key, ttl)
.is_some());
state
.update_user_group(&group.id, user_group_record(None, "unrestricted"))
.await
.expect("group should update")
.expect("group should exist after update");
assert!(state
.auth_context_cache
.get_fresh(&cache_key, ttl)
.is_none());
}
}

View File

@@ -234,13 +234,19 @@ impl AppState {
.lock()
.expect("auth wallet store should lock")
.insert(wallet.id.clone(), wallet.clone());
self.invalidate_auth_context_cache();
return Ok(Some(wallet));
}
self.data
let wallet = self
.data
.initialize_auth_user_wallet(user_id, initial_gift_usd, unlimited)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if wallet.is_some() {
self.invalidate_auth_context_cache();
}
Ok(wallet)
}
pub(crate) async fn initialize_auth_api_key_wallet(
@@ -284,13 +290,19 @@ impl AppState {
.lock()
.expect("auth wallet store should lock")
.insert(wallet.id.clone(), wallet.clone());
self.invalidate_auth_context_cache();
return Ok(Some(wallet));
}
self.data
let wallet = self
.data
.initialize_auth_api_key_wallet(api_key_id, initial_gift_usd, unlimited)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if wallet.is_some() {
self.invalidate_auth_context_cache();
}
Ok(wallet)
}
pub(crate) async fn update_auth_user_wallet_limit_mode(
@@ -310,13 +322,21 @@ impl AppState {
let _ = wallet_id;
wallet.limit_mode = limit_mode.to_string();
wallet.updated_at_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
return Ok(Some(wallet.clone()));
let wallet = wallet.clone();
drop(guard);
self.invalidate_auth_context_cache();
return Ok(Some(wallet));
}
self.data
let wallet = self
.data
.update_auth_user_wallet_limit_mode(user_id, limit_mode)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if wallet.is_some() {
self.invalidate_auth_context_cache();
}
Ok(wallet)
}
pub(crate) async fn update_auth_api_key_wallet_limit_mode(
@@ -336,13 +356,21 @@ impl AppState {
let _ = wallet_id;
wallet.limit_mode = limit_mode.to_string();
wallet.updated_at_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
return Ok(Some(wallet.clone()));
let wallet = wallet.clone();
drop(guard);
self.invalidate_auth_context_cache();
return Ok(Some(wallet));
}
self.data
let wallet = self
.data
.update_auth_api_key_wallet_limit_mode(api_key_id, limit_mode)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if wallet.is_some() {
self.invalidate_auth_context_cache();
}
Ok(wallet)
}
#[allow(clippy::too_many_arguments)]
@@ -381,10 +409,14 @@ impl AppState {
if let Some(updated_at_unix_secs) = updated_at_unix_secs {
wallet.updated_at_unix_secs = updated_at_unix_secs;
}
return Ok(Some(wallet.clone()));
let wallet = wallet.clone();
drop(guard);
self.invalidate_auth_context_cache();
return Ok(Some(wallet));
}
self.data
let wallet = self
.data
.update_auth_user_wallet_snapshot(
user_id,
balance,
@@ -399,7 +431,11 @@ impl AppState {
updated_at_unix_secs,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if wallet.is_some() {
self.invalidate_auth_context_cache();
}
Ok(wallet)
}
#[allow(clippy::too_many_arguments)]
@@ -438,10 +474,14 @@ impl AppState {
if let Some(updated_at_unix_secs) = updated_at_unix_secs {
wallet.updated_at_unix_secs = updated_at_unix_secs;
}
return Ok(Some(wallet.clone()));
let wallet = wallet.clone();
drop(guard);
self.invalidate_auth_context_cache();
return Ok(Some(wallet));
}
self.data
let wallet = self
.data
.update_auth_api_key_wallet_snapshot(
api_key_id,
balance,
@@ -456,6 +496,10 @@ impl AppState {
updated_at_unix_secs,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if wallet.is_some() {
self.invalidate_auth_context_cache();
}
Ok(wallet)
}
}

View File

@@ -57,6 +57,16 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn list_pool_key_candidate_rows_for_group_key_ids(
&self,
query: &candidate_selection::StoredPoolKeyCandidateRowsByKeyIdsQuery,
) -> Result<Vec<candidate_selection::StoredMinimalCandidateSelectionRow>, GatewayError> {
self.data
.list_pool_key_candidate_rows_for_group_key_ids(query)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn read_provider_quota_snapshot(
&self,
provider_id: &str,

Some files were not shown because too many files have changed in this diff Show More