mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 18:59:50 +08:00
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:
@@ -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(
|
||||
¤t_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(
|
||||
¤t_key,
|
||||
effect.classification,
|
||||
effect.status_code,
|
||||
current_rpm,
|
||||
effect.headers,
|
||||
observed_at_unix_secs,
|
||||
) else {
|
||||
return;
|
||||
};
|
||||
let expected = ProviderCatalogKeyAdaptiveState::from(¤t_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(¤t_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(¤t_key, current_rpm, observed_at_unix_secs)
|
||||
else {
|
||||
return;
|
||||
};
|
||||
let expected = ProviderCatalogKeyAdaptiveState::from(¤t_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();
|
||||
|
||||
@@ -3,6 +3,7 @@ use std::sync::{Mutex, OnceLock};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use aether_admin::provider::quota as admin_provider_quota_pure;
|
||||
use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyRuntimeMetadataUpdate;
|
||||
use aether_provider_pool::grok_quota_window_key_for_model;
|
||||
use aether_usage_runtime::{
|
||||
extract_gemini_file_mapping_entries, gemini_file_mapping_cache_key, normalize_gemini_file_name,
|
||||
@@ -22,6 +23,7 @@ use crate::{AppState, GatewayError};
|
||||
|
||||
const CODEX_QUOTA_CACHE_TTL_SECONDS: u64 = 30;
|
||||
const CODEX_QUOTA_CACHE_MAX_ENTRIES: usize = 4096;
|
||||
const RUNTIME_METADATA_CAS_MAX_ATTEMPTS: usize = 16;
|
||||
|
||||
type HeaderFingerprintCache = Mutex<HashMap<String, (String, Instant)>>;
|
||||
|
||||
@@ -29,6 +31,16 @@ static CODEX_QUOTA_HEADER_FINGERPRINT_CACHE: OnceLock<HeaderFingerprintCache> =
|
||||
static GROK_CHINESE_WAIT_DURATION_RE: OnceLock<Regex> = OnceLock::new();
|
||||
static GROK_ENGLISH_WAIT_DURATION_RE: OnceLock<Regex> = OnceLock::new();
|
||||
|
||||
fn upstream_metadata_namespace_value(
|
||||
upstream_metadata: Option<&Value>,
|
||||
namespace: &str,
|
||||
) -> Option<Value> {
|
||||
upstream_metadata
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|metadata| metadata.get(namespace))
|
||||
.cloned()
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub(crate) enum LocalReportEffect<'a> {
|
||||
Sync {
|
||||
@@ -172,6 +184,18 @@ fn merge_metadata_object(
|
||||
Some(Value::Object(merged))
|
||||
}
|
||||
|
||||
fn quota_status_snapshot_patch(status_snapshot: Option<&Value>) -> Value {
|
||||
let mut patch = serde_json::Map::new();
|
||||
if let Some(quota) = status_snapshot
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|snapshot| snapshot.get("quota"))
|
||||
.cloned()
|
||||
{
|
||||
patch.insert("quota".to_string(), quota);
|
||||
}
|
||||
Value::Object(patch)
|
||||
}
|
||||
|
||||
fn grok_report_context_model(report_context: Option<&Value>) -> Option<String> {
|
||||
report_context
|
||||
.and_then(|context| context.get("mapped_model"))
|
||||
@@ -321,6 +345,7 @@ async fn sync_gemini_cli_credits_from_report(
|
||||
Some(value) => value,
|
||||
None => return Ok(false),
|
||||
};
|
||||
let now_unix_secs = current_unix_secs();
|
||||
let Some(key) = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
@@ -345,22 +370,21 @@ async fn sync_gemini_cli_credits_from_report(
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let now_unix_secs = current_unix_secs();
|
||||
let mut gemini_cli_bucket = key
|
||||
.upstream_metadata
|
||||
let expected_namespace_value =
|
||||
upstream_metadata_namespace_value(key.upstream_metadata.as_ref(), "gemini_cli");
|
||||
let mut gemini_cli_bucket = expected_namespace_value
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|metadata| metadata.get("gemini_cli"))
|
||||
.and_then(Value::as_object)
|
||||
.cloned()
|
||||
.unwrap_or_else(serde_json::Map::new);
|
||||
gemini_cli_bucket.insert("credits".to_string(), credits);
|
||||
gemini_cli_bucket.insert("credits".to_string(), credits.clone());
|
||||
gemini_cli_bucket.insert("updated_at".to_string(), json!(now_unix_secs));
|
||||
|
||||
let namespace_value = Value::Object(gemini_cli_bucket);
|
||||
let updated_upstream_metadata = merge_metadata_object(
|
||||
key.upstream_metadata.as_ref(),
|
||||
"gemini_cli",
|
||||
Value::Object(gemini_cli_bucket),
|
||||
namespace_value.clone(),
|
||||
);
|
||||
let updated_status_snapshot = sync_provider_key_quota_status_snapshot(
|
||||
key.status_snapshot.as_ref(),
|
||||
@@ -368,15 +392,22 @@ async fn sync_gemini_cli_credits_from_report(
|
||||
updated_upstream_metadata.as_ref(),
|
||||
"report_effect",
|
||||
);
|
||||
let mut updated_key = key;
|
||||
updated_key.upstream_metadata = updated_upstream_metadata;
|
||||
updated_key.status_snapshot = updated_status_snapshot;
|
||||
updated_key.updated_at_unix_secs = Some(now_unix_secs);
|
||||
|
||||
Ok(state
|
||||
.update_provider_catalog_key(&updated_key)
|
||||
.await?
|
||||
.is_some())
|
||||
let persisted = state
|
||||
.update_provider_catalog_key_runtime_metadata(&ProviderCatalogKeyRuntimeMetadataUpdate {
|
||||
key_id: key_id.clone(),
|
||||
namespace: "gemini_cli".to_string(),
|
||||
expected_upstream_metadata_value: expected_namespace_value,
|
||||
upstream_metadata_value: namespace_value,
|
||||
status_snapshot_patch: quota_status_snapshot_patch(updated_status_snapshot.as_ref()),
|
||||
updated_at_unix_secs: Some(now_unix_secs),
|
||||
})
|
||||
.await?;
|
||||
if persisted {
|
||||
return Ok(true);
|
||||
}
|
||||
// Credits returned by the provider are an authoritative snapshot; do not
|
||||
// replay it over a newer local namespace after a CAS conflict.
|
||||
Ok(false)
|
||||
}
|
||||
|
||||
fn grok_quota_reset_after_seconds(
|
||||
@@ -479,70 +510,85 @@ async fn sync_grok_quota_from_report_context(
|
||||
return Ok(false);
|
||||
};
|
||||
|
||||
let Some(key) = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&key.provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
if !provider.provider_type.trim().eq_ignore_ascii_case("grok") {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let Some(grok_bucket) = key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|metadata| metadata.get("grok"))
|
||||
.and_then(Value::as_object)
|
||||
.cloned()
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
|
||||
let mut updated_grok_bucket = grok_bucket;
|
||||
let now_unix_secs = current_unix_secs();
|
||||
if !grok_apply_quota_feedback(
|
||||
&mut updated_grok_bucket,
|
||||
model.as_str(),
|
||||
status_code,
|
||||
grok_quota_reset_after_seconds(body_json, report_context),
|
||||
now_unix_secs,
|
||||
) {
|
||||
return Ok(false);
|
||||
for attempt in 0..RUNTIME_METADATA_CAS_MAX_ATTEMPTS {
|
||||
let Some(key) = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&key.provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
if !provider.provider_type.trim().eq_ignore_ascii_case("grok") {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let expected_namespace_value =
|
||||
upstream_metadata_namespace_value(key.upstream_metadata.as_ref(), "grok");
|
||||
let Some(grok_bucket) = expected_namespace_value
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.cloned()
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
|
||||
let mut updated_grok_bucket = grok_bucket;
|
||||
if !grok_apply_quota_feedback(
|
||||
&mut updated_grok_bucket,
|
||||
model.as_str(),
|
||||
status_code,
|
||||
grok_quota_reset_after_seconds(body_json, report_context),
|
||||
now_unix_secs,
|
||||
) {
|
||||
return Ok(false);
|
||||
}
|
||||
grok_mark_quota_bucket_updated(&mut updated_grok_bucket, now_unix_secs);
|
||||
|
||||
let namespace_value = Value::Object(updated_grok_bucket);
|
||||
let updated_upstream_metadata = merge_metadata_object(
|
||||
key.upstream_metadata.as_ref(),
|
||||
"grok",
|
||||
namespace_value.clone(),
|
||||
);
|
||||
let updated_status_snapshot = sync_provider_key_quota_status_snapshot(
|
||||
key.status_snapshot.as_ref(),
|
||||
provider.provider_type.as_str(),
|
||||
updated_upstream_metadata.as_ref(),
|
||||
"report_effect",
|
||||
);
|
||||
let persisted = state
|
||||
.update_provider_catalog_key_runtime_metadata(
|
||||
&ProviderCatalogKeyRuntimeMetadataUpdate {
|
||||
key_id: key_id.clone(),
|
||||
namespace: "grok".to_string(),
|
||||
expected_upstream_metadata_value: expected_namespace_value,
|
||||
upstream_metadata_value: namespace_value,
|
||||
status_snapshot_patch: quota_status_snapshot_patch(
|
||||
updated_status_snapshot.as_ref(),
|
||||
),
|
||||
updated_at_unix_secs: Some(now_unix_secs),
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
if persisted {
|
||||
return Ok(true);
|
||||
}
|
||||
if attempt + 1 < RUNTIME_METADATA_CAS_MAX_ATTEMPTS {
|
||||
let backoff_us = 50_u64.saturating_mul((attempt + 1) as u64).min(1_000);
|
||||
tokio::time::sleep(Duration::from_micros(backoff_us)).await;
|
||||
}
|
||||
}
|
||||
grok_mark_quota_bucket_updated(&mut updated_grok_bucket, now_unix_secs);
|
||||
|
||||
let updated_upstream_metadata = merge_metadata_object(
|
||||
key.upstream_metadata.as_ref(),
|
||||
"grok",
|
||||
Value::Object(updated_grok_bucket),
|
||||
);
|
||||
let updated_status_snapshot = sync_provider_key_quota_status_snapshot(
|
||||
key.status_snapshot.as_ref(),
|
||||
provider.provider_type.as_str(),
|
||||
updated_upstream_metadata.as_ref(),
|
||||
"report_effect",
|
||||
);
|
||||
let mut updated_key = key;
|
||||
updated_key.upstream_metadata = updated_upstream_metadata;
|
||||
updated_key.status_snapshot = updated_status_snapshot;
|
||||
updated_key.updated_at_unix_secs = Some(now_unix_secs);
|
||||
|
||||
Ok(state
|
||||
.update_provider_catalog_key_runtime_state(&updated_key)
|
||||
.await?
|
||||
.is_some())
|
||||
Ok(false)
|
||||
}
|
||||
|
||||
async fn apply_local_sync_report_effect(state: &AppState, payload: &GatewaySyncReportRequest) {
|
||||
@@ -820,7 +866,7 @@ async fn sync_codex_quota_from_response_headers(
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint, now);
|
||||
set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint.clone(), now);
|
||||
return Ok(false);
|
||||
};
|
||||
|
||||
@@ -830,53 +876,56 @@ async fn sync_codex_quota_from_response_headers(
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint, now);
|
||||
set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint.clone(), now);
|
||||
return Ok(false);
|
||||
};
|
||||
if !provider.provider_type.trim().eq_ignore_ascii_case("codex") {
|
||||
set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint, now);
|
||||
set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint.clone(), now);
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let current_codex = key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|metadata| metadata.get("codex"))
|
||||
.and_then(Value::as_object)
|
||||
.cloned()
|
||||
let expected_namespace_value =
|
||||
upstream_metadata_namespace_value(key.upstream_metadata.as_ref(), "codex");
|
||||
let current_codex = expected_namespace_value
|
||||
.clone()
|
||||
.and_then(|value| value.as_object().cloned())
|
||||
.unwrap_or_else(serde_json::Map::new);
|
||||
let current_codex = Value::Object(current_codex);
|
||||
let Some(current_fingerprint) = fingerprint_codex_payload(¤t_codex) else {
|
||||
set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint, now);
|
||||
set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint.clone(), now);
|
||||
return Ok(false);
|
||||
};
|
||||
if current_fingerprint == incoming_fingerprint {
|
||||
set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint, now);
|
||||
set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint.clone(), now);
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let updated_upstream_metadata =
|
||||
merge_metadata_object(key.upstream_metadata.as_ref(), "codex", parsed);
|
||||
merge_metadata_object(key.upstream_metadata.as_ref(), "codex", parsed.clone());
|
||||
let updated_status_snapshot = sync_provider_key_quota_status_snapshot(
|
||||
key.status_snapshot.as_ref(),
|
||||
provider.provider_type.as_str(),
|
||||
updated_upstream_metadata.as_ref(),
|
||||
"response_headers",
|
||||
);
|
||||
let mut updated_key = key;
|
||||
updated_key.upstream_metadata = updated_upstream_metadata;
|
||||
updated_key.status_snapshot = updated_status_snapshot;
|
||||
updated_key.updated_at_unix_secs = Some(now_unix_secs);
|
||||
|
||||
let updated = state
|
||||
.update_provider_catalog_key_runtime_state(&updated_key)
|
||||
.await?
|
||||
.is_some();
|
||||
.update_provider_catalog_key_runtime_metadata(&ProviderCatalogKeyRuntimeMetadataUpdate {
|
||||
key_id: key_id.clone(),
|
||||
namespace: "codex".to_string(),
|
||||
expected_upstream_metadata_value: expected_namespace_value,
|
||||
upstream_metadata_value: parsed.clone(),
|
||||
status_snapshot_patch: quota_status_snapshot_patch(updated_status_snapshot.as_ref()),
|
||||
updated_at_unix_secs: Some(now_unix_secs),
|
||||
})
|
||||
.await?;
|
||||
if updated {
|
||||
set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint, now);
|
||||
set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint.clone(), now);
|
||||
return Ok(true);
|
||||
}
|
||||
Ok(updated)
|
||||
// Response headers describe an authoritative quota snapshot. A CAS
|
||||
// conflict means a newer snapshot/delta won, so avoid replaying stale
|
||||
// data over it.
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -892,7 +941,86 @@ pub(crate) fn clear_local_report_effect_caches_for_tests() {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::data::GatewayDataState;
|
||||
|
||||
#[tokio::test]
|
||||
async fn gemini_report_metadata_write_preserves_adaptive_and_other_provider_state() {
|
||||
let provider = StoredProviderCatalogProvider::new(
|
||||
"gemini-provider".to_string(),
|
||||
"Gemini CLI".to_string(),
|
||||
None,
|
||||
"gemini_cli".to_string(),
|
||||
)
|
||||
.expect("provider should build");
|
||||
let mut key = StoredProviderCatalogKey::new(
|
||||
"gemini-key".to_string(),
|
||||
"gemini-provider".to_string(),
|
||||
"Gemini Key".to_string(),
|
||||
"oauth".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build");
|
||||
key.learned_rpm_limit = Some(12);
|
||||
key.rpm_429_count = Some(3);
|
||||
key.upstream_metadata = Some(json!({
|
||||
"gemini_cli": {"credits":{"remaining":9}},
|
||||
"codex": {"remaining":7}
|
||||
}));
|
||||
key.status_snapshot = Some(json!({
|
||||
"quota": {"source":"old"},
|
||||
"observation_count": 4,
|
||||
"learning_confidence": 0.7,
|
||||
"oauth": {"invalid":false}
|
||||
}));
|
||||
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
vec![],
|
||||
vec![key],
|
||||
));
|
||||
let state = AppState::new()
|
||||
.expect("gateway state should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_repository_for_tests(repository),
|
||||
);
|
||||
|
||||
assert!(sync_gemini_cli_credits_from_report(
|
||||
&state,
|
||||
Some(&json!({"key_id":"gemini-key"})),
|
||||
Some(json!({"remaining":3,"total":10}))
|
||||
)
|
||||
.await
|
||||
.expect("report metadata should update"));
|
||||
|
||||
let stored = state
|
||||
.read_provider_catalog_keys_by_ids(&["gemini-key".to_string()])
|
||||
.await
|
||||
.expect("key should reload")
|
||||
.pop()
|
||||
.expect("key should exist");
|
||||
assert_eq!(stored.learned_rpm_limit, Some(12));
|
||||
assert_eq!(stored.rpm_429_count, Some(3));
|
||||
assert_eq!(
|
||||
stored.upstream_metadata.as_ref().unwrap()["codex"],
|
||||
json!({"remaining":7})
|
||||
);
|
||||
assert_eq!(
|
||||
stored.upstream_metadata.as_ref().unwrap()["gemini_cli"]["credits"]["remaining"],
|
||||
json!(3)
|
||||
);
|
||||
let status = stored.status_snapshot.expect("status should exist");
|
||||
assert_eq!(status["observation_count"], json!(4));
|
||||
assert_eq!(status["learning_confidence"], json!(0.7));
|
||||
assert_eq!(status["oauth"], json!({"invalid":false}));
|
||||
assert_eq!(status["quota"]["provider_type"], json!("gemini_cli"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn grok_quota_feedback_decrements_the_matching_window() {
|
||||
|
||||
Reference in New Issue
Block a user