perf(gateway): scale request hot paths for 20k streams

Shard and singleflight hot-path caches, batch and prioritize candidate and usage lifecycle persistence, and extend database and pressure-test instrumentation for 20k concurrent streams.
This commit is contained in:
elky
2026-07-22 02:11:08 +08:00
parent 7756c0913f
commit fc92c4f431
124 changed files with 36325 additions and 3217 deletions
+637 -196
View File
@@ -1,5 +1,6 @@
use std::collections::BTreeMap;
use std::sync::LazyLock;
use std::collections::{BTreeMap, HashMap};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, LazyLock, Mutex as StdMutex, Weak};
use std::time::Duration;
use aether_admin::provider::quota as admin_provider_quota_pure;
@@ -8,6 +9,10 @@ use aether_contracts::{ExecutionPlan, ExecutionTelemetry};
use aether_data_contracts::repository::pool_scores::{
PoolMemberHardState, PoolMemberIdentity, PoolMemberScheduleFeedback,
};
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
ProviderCatalogKeyHealthStateUpdate,
};
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,
@@ -17,6 +22,7 @@ use aether_usage_runtime::{
GatewayStreamReportRequest, GatewaySyncReportRequest, TerminalUsageOutcome,
};
use serde_json::Value;
use tokio::sync::Mutex as TokioMutex;
use tracing::warn;
use super::{
@@ -44,21 +50,116 @@ use crate::{
const POOL_SCORE_FEEDBACK_GATE_MAX_ENTRIES: usize = 50_000;
const HEALTH_SUCCESS_PERSIST_GATE_MAX_ENTRIES: usize = 50_000;
const ADAPTIVE_SUCCESS_PERSIST_GATE_MAX_ENTRIES: usize = 50_000;
const POOL_SCORE_SUCCESS_FEEDBACK_MIN_INTERVAL_ENV: &str =
"AETHER_GATEWAY_POOL_SCORE_SUCCESS_FEEDBACK_MIN_INTERVAL_SECS";
const POOL_SCORE_FAILURE_FEEDBACK_MIN_INTERVAL_ENV: &str =
"AETHER_GATEWAY_POOL_SCORE_FAILURE_FEEDBACK_MIN_INTERVAL_SECS";
const HEALTH_SUCCESS_PERSIST_MIN_INTERVAL_ENV: &str =
"AETHER_GATEWAY_PROVIDER_KEY_HEALTH_SUCCESS_PERSIST_MIN_INTERVAL_SECS";
const ADAPTIVE_SUCCESS_PERSIST_MIN_INTERVAL_ENV: &str =
"AETHER_GATEWAY_PROVIDER_KEY_ADAPTIVE_SUCCESS_PERSIST_MIN_INTERVAL_SECS";
const DEFAULT_POOL_SCORE_SUCCESS_FEEDBACK_MIN_INTERVAL_SECS: u64 = 5;
const DEFAULT_POOL_SCORE_FAILURE_FEEDBACK_MIN_INTERVAL_SECS: u64 = 1;
const DEFAULT_HEALTH_SUCCESS_PERSIST_MIN_INTERVAL_SECS: u64 = 5;
const DEFAULT_ADAPTIVE_SUCCESS_PERSIST_MIN_INTERVAL_SECS: u64 = 5;
const MAX_POOL_SCORE_FEEDBACK_MIN_INTERVAL_SECS: u64 = 300;
const PROVIDER_KEY_EFFECT_LOCK_PRUNE_THRESHOLD: usize = 8_192;
// Same-process writers are serialized by the per-key lock. Keep remote-writer
// retries bounded so request/report completion cannot accumulate a long DB tail.
const PROVIDER_KEY_STATE_CAS_MAX_ATTEMPTS: usize = 4;
#[derive(Debug)]
struct ProviderKeyEffectLockPoolState {
entries: HashMap<String, Weak<TokioMutex<()>>>,
accesses_since_prune: usize,
next_growth_prune_at: usize,
#[cfg(test)]
prune_count: usize,
}
impl ProviderKeyEffectLockPoolState {
fn new(min_prune_threshold: usize) -> Self {
Self {
entries: HashMap::new(),
accesses_since_prune: 0,
next_growth_prune_at: min_prune_threshold,
#[cfg(test)]
prune_count: 0,
}
}
}
#[derive(Debug)]
struct ProviderKeyEffectLockPool {
state: StdMutex<ProviderKeyEffectLockPoolState>,
min_prune_threshold: usize,
}
impl Default for ProviderKeyEffectLockPool {
fn default() -> Self {
Self::new(PROVIDER_KEY_EFFECT_LOCK_PRUNE_THRESHOLD)
}
}
impl ProviderKeyEffectLockPool {
fn new(min_prune_threshold: usize) -> Self {
let min_prune_threshold = min_prune_threshold.max(1);
Self {
state: StdMutex::new(ProviderKeyEffectLockPoolState::new(min_prune_threshold)),
min_prune_threshold,
}
}
fn lock_for(&self, key_id: &str) -> Arc<TokioMutex<()>> {
let mut state = self
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
state.accesses_since_prune = state.accesses_since_prune.saturating_add(1);
let entry_count = state.entries.len();
let growth_prune_due = entry_count >= state.next_growth_prune_at;
let maintenance_prune_due = entry_count >= self.min_prune_threshold
&& state.accesses_since_prune >= entry_count.max(self.min_prune_threshold);
if growth_prune_due || maintenance_prune_due {
self.prune_inactive_locks(&mut state);
}
if let Some(existing) = state.entries.get(key_id).and_then(Weak::upgrade) {
return existing;
}
let lock = Arc::new(TokioMutex::new(()));
state
.entries
.insert(key_id.to_string(), Arc::downgrade(&lock));
lock
}
fn prune_inactive_locks(&self, state: &mut ProviderKeyEffectLockPoolState) {
state.entries.retain(|_, lock| lock.strong_count() > 0);
let active_entries = state.entries.len();
state.next_growth_prune_at = if active_entries < self.min_prune_threshold {
self.min_prune_threshold
} else {
active_entries.saturating_mul(2)
};
state.accesses_since_prune = 0;
#[cfg(test)]
{
state.prune_count = state.prune_count.saturating_add(1);
}
}
}
static POOL_SCORE_FEEDBACK_GATE: LazyLock<ExpiringMap<String, ()>> =
LazyLock::new(ExpiringMap::new);
static HEALTH_SUCCESS_PERSIST_GATE: LazyLock<ExpiringMap<String, ()>> =
LazyLock::new(ExpiringMap::new);
static ADAPTIVE_SUCCESS_PERSIST_GATE: LazyLock<ExpiringMap<String, u64>> =
LazyLock::new(ExpiringMap::new);
static ADAPTIVE_SUCCESS_PERSIST_GATE_NEXT_TOKEN: AtomicU64 = AtomicU64::new(1);
static PROVIDER_KEY_EFFECT_LOCKS: LazyLock<ProviderKeyEffectLockPool> =
LazyLock::new(ProviderKeyEffectLockPool::default);
static POOL_SCORE_SUCCESS_FEEDBACK_MIN_INTERVAL: LazyLock<Duration> = LazyLock::new(|| {
pool_score_feedback_interval_from_env(
POOL_SCORE_SUCCESS_FEEDBACK_MIN_INTERVAL_ENV,
@@ -77,6 +178,12 @@ static HEALTH_SUCCESS_PERSIST_MIN_INTERVAL: LazyLock<Duration> = LazyLock::new(|
DEFAULT_HEALTH_SUCCESS_PERSIST_MIN_INTERVAL_SECS,
)
});
static ADAPTIVE_SUCCESS_PERSIST_MIN_INTERVAL: LazyLock<Duration> = LazyLock::new(|| {
pool_score_feedback_interval_from_env(
ADAPTIVE_SUCCESS_PERSIST_MIN_INTERVAL_ENV,
DEFAULT_ADAPTIVE_SUCCESS_PERSIST_MIN_INTERVAL_SECS,
)
});
#[derive(Debug, Clone, Copy)]
pub(crate) struct LocalExecutionEffectContext<'a> {
@@ -495,15 +602,9 @@ async fn record_adaptive_rate_limit_effect(
context: LocalExecutionEffectContext<'_>,
effect: LocalAdaptiveRateLimitEffect<'_>,
) {
let effect_lock = PROVIDER_KEY_EFFECT_LOCKS.lock_for(&context.plan.key_id);
let _effect_guard = effect_lock.lock().await;
let observed_at_unix_secs = current_unix_secs();
let Some(current_key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&context.plan.key_id))
.await
.ok()
.and_then(|mut keys| keys.drain(..).next())
else {
return;
};
let current_rpm = state
.read_recent_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT)
.await
@@ -515,38 +616,88 @@ async fn record_adaptive_rate_limit_effect(
observed_at_unix_secs,
) as u32
});
let Some(projection) = project_local_adaptive_rate_limit(
&current_key,
effect.classification,
effect.status_code,
current_rpm,
effect.headers,
observed_at_unix_secs,
) else {
return;
};
let mut updated_key = current_key.clone();
updated_key.rpm_429_count = Some(projection.rpm_429_count);
updated_key.learned_rpm_limit = projection.learned_rpm_limit;
updated_key.last_429_at_unix_secs = Some(projection.last_429_at_unix_secs);
updated_key.last_429_type = Some(projection.last_429_type);
updated_key.adjustment_history = projection.adjustment_history;
updated_key.utilization_samples = projection.utilization_samples;
updated_key.last_probe_increase_at_unix_secs = projection.last_probe_increase_at_unix_secs;
updated_key.last_rpm_peak = projection.last_rpm_peak;
updated_key.status_snapshot = Some(projection.status_snapshot);
updated_key.updated_at_unix_secs = Some(observed_at_unix_secs);
if let Err(err) = state
.update_provider_catalog_key_runtime_state(&updated_key)
.await
{
warn!(
"gateway orchestration effects: failed to persist adaptive rate-limit projection for provider {} endpoint {} key {}: {:?}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id, err
);
for _ in 0..PROVIDER_KEY_STATE_CAS_MAX_ATTEMPTS {
let Some(current_key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&context.plan.key_id))
.await
.ok()
.and_then(|mut keys| keys.drain(..).next())
else {
return;
};
let Some(projection) = project_local_adaptive_rate_limit(
&current_key,
effect.classification,
effect.status_code,
current_rpm,
effect.headers,
observed_at_unix_secs,
) else {
return;
};
let expected = ProviderCatalogKeyAdaptiveState::from(&current_key);
let mut next = expected.clone();
next.rpm_429_count = Some(projection.rpm_429_count);
next.learned_rpm_limit = projection.learned_rpm_limit;
next.last_429_at_unix_secs = Some(projection.last_429_at_unix_secs);
next.last_429_type = Some(projection.last_429_type);
next.adjustment_history = projection.adjustment_history;
next.utilization_samples = projection.utilization_samples;
next.last_probe_increase_at_unix_secs = projection.last_probe_increase_at_unix_secs;
next.last_rpm_peak = projection.last_rpm_peak;
let update = ProviderCatalogKeyAdaptiveStateUpdate {
key_id: context.plan.key_id.clone(),
expected,
next,
status_snapshot_patch: adaptive_status_snapshot_patch(&projection.status_snapshot),
updated_at_unix_secs: Some(observed_at_unix_secs),
};
provider_key_adaptive_success_persist_gate_reset(&context.plan.key_id);
match state
.compare_and_update_provider_catalog_key_adaptive_state(&update)
.await
{
Ok(true) => return,
Ok(false) => tokio::task::yield_now().await,
Err(err) => {
warn!(
"gateway orchestration effects: failed to persist adaptive rate-limit projection for provider {} endpoint {} key {}: {:?}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id, err
);
return;
}
}
}
warn!(
"gateway orchestration effects: adaptive rate-limit CAS retries exhausted for provider {} endpoint {} key {}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id
);
}
fn adaptive_status_snapshot_patch(status_snapshot: &Value) -> Value {
const OWNED_FIELDS: [&str; 6] = [
"observation_count",
"header_observation_count",
"latest_upstream_limit",
"learning_confidence",
"enforcement_active",
"known_boundary",
];
let Some(snapshot) = status_snapshot.as_object() else {
return serde_json::json!({});
};
Value::Object(
OWNED_FIELDS
.into_iter()
.filter_map(|field| {
snapshot
.get(field)
.cloned()
.map(|value| (field.to_string(), value))
})
.collect(),
)
}
async fn record_adaptive_success_effect(
@@ -571,6 +722,19 @@ async fn record_adaptive_success_effect(
{
return;
}
let Some(gate_token) = provider_key_adaptive_success_persist_gate_admit(&context.plan.key_id)
else {
return;
};
let effect_lock = PROVIDER_KEY_EFFECT_LOCKS.lock_for(&context.plan.key_id);
let _effect_guard = effect_lock.lock().await;
if !provider_key_adaptive_success_persist_gate_admission_is_current(
&context.plan.key_id,
gate_token,
) {
return;
}
let Some(recent_candidates) = state
.read_recent_request_candidates(ADAPTIVE_RPM_RECENT_CANDIDATE_LIMIT)
.await
@@ -583,29 +747,101 @@ async fn record_adaptive_success_effect(
&context.plan.key_id,
observed_at_unix_secs,
) as u32;
let Some(projection) =
project_local_adaptive_success(&current_key, current_rpm, observed_at_unix_secs)
else {
return;
};
let mut updated_key = current_key.clone();
updated_key.learned_rpm_limit = projection.learned_rpm_limit;
updated_key.adjustment_history = projection.adjustment_history;
updated_key.utilization_samples = projection.utilization_samples;
updated_key.last_probe_increase_at_unix_secs = projection.last_probe_increase_at_unix_secs;
updated_key.status_snapshot = Some(projection.status_snapshot);
updated_key.updated_at_unix_secs = Some(observed_at_unix_secs);
if let Err(err) = state
.update_provider_catalog_key_runtime_state(&updated_key)
.await
{
warn!(
"gateway orchestration effects: failed to persist adaptive success projection for provider {} endpoint {} key {}: {:?}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id, err
);
for _ in 0..PROVIDER_KEY_STATE_CAS_MAX_ATTEMPTS {
let Some(current_key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&context.plan.key_id))
.await
.ok()
.and_then(|mut keys| keys.drain(..).next())
else {
return;
};
if current_key.rpm_limit.is_some()
|| current_key
.learned_rpm_limit
.filter(|value| *value > 0)
.is_none()
{
return;
}
let Some(projection) =
project_local_adaptive_success(&current_key, current_rpm, observed_at_unix_secs)
else {
return;
};
let expected = ProviderCatalogKeyAdaptiveState::from(&current_key);
let mut next = expected.clone();
next.learned_rpm_limit = projection.learned_rpm_limit;
next.adjustment_history = projection.adjustment_history;
next.utilization_samples = projection.utilization_samples;
next.last_probe_increase_at_unix_secs = projection.last_probe_increase_at_unix_secs;
let update = ProviderCatalogKeyAdaptiveStateUpdate {
key_id: context.plan.key_id.clone(),
expected,
next,
status_snapshot_patch: adaptive_status_snapshot_patch(&projection.status_snapshot),
updated_at_unix_secs: Some(observed_at_unix_secs),
};
match state
.compare_and_update_provider_catalog_key_adaptive_state(&update)
.await
{
Ok(true) => return,
Ok(false) => tokio::task::yield_now().await,
Err(err) => {
warn!(
"gateway orchestration effects: failed to persist adaptive success projection for provider {} endpoint {} key {}: {:?}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id, err
);
return;
}
}
}
warn!(
"gateway orchestration effects: adaptive success CAS retries exhausted for provider {} endpoint {} key {}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id
);
}
fn provider_key_adaptive_success_persist_gate_admit(key_id: &str) -> Option<u64> {
if cfg!(test) {
return Some(0);
}
let interval = *ADAPTIVE_SUCCESS_PERSIST_MIN_INTERVAL;
if interval.is_zero() {
return Some(0);
}
let token = ADAPTIVE_SUCCESS_PERSIST_GATE_NEXT_TOKEN.fetch_add(1, Ordering::Relaxed);
ADAPTIVE_SUCCESS_PERSIST_GATE
.insert_if_absent_fresh(
provider_key_adaptive_success_persist_gate_key(key_id),
token,
interval,
ADAPTIVE_SUCCESS_PERSIST_GATE_MAX_ENTRIES,
)
.then_some(token)
}
fn provider_key_adaptive_success_persist_gate_admission_is_current(
key_id: &str,
token: u64,
) -> bool {
if cfg!(test) || ADAPTIVE_SUCCESS_PERSIST_MIN_INTERVAL.is_zero() {
return true;
}
ADAPTIVE_SUCCESS_PERSIST_GATE.get_fresh(
&provider_key_adaptive_success_persist_gate_key(key_id),
*ADAPTIVE_SUCCESS_PERSIST_MIN_INTERVAL,
) == Some(token)
}
fn provider_key_adaptive_success_persist_gate_reset(key_id: &str) {
ADAPTIVE_SUCCESS_PERSIST_GATE.remove(&provider_key_adaptive_success_persist_gate_key(key_id));
}
fn provider_key_adaptive_success_persist_gate_key(key_id: &str) -> String {
format!("adaptive-success:{key_id}")
}
async fn record_health_failure_effect(
@@ -618,65 +854,73 @@ async fn record_health_failure_effect(
return;
}
let Some(current_key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&context.plan.key_id))
.await
.ok()
.and_then(|mut keys| keys.drain(..).next())
else {
return;
};
let effect_lock = PROVIDER_KEY_EFFECT_LOCKS.lock_for(&context.plan.key_id);
let _effect_guard = effect_lock.lock().await;
let is_pool_provider = local_execution_plan_uses_pool(state, context.plan).await;
let observed_at_unix_secs = current_unix_secs();
let Some(health_by_format) = project_local_failure_health(
current_key.health_by_format.as_ref(),
api_format,
effect.classification,
effect.status_code,
observed_at_unix_secs,
) else {
return;
};
let consecutive_failures = health_by_format
.get(api_format)
.and_then(|value| value.get("consecutive_failures"))
.and_then(Value::as_u64)
.unwrap_or(0);
let circuit_breaker_update_owned = if is_pool_provider {
None
} else {
project_local_key_circuit_failure(
current_key.circuit_breaker_by_format.as_ref(),
api_format,
observed_at_unix_secs,
consecutive_failures,
current_key.max_probe_interval_minutes,
)
};
let circuit_breaker_update = if is_pool_provider {
None
} else {
circuit_breaker_update_owned
.as_ref()
.or(current_key.circuit_breaker_by_format.as_ref())
};
provider_key_health_success_persist_gate_reset(&context.plan.key_id, api_format);
if let Err(err) = state
.update_provider_catalog_key_success_health_state(
&context.plan.key_id,
current_key.is_active,
Some(&health_by_format),
circuit_breaker_update,
)
.await
{
warn!(
"gateway orchestration effects: failed to persist health failure projection for provider {} endpoint {} key {}: {:?}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id, err
);
for _ in 0..PROVIDER_KEY_STATE_CAS_MAX_ATTEMPTS {
let Some(current_key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&context.plan.key_id))
.await
.ok()
.and_then(|mut keys| keys.drain(..).next())
else {
return;
};
let Some(health_by_format) = project_local_failure_health(
current_key.health_by_format.as_ref(),
api_format,
effect.classification,
effect.status_code,
observed_at_unix_secs,
) else {
return;
};
let consecutive_failures = health_by_format
.get(api_format)
.and_then(|value| value.get("consecutive_failures"))
.and_then(Value::as_u64)
.unwrap_or(0);
let circuit_breaker_by_format = if is_pool_provider {
None
} else {
project_local_key_circuit_failure(
current_key.circuit_breaker_by_format.as_ref(),
api_format,
observed_at_unix_secs,
consecutive_failures,
current_key.max_probe_interval_minutes,
)
.or_else(|| current_key.circuit_breaker_by_format.clone())
};
let update = ProviderCatalogKeyHealthStateUpdate {
key_id: context.plan.key_id.clone(),
expected_health_by_format: current_key.health_by_format,
expected_circuit_breaker_by_format: current_key.circuit_breaker_by_format,
health_by_format: Some(health_by_format),
circuit_breaker_by_format,
};
match state
.compare_and_update_provider_catalog_key_health_state(&update)
.await
{
Ok(true) => return,
Ok(false) => tokio::task::yield_now().await,
Err(err) => {
warn!(
"gateway orchestration effects: failed to persist health failure projection for provider {} endpoint {} key {}: {:?}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id, err
);
return;
}
}
}
warn!(
"gateway orchestration effects: health failure CAS retries exhausted for provider {} endpoint {} key {}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id
);
}
async fn record_health_success_effect(
@@ -691,66 +935,86 @@ async fn record_health_success_effect(
return;
}
let Some(current_key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&context.plan.key_id))
.await
.ok()
.and_then(|mut keys| keys.drain(..).next())
else {
return;
};
// Health updates replace both JSON snapshots in one write. Serialize the success
// read/project/write with failure and circuit-clear effects for this provider key so a
// stale success snapshot cannot overwrite a newer failure counter or open circuit.
let effect_lock = PROVIDER_KEY_EFFECT_LOCKS.lock_for(&context.plan.key_id);
let _effect_guard = effect_lock.lock().await;
let is_pool_provider = local_execution_plan_uses_pool(state, context.plan).await;
let Some(health_by_format) =
project_local_success_health(current_key.health_by_format.as_ref(), api_format)
else {
return;
};
let circuit_breaker_update_owned = if is_pool_provider {
None
} else {
current_key
.circuit_breaker_by_format
.as_ref()
.and_then(|current| project_local_key_circuit_closed(Some(current), api_format))
};
if current_key.health_by_format.as_ref() == Some(&health_by_format)
&& ((is_pool_provider && current_key.circuit_breaker_by_format.is_none())
|| (!is_pool_provider
&& circuit_breaker_update_owned.as_ref()
== current_key.circuit_breaker_by_format.as_ref()))
{
return;
}
let circuit_breaker_update = if is_pool_provider {
None
} else {
circuit_breaker_update_owned
.as_ref()
.or(current_key.circuit_breaker_by_format.as_ref())
};
let mut persist_gate_checked = false;
if !provider_key_health_success_persist_gate_allows(
&context.plan.key_id,
api_format,
circuit_breaker_update_owned.is_some(),
) {
return;
}
if let Err(err) = state
.update_provider_catalog_key_health_state(
&context.plan.key_id,
current_key.is_active,
Some(&health_by_format),
circuit_breaker_update,
)
.await
{
warn!(
"gateway orchestration effects: failed to persist health success projection for provider {} endpoint {} key {}: {:?}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id, err
);
for _ in 0..PROVIDER_KEY_STATE_CAS_MAX_ATTEMPTS {
let Some(current_key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&context.plan.key_id))
.await
.ok()
.and_then(|mut keys| keys.drain(..).next())
else {
return;
};
let Some(health_by_format) =
project_local_success_health(current_key.health_by_format.as_ref(), api_format)
else {
return;
};
let circuit_breaker_update_owned = if is_pool_provider {
None
} else {
current_key
.circuit_breaker_by_format
.as_ref()
.and_then(|current| project_local_key_circuit_closed(Some(current), api_format))
};
if current_key.health_by_format.as_ref() == Some(&health_by_format)
&& ((is_pool_provider && current_key.circuit_breaker_by_format.is_none())
|| (!is_pool_provider
&& circuit_breaker_update_owned.as_ref()
== current_key.circuit_breaker_by_format.as_ref()))
{
return;
}
if !persist_gate_checked {
if !provider_key_health_success_persist_gate_allows(
&context.plan.key_id,
api_format,
circuit_breaker_update_owned.is_some(),
) {
return;
}
persist_gate_checked = true;
}
let circuit_breaker_by_format = if is_pool_provider {
None
} else {
circuit_breaker_update_owned.or_else(|| current_key.circuit_breaker_by_format.clone())
};
let update = ProviderCatalogKeyHealthStateUpdate {
key_id: context.plan.key_id.clone(),
expected_health_by_format: current_key.health_by_format,
expected_circuit_breaker_by_format: current_key.circuit_breaker_by_format,
health_by_format: Some(health_by_format),
circuit_breaker_by_format,
};
match state
.compare_and_update_provider_catalog_key_health_state(&update)
.await
{
Ok(true) => return,
Ok(false) => tokio::task::yield_now().await,
Err(err) => {
warn!(
"gateway orchestration effects: failed to persist health success projection for provider {} endpoint {} key {}: {:?}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id, err
);
return;
}
}
}
warn!(
"gateway orchestration effects: health success CAS retries exhausted for provider {} endpoint {} key {}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id
);
}
fn provider_key_health_success_persist_gate_allows(
@@ -871,32 +1135,47 @@ async fn clear_pool_key_circuit_breaker(
state: &AppState,
context: LocalExecutionEffectContext<'_>,
) {
let Some(current_key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&context.plan.key_id))
.await
.ok()
.and_then(|mut keys| keys.drain(..).next())
else {
return;
};
if current_key.circuit_breaker_by_format.is_none() {
return;
}
let effect_lock = PROVIDER_KEY_EFFECT_LOCKS.lock_for(&context.plan.key_id);
let _effect_guard = effect_lock.lock().await;
if let Err(err) = state
.update_provider_catalog_key_health_state(
&context.plan.key_id,
current_key.is_active,
current_key.health_by_format.as_ref(),
None,
)
.await
{
warn!(
"gateway orchestration effects: failed to clear pool key circuit for provider {} endpoint {} key {}: {:?}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id, err
);
for _ in 0..PROVIDER_KEY_STATE_CAS_MAX_ATTEMPTS {
let Some(current_key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&context.plan.key_id))
.await
.ok()
.and_then(|mut keys| keys.drain(..).next())
else {
return;
};
if current_key.circuit_breaker_by_format.is_none() {
return;
}
let update = ProviderCatalogKeyHealthStateUpdate {
key_id: context.plan.key_id.clone(),
expected_health_by_format: current_key.health_by_format.clone(),
expected_circuit_breaker_by_format: current_key.circuit_breaker_by_format,
health_by_format: current_key.health_by_format,
circuit_breaker_by_format: None,
};
match state
.compare_and_update_provider_catalog_key_health_state(&update)
.await
{
Ok(true) => return,
Ok(false) => tokio::task::yield_now().await,
Err(err) => {
warn!(
"gateway orchestration effects: failed to clear pool key circuit for provider {} endpoint {} key {}: {:?}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id, err
);
return;
}
}
}
warn!(
"gateway orchestration effects: clear pool key circuit CAS retries exhausted for provider {} endpoint {} key {}",
context.plan.provider_id, context.plan.endpoint_id, context.plan.key_id
);
}
async fn record_oauth_invalidation_effect(
@@ -1292,6 +1571,7 @@ mod tests {
LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect,
LocalExecutionEffect, LocalExecutionEffectContext, LocalHealthFailureEffect,
LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalPoolErrorEffect,
ProviderKeyEffectLockPool,
};
use crate::data::{GatewayDataConfig, GatewayDataState};
use crate::orchestration::LocalFailoverClassification;
@@ -1363,6 +1643,37 @@ mod tests {
));
}
#[test]
fn provider_key_effect_lock_pool_prunes_geometrically_and_keeps_active_locks() {
let pool = ProviderKeyEffectLockPool::new(4);
let locks = (0..10)
.map(|index| pool.lock_for(&format!("key-{index}")))
.collect::<Vec<_>>();
{
let state = pool.state.lock().expect("effect lock pool should lock");
assert_eq!(state.entries.len(), 10);
assert_eq!(state.prune_count, 2);
assert_eq!(state.next_growth_prune_at, 16);
}
for (index, expected) in locks.iter().enumerate() {
let existing = pool.lock_for(&format!("key-{index}"));
assert!(Arc::ptr_eq(expected, &existing));
}
let hot_lock = Arc::clone(&locks[0]);
drop(locks);
for _ in 0..32 {
let existing = pool.lock_for("key-0");
assert!(Arc::ptr_eq(&hot_lock, &existing));
}
let state = pool.state.lock().expect("effect lock pool should lock");
assert_eq!(state.entries.len(), 1);
assert!(state.prune_count >= 3);
}
fn session_affinity() -> ClientSessionAffinity {
ClientSessionAffinity::new(
Some("generic".to_string()),
@@ -2538,6 +2849,45 @@ mod tests {
);
}
#[tokio::test]
async fn runtime_health_failure_does_not_reactivate_admin_disabled_key() {
let mut key = sample_health_key();
key.is_active = false;
let state = health_state_with_key(key);
let plan = sample_plan();
apply_local_execution_effect(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: None,
},
LocalExecutionEffect::HealthFailure(LocalHealthFailureEffect {
status_code: 503,
classification: LocalFailoverClassification::RetryUpstreamFailure,
}),
)
.await;
let stored_key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await
.expect("provider catalog keys should load")
.into_iter()
.next()
.expect("stored key should exist");
assert!(!stored_key.is_active);
assert_eq!(
stored_key
.health_by_format
.as_ref()
.and_then(|value| value.get("openai:chat"))
.and_then(|value| value.get("consecutive_failures"))
.and_then(Value::as_u64),
Some(1)
);
}
#[tokio::test]
async fn health_failure_opens_circuit_after_eight_consecutive_failures() {
let state = health_state();
@@ -2583,6 +2933,59 @@ mod tests {
);
}
#[tokio::test]
async fn concurrent_health_failures_for_one_key_do_not_lose_updates() {
let state = health_state();
let plan = sample_plan();
let mut tasks = tokio::task::JoinSet::new();
for _ in 0..8 {
let state = state.clone();
let plan = plan.clone();
tasks.spawn(async move {
apply_local_execution_effect(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: None,
},
LocalExecutionEffect::HealthFailure(LocalHealthFailureEffect {
status_code: 503,
classification: LocalFailoverClassification::RetryUpstreamFailure,
}),
)
.await;
});
}
while let Some(result) = tasks.join_next().await {
result.expect("health failure task should complete");
}
let stored_key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await
.expect("provider catalog keys should load")
.into_iter()
.next()
.expect("stored key should exist");
let circuit = stored_key
.circuit_breaker_by_format
.as_ref()
.and_then(|value| value.get("openai:chat"))
.expect("format circuit should be stored");
assert_eq!(
stored_key
.health_by_format
.as_ref()
.and_then(|value| value.get("openai:chat"))
.and_then(|value| value.get("consecutive_failures"))
.and_then(Value::as_u64),
Some(8)
);
assert_eq!(circuit["open"], json!(true));
assert_eq!(circuit["reason"], json!("consecutive_failures_8"));
}
#[tokio::test]
async fn pool_health_failure_does_not_open_key_circuit_after_eight_consecutive_failures() {
let state = pool_health_state();
@@ -2879,6 +3282,44 @@ mod tests {
);
}
#[tokio::test]
async fn adaptive_rate_limit_effect_preserves_quota_status_owned_by_reports() {
let mut key = sample_adaptive_key();
key.status_snapshot = Some(json!({
"quota": {"remaining": 9},
"oauth": {"invalid": false},
"observation_count": 0
}));
let state = health_state_with_key(key);
let plan = sample_plan();
apply_local_execution_effect(
&state,
LocalExecutionEffectContext {
plan: &plan,
report_context: None,
},
LocalExecutionEffect::AdaptiveRateLimit(LocalAdaptiveRateLimitEffect {
status_code: 429,
classification: LocalFailoverClassification::RetryUpstreamFailure,
headers: None,
}),
)
.await;
let stored_key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await
.expect("provider catalog keys should load")
.into_iter()
.next()
.expect("stored key should exist");
let status = stored_key.status_snapshot.expect("status should exist");
assert_eq!(status["quota"], json!({"remaining":9}));
assert_eq!(status["oauth"], json!({"invalid":false}));
assert_eq!(status["observation_count"], json!(1));
}
#[tokio::test]
async fn adaptive_rate_limit_effect_ignores_fixed_limit_key() {
let state = fixed_limit_state();