use std::collections::{BTreeMap, BTreeSet}; use std::future::Future; use aether_admin::provider::{ pool as admin_provider_pool_pure, status as admin_provider_status_pure, }; use aether_data_contracts::repository::candidates::StoredRequestCandidate; use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey; use aether_scheduler_core::{ auth_api_key_concurrency_limit_reached, build_provider_concurrent_limit_map, candidate_is_selectable_with_runtime_state, candidate_runtime_skip_reason_with_state, effective_provider_key_rpm_limit, CandidateRuntimeSelectabilityInput, }; use std::time::{Duration, SystemTime, UNIX_EPOCH}; use tokio::time::Instant; use crate::data::auth::GatewayAuthApiKeySnapshot; use crate::GatewayError; use super::{SchedulerMinimalCandidateSelectionCandidate, SchedulerRuntimeState}; pub(super) use aether_scheduler_core::should_skip_provider_quota; pub(super) struct CandidateRuntimeSelectionSnapshot { pub(super) recent_candidates: Vec, pub(super) provider_concurrent_limits: BTreeMap, pub(super) provider_key_rpm_states: BTreeMap, pub(super) pool_provider_ids: BTreeSet, provider_quota_blocks_requests: BTreeMap, key_account_quota_exhausted: BTreeMap, key_oauth_invalid: BTreeMap, provider_key_rpm_reset_ats: BTreeMap>, } pub(super) async fn read_candidate_runtime_selection_snapshot( state: &(impl SchedulerRuntimeState + ?Sized), candidates: &[SchedulerMinimalCandidateSelectionCandidate], auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, now_unix_secs: u64, ) -> Result { let provider_concurrent_limits = read_provider_concurrent_limits(state, candidates).await?; let provider_pool_state = read_provider_pool_state_map(state, candidates).await?; let pool_provider_ids = provider_pool_state .iter() .filter_map(|(provider_id, state)| state.pool_enabled.then_some(provider_id.clone())) .collect::>(); let provider_key_rpm_states = read_provider_key_rpm_states(state, candidates).await?; let recent_candidates = if runtime_snapshot_requires_recent_candidates( auth_snapshot, &provider_concurrent_limits, &provider_key_rpm_states, now_unix_secs, ) { state.read_recent_request_candidates(128).await? } else { Vec::new() }; let key_account_quota_exhausted = read_key_account_quota_exhaustion_map( candidates, &provider_key_rpm_states, &provider_pool_state, ); let key_oauth_invalid = read_key_oauth_invalid_map(candidates, &provider_key_rpm_states, now_unix_secs); let provider_quota_blocks_requests = read_provider_quota_block_map(state, candidates, now_unix_secs).await?; let provider_key_rpm_reset_ats = read_provider_key_rpm_reset_at_map(state, candidates, now_unix_secs); Ok(CandidateRuntimeSelectionSnapshot { recent_candidates, provider_concurrent_limits, provider_key_rpm_states, pool_provider_ids, provider_quota_blocks_requests, key_account_quota_exhausted, key_oauth_invalid, provider_key_rpm_reset_ats, }) } fn runtime_snapshot_requires_recent_candidates( auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, provider_concurrent_limits: &BTreeMap, provider_key_rpm_states: &BTreeMap, now_unix_secs: u64, ) -> bool { if auth_snapshot .and_then(|snapshot| snapshot.api_key_concurrent_limit) .is_some_and(|limit| limit > 0) { return true; } if provider_concurrent_limits.values().any(|limit| *limit > 0) { return true; } provider_key_rpm_states.values().any(|key| { key.concurrent_limit.is_some_and(|limit| limit > 0) || effective_provider_key_rpm_limit(key, now_unix_secs).is_some() }) } pub(super) fn auth_snapshot_concurrency_limit_reached( auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, snapshot: &CandidateRuntimeSelectionSnapshot, now_unix_secs: u64, ) -> bool { auth_snapshot_concurrency_limit(auth_snapshot).is_some_and(|(api_key_id, limit)| { auth_api_key_concurrency_limit_reached( &snapshot.recent_candidates, now_unix_secs, api_key_id, limit, ) }) } fn auth_snapshot_concurrency_limit( auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, ) -> Option<(&str, usize)> { let snapshot = auth_snapshot?; let limit = usize::try_from(snapshot.api_key_concurrent_limit?).ok()?; (limit > 0).then_some((snapshot.api_key_id.as_str(), limit)) } async fn read_auth_api_key_concurrency_limit_reached( state: &(impl SchedulerRuntimeState + ?Sized), auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, ) -> Result { let Some((api_key_id, limit)) = auth_snapshot_concurrency_limit(auth_snapshot) else { return Ok(false); }; let recent_candidates = state.read_recent_request_candidates(128).await?; Ok(auth_api_key_concurrency_limit_reached( &recent_candidates, crate::clock::current_unix_secs(), api_key_id, limit, )) } /// A retry always rebuilds candidates, including at the deadline. Only the /// intervening polls omit catalog, quota and ranking work while auth is blocked. pub(crate) async fn wait_for_auth_api_key_concurrency_retry( state: &(impl SchedulerRuntimeState + ?Sized), auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, deadline: Instant, poll_interval: Duration, ) -> Result { if Instant::now() >= deadline { return Ok(false); } let poll_interval = poll_interval.max(Duration::from_millis(1)); loop { let remaining = deadline.saturating_duration_since(Instant::now()); tokio::time::sleep(poll_interval.min(remaining)).await; if Instant::now() >= deadline || !read_auth_api_key_concurrency_limit_reached(state, auth_snapshot).await? { return Ok(true); } } } pub(crate) async fn select_with_auth_concurrency_wait( state: &(impl SchedulerRuntimeState + ?Sized), auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, now_unix_secs: u64, wait_timeout: Duration, poll_interval: Duration, mut select: Select, ) -> Result where Select: FnMut(u64) -> Selection, Selection: Future>, { let deadline = Instant::now() + wait_timeout; let mut attempt_now_unix_secs = now_unix_secs; loop { let (result, auth_limit_blocked) = select(attempt_now_unix_secs).await?; if !auth_limit_blocked || !wait_for_auth_api_key_concurrency_retry( state, auth_snapshot, deadline, poll_interval, ) .await? { return Ok(result); } attempt_now_unix_secs = crate::clock::current_unix_secs(); } } pub(super) fn is_candidate_selectable( candidate: &SchedulerMinimalCandidateSelectionCandidate, snapshot: &CandidateRuntimeSelectionSnapshot, now_unix_secs: u64, ) -> bool { let pool_group = snapshot .pool_provider_ids .contains(candidate.provider_id.as_str()); candidate_is_selectable_with_runtime_state(CandidateRuntimeSelectabilityInput { candidate, recent_candidates: &snapshot.recent_candidates, provider_concurrent_limits: &snapshot.provider_concurrent_limits, provider_key_rpm_states: &snapshot.provider_key_rpm_states, now_unix_secs, provider_quota_blocks_requests: snapshot .provider_quota_blocks_requests .get(candidate.provider_id.as_str()) .copied() .unwrap_or(false), account_quota_exhausted: !pool_group && snapshot .key_account_quota_exhausted .get(candidate.key_id.as_str()) .copied() .unwrap_or(false), oauth_invalid: !pool_group && snapshot .key_oauth_invalid .get(candidate.key_id.as_str()) .copied() .unwrap_or(false), enforce_key_circuit_breaker: !pool_group, rpm_reset_at: (!pool_group) .then(|| { snapshot .provider_key_rpm_reset_ats .get(candidate.key_id.as_str()) .copied() .flatten() }) .flatten(), }) } pub(super) fn current_candidate_runtime_skip_reason( candidate: &SchedulerMinimalCandidateSelectionCandidate, snapshot: &CandidateRuntimeSelectionSnapshot, now_unix_secs: u64, ) -> Option<&'static str> { let pool_group = snapshot .pool_provider_ids .contains(candidate.provider_id.as_str()); let provider_quota_blocks_requests = snapshot .provider_quota_blocks_requests .get(candidate.provider_id.as_str()) .copied() .unwrap_or(false); let rpm_reset_at = (!pool_group) .then(|| { snapshot .provider_key_rpm_reset_ats .get(candidate.key_id.as_str()) .copied() .flatten() }) .flatten(); candidate_runtime_skip_reason_with_state(CandidateRuntimeSelectabilityInput { candidate, recent_candidates: &snapshot.recent_candidates, provider_concurrent_limits: &snapshot.provider_concurrent_limits, provider_key_rpm_states: &snapshot.provider_key_rpm_states, now_unix_secs, provider_quota_blocks_requests, account_quota_exhausted: !pool_group && snapshot .key_account_quota_exhausted .get(candidate.key_id.as_str()) .copied() .unwrap_or(false), oauth_invalid: !pool_group && snapshot .key_oauth_invalid .get(candidate.key_id.as_str()) .copied() .unwrap_or(false), enforce_key_circuit_breaker: !pool_group, rpm_reset_at, }) } pub(super) async fn read_provider_concurrent_limits( state: &(impl SchedulerRuntimeState + ?Sized), candidates: &[SchedulerMinimalCandidateSelectionCandidate], ) -> Result, GatewayError> { let provider_ids = candidates .iter() .map(|candidate| candidate.provider_id.clone()) .collect::>() .into_iter() .collect::>(); if provider_ids.is_empty() { return Ok(BTreeMap::new()); } let providers = state .read_provider_catalog_providers_by_ids(&provider_ids) .await?; Ok(build_provider_concurrent_limit_map(providers)) } pub(super) async fn read_provider_key_rpm_states( state: &(impl SchedulerRuntimeState + ?Sized), candidates: &[SchedulerMinimalCandidateSelectionCandidate], ) -> Result, GatewayError> { let key_ids = candidates .iter() .map(|candidate| candidate.key_id.clone()) .collect::>() .into_iter() .collect::>(); if key_ids.is_empty() { return Ok(BTreeMap::new()); } let keys = state.read_provider_catalog_keys_by_ids(&key_ids).await?; Ok(keys .into_iter() .map(|key| (key.id.clone(), key)) .collect::>()) } async fn read_provider_quota_block_map( state: &(impl SchedulerRuntimeState + ?Sized), candidates: &[SchedulerMinimalCandidateSelectionCandidate], now_unix_secs: u64, ) -> Result, GatewayError> { let provider_ids = candidates .iter() .map(|candidate| candidate.provider_id.clone()) .collect::>() .into_iter() .collect::>(); let mut quota_blocks = BTreeMap::new(); for provider_id in provider_ids { let blocks_requests = state .read_provider_quota_snapshot(&provider_id) .await? .as_ref() .is_some_and(|quota| should_skip_provider_quota(quota, now_unix_secs)); quota_blocks.insert(provider_id, blocks_requests); } Ok(quota_blocks) } #[derive(Debug, Clone, Copy, Default)] struct ProviderPoolState { pool_enabled: bool, skip_exhausted_accounts: bool, reserve_minimum_quota: bool, } async fn read_provider_pool_state_map( state: &(impl SchedulerRuntimeState + ?Sized), candidates: &[SchedulerMinimalCandidateSelectionCandidate], ) -> Result, GatewayError> { let provider_ids = candidates .iter() .map(|candidate| candidate.provider_id.clone()) .collect::>() .into_iter() .collect::>(); if provider_ids.is_empty() { return Ok(BTreeMap::new()); } let providers = state .read_provider_catalog_providers_by_ids(&provider_ids) .await?; Ok(providers .into_iter() .map(|provider| { let pool_advanced = provider .config .as_ref() .and_then(|value| value.get("pool_advanced")); let skip_exhausted_accounts = pool_advanced .and_then(serde_json::Value::as_object) .and_then(|value| value.get("skip_exhausted_accounts")) .and_then(serde_json::Value::as_bool) .unwrap_or(false); let reserve_minimum_quota = pool_advanced .and_then(serde_json::Value::as_object) .and_then(|value| value.get("reserve_minimum_quota")) .and_then(serde_json::Value::as_bool) .unwrap_or(false); ( provider.id, ProviderPoolState { pool_enabled: pool_advanced.is_some(), skip_exhausted_accounts, reserve_minimum_quota, }, ) }) .collect()) } fn read_key_account_quota_exhaustion_map( candidates: &[SchedulerMinimalCandidateSelectionCandidate], provider_key_rpm_states: &BTreeMap, provider_pool_state: &BTreeMap, ) -> BTreeMap { candidates .iter() .map(|candidate| { let exhausted = provider_key_rpm_states .get(candidate.key_id.as_str()) .is_some_and(|key| { let account_exhausted = admin_provider_pool_pure::admin_pool_key_model_quota_exhausted( key, candidate.provider_type.as_str(), candidate.selected_provider_model_name.as_str(), ) .unwrap_or_else(|| { admin_provider_pool_pure::admin_pool_key_account_quota_exhausted( key, candidate.provider_type.as_str(), ) }); let hard_blocked = admin_provider_pool_pure::admin_pool_key_model_quota_hard_blocked( key, candidate.provider_type.as_str(), candidate.selected_provider_model_name.as_str(), ); let pool_state = provider_pool_state .get(candidate.provider_id.as_str()) .copied() .unwrap_or_default(); let reserve_reached = pool_state.reserve_minimum_quota && admin_provider_pool_pure::admin_pool_key_minimum_quota_reached( key, candidate.provider_type.as_str(), Some(candidate.selected_provider_model_name.as_str()), ); hard_blocked || reserve_reached || (pool_state.skip_exhausted_accounts && account_exhausted) }); (candidate.key_id.clone(), exhausted) }) .collect() } fn read_key_oauth_invalid_map( candidates: &[SchedulerMinimalCandidateSelectionCandidate], provider_key_rpm_states: &BTreeMap, now_unix_secs: u64, ) -> BTreeMap { candidates .iter() .map(|candidate| { let oauth_invalid = provider_key_rpm_states .get(candidate.key_id.as_str()) .is_some_and(|key| { key_requires_oauth_reauth(key, candidate.provider_type.as_str(), now_unix_secs) }); (candidate.key_id.clone(), oauth_invalid) }) .collect() } fn key_requires_oauth_reauth( key: &StoredProviderCatalogKey, provider_type: &str, now_unix_secs: u64, ) -> bool { if !key.auth_type.trim().eq_ignore_ascii_case("oauth") { return false; } let invalid_reason = key .oauth_invalid_reason .as_deref() .map(str::trim) .unwrap_or_default(); if !invalid_reason.is_empty() { return oauth_invalid_reason_blocks_scheduling( key, provider_type, invalid_reason, now_unix_secs, ); } false } fn oauth_invalid_reason_blocks_scheduling( key: &StoredProviderCatalogKey, provider_type: &str, invalid_reason: &str, now_unix_secs: u64, ) -> bool { let trimmed_reason = invalid_reason.trim(); let account_state = admin_provider_status_pure::resolve_pool_account_state( Some(provider_type), key.upstream_metadata.as_ref(), Some(trimmed_reason), ); if account_state.blocked && !account_state.recoverable && account_state .code .as_deref() .is_some_and(oauth_account_state_code_is_hard_block) { return true; } if oauth_invalid_reason_has_tag(trimmed_reason, "[REFRESH_FAILED]") { return oauth_access_token_expired(key, now_unix_secs); } false } fn oauth_invalid_reason_has_tag(reason: &str, tag: &str) -> bool { reason .lines() .map(str::trim) .any(|line| line.starts_with(tag)) } fn oauth_access_token_expired(key: &StoredProviderCatalogKey, now_unix_secs: u64) -> bool { let now_unix_secs = if now_unix_secs == 0 { SystemTime::now() .duration_since(UNIX_EPOCH) .ok() .map(|duration| duration.as_secs()) .unwrap_or(0) } else { now_unix_secs }; key.expires_at_unix_secs .is_none_or(|expires_at| expires_at == 0 || expires_at <= now_unix_secs) } fn oauth_account_state_code_is_hard_block(code: &str) -> bool { matches!( code.trim().to_ascii_lowercase().as_str(), "account_banned" | "account_suspended" | "account_disabled" | "workspace_deactivated" | "account_forbidden" | "account_blocked" | "account_verification" | "oauth_token_invalid" ) } fn read_provider_key_rpm_reset_at_map( state: &(impl SchedulerRuntimeState + ?Sized), candidates: &[SchedulerMinimalCandidateSelectionCandidate], now_unix_secs: u64, ) -> BTreeMap> { candidates .iter() .map(|candidate| { ( candidate.key_id.clone(), state.provider_key_rpm_reset_at(candidate.key_id.as_str(), now_unix_secs), ) }) .collect::>() } #[cfg(test)] mod reserve_minimum_quota_tests { use super::*; use serde_json::json; #[test] fn reserve_minimum_quota_is_independent_of_skip_exhausted_accounts() { let candidate = SchedulerMinimalCandidateSelectionCandidate { provider_id: "provider-codex".to_string(), provider_name: "codex".to_string(), provider_type: "codex".to_string(), provider_priority: 0, endpoint_id: "endpoint-codex".to_string(), endpoint_api_format: "openai:responses".to_string(), key_id: "key-codex".to_string(), key_name: "codex".to_string(), key_auth_type: "oauth".to_string(), key_internal_priority: 0, key_global_priority_for_format: None, key_capabilities: None, model_id: "model-codex".to_string(), global_model_id: "global-model-codex".to_string(), global_model_name: "gpt-5".to_string(), selected_provider_model_name: "gpt-5".to_string(), supports_streaming: true, mapping_matched_model: None, }; let mut key = StoredProviderCatalogKey::new( candidate.key_id.clone(), candidate.provider_id.clone(), "codex".to_string(), "oauth".to_string(), None, true, ) .expect("key should build"); for reserve_enabled in [false, true] { for used_percent in [99.0, 98.0] { key.upstream_metadata = Some(json!({"codex": {"primary_used_percent": used_percent}})); let exhausted = read_key_account_quota_exhaustion_map( std::slice::from_ref(&candidate), &BTreeMap::from([(key.id.clone(), key.clone())]), &BTreeMap::from([( candidate.provider_id.clone(), ProviderPoolState { pool_enabled: true, reserve_minimum_quota: reserve_enabled, skip_exhausted_accounts: false, }, )]), ); assert_eq!( exhausted.get(&key.id), Some(&(reserve_enabled && used_percent >= 99.0)) ); } } } }