fix(gateway): add provider key concurrency and cache affinity modes

This commit is contained in:
ZheFox
2026-09-01 08:06:58 +08:00
parent 36daba7a34
commit 9631b229b3
26 changed files with 1176 additions and 154 deletions
+34 -5
View File
@@ -205,8 +205,8 @@ async fn available_balance_capacity_usd(
.as_ref()
.is_some_and(|wallet| wallet.limit_mode.eq_ignore_ascii_case("unlimited"));
Ok(match quota.as_ref() {
Some(quota) if !quota.allow_wallet_overage => Some(quota.remaining_usd.max(0.0)),
Some(_) if wallet_is_unlimited => None,
Some(quota) if !quota.allow_wallet_overage => Some(quota.remaining_usd.max(0.0)),
Some(quota) => Some(quota.remaining_usd.max(0.0) + wallet_available_usd.unwrap_or(0.0)),
None if wallet_is_unlimited => None,
None => wallet_available_usd,
@@ -832,9 +832,10 @@ mod tests {
use serde_json::json;
use super::{
execution_plan_balance_capacity_rejection, execution_plan_cost_upper_bound_cache_key,
max_output_tokens_from_request, openai_request_input_is_self_contained,
output_choice_count_upper_bound, request_model_local_rejection, GatewayLocalAuthRejection,
available_balance_capacity_usd, execution_plan_balance_capacity_rejection,
execution_plan_cost_upper_bound_cache_key, max_output_tokens_from_request,
openai_request_input_is_self_contained, output_choice_count_upper_bound,
request_model_local_rejection, GatewayLocalAuthRejection,
};
use crate::control::{GatewayControlAuthContext, GatewayControlDecision};
use crate::data::GatewayDataState;
@@ -939,6 +940,14 @@ mod tests {
fn state_with_quota_and_wallet(
quota: UserDailyQuotaAvailabilityRecord,
context: StoredBillingModelContext,
) -> AppState {
state_with_quota_context_and_wallet(quota, context, sample_wallet("user-1", 30.0))
}
fn state_with_quota_context_and_wallet(
quota: UserDailyQuotaAvailabilityRecord,
context: StoredBillingModelContext,
wallet: StoredWalletSnapshot,
) -> AppState {
let candidate_repository =
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
@@ -952,7 +961,7 @@ mod tests {
AppState::new()
.expect("state should build")
.with_data_state_for_tests(data)
.with_auth_wallets_for_tests(vec![sample_wallet("user-1", 30.0)])
.with_auth_wallets_for_tests(vec![wallet])
}
fn state_with_model_mapping() -> AppState {
@@ -1350,6 +1359,26 @@ mod tests {
}
}
#[tokio::test]
async fn unlimited_wallet_capacity_ignores_exhausted_non_overage_quota() {
let context = billing_context_with_pricing(None, None, None, None);
let mut wallet = sample_wallet("user-1", 0.0);
wallet.limit_mode = "unlimited".to_string();
let state =
state_with_quota_context_and_wallet(quota_availability(0.0, false), context, wallet);
let decision = decision_with_allowed_models(vec!["gpt-5".to_string()]);
let auth_context = decision
.auth_context
.as_ref()
.expect("decision should include auth context");
let capacity = available_balance_capacity_usd(&state, auth_context)
.await
.expect("capacity should resolve");
assert_eq!(capacity, None);
}
#[tokio::test]
async fn positive_balance_does_not_allow_historical_invalid_processing_pricing() {
let context = billing_context_with_pricing(
@@ -1826,17 +1826,17 @@ fn pool_key_candidate_order_for_group(
})
.collect::<Vec<_>>();
let active_presets = ProviderPoolService::with_builtin_adapters()
.normalize_scheduling_presets(group.transport.provider.provider_type.as_str(), &presets)
.into_iter()
.map(|preset| preset.preset)
.collect::<Vec<_>>();
.normalize_scheduling_presets(group.transport.provider.provider_type.as_str(), &presets);
if let Some(distribution_mode) = active_presets
.iter()
.find(|preset| pool_distribution_mode_preset(preset.as_str()))
.map(String::as_str)
.find(|preset| pool_distribution_mode_preset(preset.preset.as_str()))
{
return match distribution_mode {
"cache_affinity" => StoredPoolKeyCandidateOrder::CacheAffinity,
return match distribution_mode.preset.as_str() {
"cache_affinity" => match distribution_mode.mode.as_deref() {
Some("lru") => StoredPoolKeyCandidateOrder::Lru,
Some("single_account") => StoredPoolKeyCandidateOrder::SingleAccount,
_ => StoredPoolKeyCandidateOrder::CacheAffinity,
},
"load_balance" => StoredPoolKeyCandidateOrder::LoadBalance {
seed: pool_sort_seed(),
},
@@ -2151,7 +2151,7 @@ mod tests {
}
#[test]
fn pool_scheduler_promotes_sticky_hit_before_other_sorted_keys() {
fn pool_scheduler_promotes_sticky_hit_before_lru_secondary_order() {
let key_a = sample_eligible_candidate(
"provider-pool",
"endpoint-1",
@@ -2159,7 +2159,11 @@ mod tests {
10,
Some(json!({
"pool_advanced": {
"scheduling_presets": [{"preset": "cache_affinity", "enabled": true}]
"scheduling_presets": [{
"preset": "cache_affinity",
"enabled": true,
"mode": "lru"
}]
}
})),
);
@@ -2170,7 +2174,11 @@ mod tests {
10,
Some(json!({
"pool_advanced": {
"scheduling_presets": [{"preset": "cache_affinity", "enabled": true}]
"scheduling_presets": [{
"preset": "cache_affinity",
"enabled": true,
"mode": "lru"
}]
}
})),
);
@@ -2204,6 +2212,37 @@ mod tests {
);
}
#[test]
fn cache_affinity_secondary_modes_select_distinct_candidate_orders() {
for (mode, expected) in [
("single_account", StoredPoolKeyCandidateOrder::SingleAccount),
("lru", StoredPoolKeyCandidateOrder::Lru),
] {
let group = sample_eligible_candidate(
"provider-pool",
"endpoint-1",
"key-a",
10,
Some(json!({
"pool_advanced": {
"scheduling_presets": [{
"preset": "cache_affinity",
"enabled": true,
"mode": mode
}]
}
})),
);
let config = pool_config_for_candidate(&group).expect("pool config should parse");
assert!(admin_provider_pool_cache_affinity_enabled(&config));
assert_eq!(
pool_key_candidate_order_for_group(&group, Some(&config)),
expected
);
}
}
#[test]
fn pool_scheduler_ignores_sticky_hit_without_cache_affinity() {
let key_a = sample_eligible_candidate(
@@ -126,7 +126,7 @@ use crate::orchestration::{
LocalOAuthSuccessEffect, LocalPoolErrorEffect,
};
use crate::provider_pool_demand::{
acquire_provider_pool_in_flight_guard, ProviderPoolInFlightGuard,
acquire_provider_pool_execution_guard, ProviderPoolInFlightAdmission, ProviderPoolInFlightGuard,
};
use crate::request_candidate_runtime::{
ensure_execution_request_candidate_slot, persist_local_request_candidate_status_record,
@@ -3772,6 +3772,41 @@ async fn execute_execution_runtime_stream_inner(
plan_kind,
report_context.as_ref(),
);
let candidate_started_unix_secs = current_request_candidate_unix_ms();
let provider_in_flight_started_at = Instant::now();
let mut provider_pool_in_flight_guard =
match acquire_provider_pool_execution_guard(state, &plan).await? {
ProviderPoolInFlightAdmission::Acquired(guard) => guard,
ProviderPoolInFlightAdmission::Saturated { limit } => {
if let Some(retry_scope) = retry_scope_out.as_deref_mut() {
*retry_scope = AiAttemptRetryScope::Candidate;
}
if let Some(snapshot) = request_candidate_status_snapshot.as_ref() {
record_local_request_candidate_status_snapshot(
state,
snapshot,
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Skipped,
status_code: Some(http::StatusCode::TOO_MANY_REQUESTS.as_u16()),
error_type: Some("provider_key_concurrency_limit_reached".to_string()),
error_message: Some(format!(
"provider key concurrency limit reached: {limit}"
)),
latency_ms: Some(0),
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(candidate_started_unix_secs),
},
)
.await;
}
return Ok(None);
}
};
observe_gateway_stage_trace_ms(
&mut stage_trace,
"stream_provider_in_flight",
provider_in_flight_started_at.elapsed().as_millis() as u64,
);
// Inline passthrough records its lifecycle seed after upstream headers are
// available. Avoid constructing a throwaway seed on the common path.
let mut lifecycle_seed = (!defer_stream_pending_for_direct_inline)
@@ -3781,7 +3816,6 @@ async fn execute_execution_runtime_stream_inner(
record_stream_pending_lifecycle(state, seed, &mut stage_trace).await;
lifecycle_pending_recorded = true;
}
let candidate_started_unix_secs = current_request_candidate_unix_ms();
if let Some(snapshot) = request_candidate_status_snapshot.clone() {
record_local_request_candidate_status_snapshot(
state,
@@ -3810,20 +3844,6 @@ async fn execute_execution_runtime_stream_inner(
.and_then(|context| context.candidate_index)
.map(|value| value.to_string())
.unwrap_or_else(|| "-".to_string());
let provider_in_flight_started_at = Instant::now();
let mut provider_pool_in_flight_guard = acquire_provider_pool_in_flight_guard(
state.runtime_state.clone(),
&plan.provider_id,
plan.request_id.as_str(),
plan.candidate_id.as_deref(),
key_id.as_str(),
)
.await;
observe_gateway_stage_trace_ms(
&mut stage_trace,
"stream_provider_in_flight",
provider_in_flight_started_at.elapsed().as_millis() as u64,
);
match maybe_execute_grok_stream(&plan, report_context.as_ref()).await {
Ok(Some(grok_stream)) => {
return execute_stream_from_frame_stream_with_retry_scope(
@@ -78,7 +78,9 @@ use crate::orchestration::{
LocalExecutionEffectContext, LocalHealthFailureEffect, LocalHealthSuccessEffect,
LocalOAuthInvalidationEffect, LocalOAuthSuccessEffect, LocalPoolErrorEffect,
};
use crate::provider_pool_demand::acquire_provider_pool_in_flight_guard;
use crate::provider_pool_demand::{
acquire_provider_pool_execution_guard, ProviderPoolInFlightAdmission,
};
use crate::request_candidate_runtime::{
ensure_execution_request_candidate_slot, record_local_request_candidate_extra_data,
record_local_request_candidate_status, record_local_request_candidate_status_snapshot,
@@ -1995,6 +1997,32 @@ async fn execute_execution_runtime_sync_impl(
.unwrap_or_else(|| "-".to_string());
let candidate_started_at = Instant::now();
let candidate_started_unix_secs = current_request_candidate_unix_ms();
let _provider_pool_in_flight_guard = match acquire_provider_pool_execution_guard(state, &plan)
.await?
{
ProviderPoolInFlightAdmission::Acquired(guard) => guard,
ProviderPoolInFlightAdmission::Saturated { limit } => {
if let Some(retry_scope) = retry_scope_out.as_deref_mut() {
*retry_scope = AiAttemptRetryScope::Candidate;
}
record_local_request_candidate_status(
state,
&plan,
report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Skipped,
status_code: Some(StatusCode::TOO_MANY_REQUESTS.as_u16()),
error_type: Some("provider_key_concurrency_limit_reached".to_string()),
error_message: Some(format!("provider key concurrency limit reached: {limit}")),
latency_ms: Some(0),
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(candidate_started_unix_secs),
},
)
.await;
return Ok(None);
}
};
let lifecycle_seed = build_lifecycle_usage_seed(&plan, report_context.as_ref());
let usage_data = state.usage_lifecycle_data_state().as_ref().clone();
state
@@ -2024,14 +2052,6 @@ async fn execute_execution_runtime_sync_impl(
candidate_started_at,
);
let result = (async {
let _provider_pool_in_flight_guard = acquire_provider_pool_in_flight_guard(
state.runtime_state.clone(),
&plan.provider_id,
plan_request_id.as_str(),
plan_candidate_id.as_deref(),
key_id.as_str(),
)
.await;
record_sync_execution_active(
state,
&plan,
@@ -167,6 +167,14 @@ fn parse_pool_score_rules(pool_advanced: &Map<String, Value>) -> PoolMemberScore
fn normalize_pool_preset_mode(preset: &str, raw_mode: Option<&Value>) -> Option<String> {
match preset {
"cache_affinity" => Some(
raw_mode
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| matches!(*value, "single_account" | "lru"))
.unwrap_or("single_account")
.to_string(),
),
"free_team_first" | "free_first" | "team_first" | "plus_first" | "pro_first" => {
let default_mode = match preset {
"free_team_first" => "both",
@@ -1342,6 +1342,7 @@ pub(super) fn build_admin_pool_key_payload(
json!(key.internal_priority),
);
payload.insert("rpm_limit".to_string(), json!(key.rpm_limit));
payload.insert("concurrent_limit".to_string(), json!(key.concurrent_limit));
payload.insert(
"cache_ttl_minutes".to_string(),
json!(key.cache_ttl_minutes),
@@ -11,7 +11,7 @@ use aether_contracts::ExecutionPlan;
use crate::execution_runtime::acquire_upstream_execution_gate;
use crate::provider_pool_demand::{
acquire_provider_pool_in_flight_guard, ProviderPoolInFlightGuard,
acquire_provider_pool_execution_guard, ProviderPoolInFlightAdmission, ProviderPoolInFlightGuard,
};
use crate::upstream_admission::UpstreamTargetAdmissionPermit;
use crate::{AppState, GatewayError};
@@ -41,14 +41,17 @@ impl ResponsesWebSocketTurnAdmission {
return Err(error);
}
};
let provider_pool = acquire_provider_pool_in_flight_guard(
state.runtime_state.clone(),
&plan.provider_id,
&plan.request_id,
plan.candidate_id.as_deref(),
&plan.key_id,
)
.await;
let provider_pool = match acquire_provider_pool_execution_guard(state, plan).await? {
ProviderPoolInFlightAdmission::Acquired(guard) => guard,
ProviderPoolInFlightAdmission::Saturated { limit } => {
drop(upstream_target);
drop(upstream_execution);
return Err(GatewayError::Client {
status: http::StatusCode::TOO_MANY_REQUESTS,
message: format!("上游账号并发已达上限 ({limit})"),
});
}
};
Ok(Self {
upstream_execution,
@@ -256,12 +256,16 @@ fn auth_config_has_refresh_token(auth_config: Option<&str>) -> bool {
let Ok(value) = serde_json::from_str::<Value>(auth_config) else {
return false;
};
value
.as_object()
.and_then(|object| object.get("refresh_token"))
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty())
let Some(object) = value.as_object() else {
return false;
};
["refresh_token", "refreshToken"].iter().any(|field| {
object
.get(*field)
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty())
})
}
fn now_unix_secs() -> u64 {
@@ -273,7 +277,44 @@ fn now_unix_secs() -> u64 {
#[cfg(test)]
mod tests {
use super::agent_identity_needs_task_recovery;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use super::{
agent_identity_needs_task_recovery, auth_config_has_refresh_token, oauth_refresh_candidate,
};
#[test]
fn legacy_antigravity_refresh_token_is_refreshable() {
assert!(auth_config_has_refresh_token(Some(
r#"{"refreshToken":"legacy-refresh-token"}"#,
)));
}
#[test]
fn expiring_antigravity_oauth_key_is_refresh_candidate() {
let provider = StoredProviderCatalogProvider::new(
"provider-antigravity".to_string(),
"Antigravity".to_string(),
None,
"antigravity".to_string(),
)
.expect("provider should build");
let mut key = StoredProviderCatalogKey::new(
"key-antigravity".to_string(),
provider.id.clone(),
"Antigravity OAuth".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build");
key.encrypted_auth_config = Some("encrypted-auth-config".to_string());
key.expires_at_unix_secs = Some(120);
assert!(oauth_refresh_candidate(&provider, &key, 120));
}
#[test]
fn pending_agent_identity_without_task_is_recoverable() {
+14 -5
View File
@@ -90,11 +90,15 @@ pub(crate) fn provider_key_can_refresh_oauth(
) -> bool {
auth_semantics.can_refresh_oauth()
&& (provider_key_auth_config_is_agent_identity(provider_type, auth_config)
|| auth_config
.and_then(|config| config.get("refresh_token"))
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty()))
|| auth_config.is_some_and(|config| {
["refresh_token", "refreshToken"].iter().any(|field| {
config
.get(*field)
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty())
})
}))
}
pub(crate) fn provider_key_can_export_oauth(
@@ -425,6 +429,11 @@ mod tests {
"codex",
json!({ "refresh_token": "refresh-token" }).as_object()
));
assert!(provider_key_can_refresh_oauth(
provider_key_auth_semantics(&sample_key("oauth"), "antigravity"),
"antigravity",
json!({ "refreshToken": "legacy-refresh-token" }).as_object()
));
assert!(provider_key_can_refresh_oauth(
semantics,
"codex",
+176 -11
View File
@@ -4,14 +4,20 @@ use std::sync::{
};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_runtime_state::RuntimeState;
use aether_contracts::ExecutionPlan;
use aether_runtime_state::{
RuntimeSemaphoreConfig, RuntimeSemaphoreError, RuntimeSemaphorePermit, RuntimeState,
};
use dashmap::DashMap;
use serde::{Deserialize, Serialize};
use tokio::task::JoinHandle;
use tracing::debug;
use uuid::Uuid;
use crate::{AppState, GatewayError};
const PROVIDER_POOL_IN_FLIGHT_TOKENS_PREFIX: &str = "ap:provider_pool:in_flight";
const PROVIDER_KEY_CONCURRENCY_GATE: &str = "provider_key";
const PROVIDER_POOL_DEMAND_SNAPSHOT_PREFIX: &str = "ap:provider_pool:demand";
const PROVIDER_POOL_BURST_PENDING_PREFIX: &str = "ap:quota_probe:burst_pending";
const PROVIDER_POOL_IN_FLIGHT_TOKEN_TTL_MS: u64 = 120_000;
@@ -42,10 +48,17 @@ pub(crate) struct ProviderPoolDemandSnapshot {
pub(crate) struct ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind,
provider_key_permit: Option<RuntimeSemaphorePermit>,
released: bool,
}
pub(crate) enum ProviderPoolInFlightAdmission {
Acquired(Option<ProviderPoolInFlightGuard>),
Saturated { limit: usize },
}
enum ProviderPoolInFlightGuardKind {
Disabled,
Local {
provider_id: String,
counter: Arc<AtomicUsize>,
@@ -69,6 +82,9 @@ impl ProviderPoolInFlightGuard {
return;
}
match &mut self.kind {
ProviderPoolInFlightGuardKind::Disabled => {
self.released = true;
}
ProviderPoolInFlightGuardKind::Local {
provider_id,
counter,
@@ -101,6 +117,16 @@ impl ProviderPoolInFlightGuard {
}
}
}
if self.released {
if let Some(provider_key_permit) = self.provider_key_permit.take() {
if let Err(err) = provider_key_permit.release().await {
debug!(
error = ?err,
"gateway provider pool demand: failed to release provider key permit; scheduling drop fallback"
);
}
}
}
}
}
@@ -111,6 +137,7 @@ impl Drop for ProviderPoolInFlightGuard {
}
self.released = true;
match &mut self.kind {
ProviderPoolInFlightGuardKind::Disabled => {}
ProviderPoolInFlightGuardKind::Local {
provider_id,
counter,
@@ -290,32 +317,83 @@ pub(crate) async fn acquire_provider_pool_in_flight_guard(
candidate_id: Option<&str>,
key_id: &str,
) -> Option<ProviderPoolInFlightGuard> {
acquire_provider_pool_in_flight_guard_with_key_limit(
runtime,
provider_id,
request_id,
candidate_id,
key_id,
None,
)
.await
.ok()
.flatten()
}
pub(crate) async fn acquire_provider_pool_in_flight_guard_with_key_limit(
runtime: Arc<RuntimeState>,
provider_id: &str,
request_id: &str,
candidate_id: Option<&str>,
key_id: &str,
concurrent_limit: Option<usize>,
) -> Result<Option<ProviderPoolInFlightGuard>, RuntimeSemaphoreError> {
let provider_key_permit = match concurrent_limit.filter(|limit| *limit > 0) {
Some(limit) => Some(
runtime
.keyed_semaphore(
PROVIDER_KEY_CONCURRENCY_GATE,
key_id,
limit,
RuntimeSemaphoreConfig::default(),
)?
.try_acquire()
.await?,
),
None => None,
};
let provider_id = provider_id.trim();
if provider_id.is_empty() {
return None;
return Ok(
provider_key_permit.map(|provider_key_permit| ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Disabled,
provider_key_permit: Some(provider_key_permit),
released: false,
}),
);
}
match provider_pool_in_flight_mode() {
ProviderPoolInFlightMode::Off => return None,
ProviderPoolInFlightMode::Off => {
return Ok(
provider_key_permit.map(|provider_key_permit| ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Disabled,
provider_key_permit: Some(provider_key_permit),
released: false,
}),
);
}
ProviderPoolInFlightMode::Local => {
let counter = increment_local_provider_in_flight(provider_id);
return Some(ProviderPoolInFlightGuard {
return Ok(Some(ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Local {
provider_id: provider_id.to_string(),
counter,
},
provider_key_permit,
released: false,
});
}));
}
ProviderPoolInFlightMode::Runtime if runtime.is_memory() => {
let counter = increment_local_provider_in_flight(provider_id);
return Some(ProviderPoolInFlightGuard {
return Ok(Some(ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Local {
provider_id: provider_id.to_string(),
counter,
},
provider_key_permit,
released: false,
});
}));
}
ProviderPoolInFlightMode::Runtime => {}
}
@@ -335,7 +413,13 @@ pub(crate) async fn acquire_provider_pool_in_flight_guard(
error = ?err,
"gateway provider pool demand: failed to acquire in-flight token"
);
return None;
return Ok(
provider_key_permit.map(|provider_key_permit| ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Disabled,
provider_key_permit: Some(provider_key_permit),
released: false,
}),
);
}
Err(_) => {
debug!(
@@ -343,7 +427,13 @@ pub(crate) async fn acquire_provider_pool_in_flight_guard(
timeout_ms = provider_pool_in_flight_acquire_timeout().as_millis() as u64,
"gateway provider pool demand: skipped in-flight token after acquire timeout"
);
return None;
return Ok(
provider_key_permit.map(|provider_key_permit| ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Disabled,
provider_key_permit: Some(provider_key_permit),
released: false,
}),
);
}
}
@@ -355,7 +445,7 @@ pub(crate) async fn acquire_provider_pool_in_flight_guard(
stop_renewal.clone(),
);
Some(ProviderPoolInFlightGuard {
Ok(Some(ProviderPoolInFlightGuard {
kind: ProviderPoolInFlightGuardKind::Runtime {
runtime,
tokens_key,
@@ -363,8 +453,39 @@ pub(crate) async fn acquire_provider_pool_in_flight_guard(
stop_renewal,
renew_handle: Some(renew_handle),
},
provider_key_permit,
released: false,
})
}))
}
pub(crate) async fn acquire_provider_pool_execution_guard(
state: &AppState,
plan: &ExecutionPlan,
) -> Result<ProviderPoolInFlightAdmission, GatewayError> {
let concurrent_limit = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await?
.into_iter()
.find(|key| key.id == plan.key_id)
.and_then(|key| key.concurrent_limit)
.filter(|limit| *limit > 0)
.and_then(|limit| usize::try_from(limit).ok());
match acquire_provider_pool_in_flight_guard_with_key_limit(
state.runtime_state.clone(),
&plan.provider_id,
&plan.request_id,
plan.candidate_id.as_deref(),
&plan.key_id,
concurrent_limit,
)
.await
{
Ok(guard) => Ok(ProviderPoolInFlightAdmission::Acquired(guard)),
Err(RuntimeSemaphoreError::Saturated { limit, .. }) => {
Ok(ProviderPoolInFlightAdmission::Saturated { limit })
}
Err(error) => Err(GatewayError::Internal(error.to_string())),
}
}
pub(crate) async fn provider_pool_live_in_flight_count(
@@ -586,6 +707,50 @@ mod tests {
);
}
#[tokio::test]
async fn provider_key_limit_rejects_concurrent_guard_until_release() {
let runtime = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let first = acquire_provider_pool_in_flight_guard_with_key_limit(
runtime.clone(),
"provider-limit",
"request-1",
Some("candidate-1"),
"key-limit",
Some(1),
)
.await
.expect("first admission should resolve")
.expect("first guard should be acquired");
let second = acquire_provider_pool_in_flight_guard_with_key_limit(
runtime.clone(),
"provider-limit",
"request-2",
Some("candidate-2"),
"key-limit",
Some(1),
)
.await;
assert!(matches!(
second,
Err(RuntimeSemaphoreError::Saturated { limit: 1, .. })
));
first.release().await;
let replacement = acquire_provider_pool_in_flight_guard_with_key_limit(
runtime,
"provider-limit",
"request-3",
Some("candidate-3"),
"key-limit",
Some(1),
)
.await
.expect("replacement admission should resolve")
.expect("replacement guard should acquire after release");
drop(replacement);
}
#[tokio::test]
async fn demand_snapshot_uses_instant_in_flight_for_fast_rise_and_ema_for_fall() {
let runtime = RuntimeState::memory(MemoryRuntimeStateConfig::default());
+89 -4
View File
@@ -33,7 +33,9 @@ use std::sync::atomic::Ordering;
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_crypto::encrypt_python_fernet_plaintext;
use aether_crypto::{
decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY,
};
const LOCAL_OAUTH_HTTP_TIMEOUT_MS: u64 = 30_000;
const REMOTE_OAUTH_REFRESH_WAIT_TIMEOUT: Duration = Duration::from_secs(35);
@@ -3395,7 +3397,10 @@ mod tests {
use std::sync::Arc;
use std::time::{Duration, Instant};
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
use aether_crypto::{
decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext,
DEVELOPMENT_ENCRYPTION_KEY,
};
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyListQuery,
@@ -3478,9 +3483,17 @@ mod tests {
fn codex_oauth_state(
auth_config: &serde_json::Value,
access_token: &str,
) -> (AppState, Arc<InMemoryProviderCatalogReadRepository>, String) {
provider_oauth_state("codex", auth_config, access_token)
}
fn provider_oauth_state(
provider_type: &str,
auth_config: &serde_json::Value,
access_token: &str,
) -> (AppState, Arc<InMemoryProviderCatalogReadRepository>, String) {
let mut provider = sample_provider();
provider.provider_type = "codex".to_string();
provider.provider_type = provider_type.to_string();
let encrypted_auth_config =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, &auth_config.to_string())
.expect("auth config should encrypt");
@@ -3490,7 +3503,7 @@ mod tests {
let key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"Codex OAuth".to_string(),
format!("{provider_type} OAuth"),
"oauth".to_string(),
None,
true,
@@ -4564,6 +4577,78 @@ mod tests {
);
}
#[tokio::test]
async fn antigravity_refresh_entry_persists_tokens_and_expiry() {
let initial_config = json!({
"provider_type": "antigravity",
"refreshToken": "legacy-refresh-token",
"expires_at": 1,
});
let (state, repository, _) =
provider_oauth_state("antigravity", &initial_config, "stale-access-token");
let transport = state
.read_provider_transport_snapshot("provider-1", "endpoint-1", "key-1")
.await
.expect("transport should load")
.expect("transport should exist");
let expected_credential_fence = state
.capture_provider_transport_credential_fence(&transport)
.await
.expect("credential fence should load")
.expect("credential fence should match");
let expires_at = 4_102_555_900;
let refreshed_entry = crate::provider_transport::CachedOAuthEntry {
provider_type: "antigravity".to_string(),
auth_header_name: "authorization".to_string(),
auth_header_value: "Bearer fresh-access-token".to_string(),
expires_at_unix_secs: Some(expires_at),
metadata: Some(json!({
"provider_type": "antigravity",
"refresh_token": "legacy-refresh-token",
"expires_at": expires_at,
})),
source_fingerprint: None,
};
state
.persist_local_oauth_refresh_entry(
&transport,
&refreshed_entry,
Some(&expected_credential_fence),
)
.await
.expect("Antigravity refresh should persist");
let stored = repository
.list_keys_by_ids(&["key-1".to_string()])
.await
.expect("key should reload")
.pop()
.expect("key should remain");
let access_token = decrypt_python_fernet_ciphertext(
DEVELOPMENT_ENCRYPTION_KEY,
stored
.encrypted_api_key
.as_deref()
.expect("access token should persist"),
)
.expect("access token should decrypt");
let auth_config = decrypt_python_fernet_ciphertext(
DEVELOPMENT_ENCRYPTION_KEY,
stored
.encrypted_auth_config
.as_deref()
.expect("auth config should persist"),
)
.expect("auth config should decrypt");
let auth_config: serde_json::Value =
serde_json::from_str(&auth_config).expect("auth config should parse");
assert_eq!(access_token, "fresh-access-token");
assert_eq!(auth_config["refresh_token"], json!("legacy-refresh-token"));
assert_eq!(stored.expires_at_unix_secs, Some(expires_at));
}
#[tokio::test]
async fn agent_auth_config_fence_rejects_metadata_only_rewrite() {
let initial_config = json!({
@@ -629,6 +629,7 @@ async fn gateway_pool_list_includes_usage_totals_and_nullable_lru_score() {
"sk-usage",
);
key.name = "usage key".to_string();
key.concurrent_limit = Some(5);
key.request_count = Some(1566);
key.total_tokens = 187_327_321;
key.total_cost_usd = 93.1319297;
@@ -661,6 +662,7 @@ async fn gateway_pool_list_includes_usage_totals_and_nullable_lru_score() {
.expect("json body should parse");
let keys = payload["keys"].as_array().expect("keys should be array");
assert_eq!(keys.len(), 1);
assert_eq!(keys[0]["concurrent_limit"], json!(5));
assert_eq!(keys[0]["request_count"], json!(1566));
assert_eq!(keys[0]["total_tokens"], json!(187_327_321u64));
assert_eq!(keys[0]["total_cost_usd"], json!("93.13192970"));
@@ -3643,6 +3645,7 @@ async fn gateway_batch_updates_shared_pool_key_configuration() {
"api_formats": ["openai:responses"],
"internal_priority": 7,
"rpm_limit": null,
"concurrent_limit": 6,
"auto_fetch_models": false,
"allowed_models": ["gpt-5.6-sol", "gpt-5.6-luna"],
"locked_models": [],
@@ -3668,6 +3671,7 @@ async fn gateway_batch_updates_shared_pool_key_configuration() {
assert_eq!(key.allow_auth_channel_mismatch_formats, Some(json!([])));
assert_eq!(key.internal_priority, 7);
assert_eq!(key.rpm_limit, None);
assert_eq!(key.concurrent_limit, Some(6));
assert_eq!(key.learned_rpm_limit, None);
assert!(!key.auto_fetch_models);
assert_eq!(
@@ -55,6 +55,9 @@ async fn resolve_wallet_auth_gate_with_cache(
None => WalletAccessDecision::wallet_unavailable(None),
};
if !auth_snapshot.api_key_is_standalone {
let wallet_is_unlimited = wallet
.as_ref()
.is_some_and(|wallet| wallet.limit_mode.eq_ignore_ascii_case("unlimited"));
let quota = if use_cache {
state
.find_user_daily_quota_availability_for_auth(&auth_snapshot.user_id)
@@ -71,7 +74,11 @@ async fn resolve_wallet_auth_gate_with_cache(
quota.remaining_usd,
))));
}
if decision.failure.is_none() && !quota.allow_wallet_overage && !has_remaining_quota {
if !wallet_is_unlimited
&& decision.failure.is_none()
&& !quota.allow_wallet_overage
&& !has_remaining_quota
{
return Ok(Some(WalletAccessDecision::balance_denied(Some(0.0))));
}
}
@@ -264,6 +271,23 @@ mod tests {
assert_eq!(decision.remaining, Some(4.0));
}
#[tokio::test]
async fn unlimited_wallet_ignores_exhausted_non_overage_quota() {
let mut wallet = empty_user_wallet();
wallet.limit_mode = "unlimited".to_string();
let state = state_with_wallet_and_quota(wallet, Some(quota_availability(10.0, 0.0, false)));
let auth_snapshot = ordinary_user_api_key_snapshot();
let decision = resolve_wallet_auth_gate(&state, &auth_snapshot)
.await
.expect("wallet gate should resolve")
.expect("wallet gate should return a decision");
assert!(decision.allowed);
assert_eq!(decision.failure, None);
assert_eq!(decision.remaining, None);
}
#[tokio::test]
async fn disabled_auth_capacity_cache_still_gates_wallet_reads() {
let mut state = state_with_wallet_and_quota(empty_user_wallet(), None);