mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
@@ -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,
|
||||
|
||||
@@ -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::{
|
||||
|
||||
@@ -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(),
|
||||
),
|
||||
}]),
|
||||
};
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
|
||||
|
||||
@@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
161
apps/aether-gateway/src/ai_serving/planner/pool_scores.rs
Normal file
161
apps/aether-gateway/src/ai_serving/planner/pool_scores.rs
Normal 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
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
{
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -28,4 +28,8 @@ impl AuthContextCache {
|
||||
self.entries
|
||||
.insert(cache_key, auth_context, ttl, max_entries);
|
||||
}
|
||||
|
||||
pub(crate) fn clear(&self) {
|
||||
self.entries.clear();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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(&[]);
|
||||
|
||||
@@ -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");
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
|
||||
123
apps/aether-gateway/src/data/state/pool_scores.rs
Normal file
123
apps/aether-gateway/src/data/state/pool_scores.rs
Normal 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),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
3
apps/aether-gateway/src/dispatch/mod.rs
Normal file
3
apps/aether-gateway/src/dispatch/mod.rs
Normal file
@@ -0,0 +1,3 @@
|
||||
pub(crate) mod pool;
|
||||
pub(crate) mod pool_scheduler;
|
||||
pub(crate) mod refs;
|
||||
5
apps/aether-gateway/src/dispatch/pool.rs
Normal file
5
apps/aether-gateway/src/dispatch/pool.rs
Normal file
@@ -0,0 +1,5 @@
|
||||
use aether_dispatch_core::PoolWindowConfig;
|
||||
|
||||
pub(crate) fn default_pool_window_config() -> PoolWindowConfig {
|
||||
PoolWindowConfig::default()
|
||||
}
|
||||
2630
apps/aether-gateway/src/dispatch/pool_scheduler.rs
Normal file
2630
apps/aether-gateway/src/dispatch/pool_scheduler.rs
Normal file
File diff suppressed because it is too large
Load Diff
213
apps/aether-gateway/src/dispatch/refs.rs
Normal file
213
apps/aether-gateway/src/dispatch/refs.rs
Normal 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,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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/") =>
|
||||
|
||||
@@ -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> {
|
||||
|
||||
@@ -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?
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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!({
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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::{
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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(());
|
||||
};
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -36,6 +36,7 @@ mod clock;
|
||||
mod constants;
|
||||
mod control;
|
||||
mod data;
|
||||
mod dispatch;
|
||||
mod error;
|
||||
mod execution_runtime;
|
||||
mod executor;
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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::*;
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}))
|
||||
}
|
||||
@@ -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(),
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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::{
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user