mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-05 00:47:48 +08:00
fix(gateway): add provider key concurrency and cache affinity modes
This commit is contained in:
@@ -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() {
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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());
|
||||
|
||||
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user