mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-07 18:07:47 +08:00
fix(codex): fence concurrent quota updates
This commit is contained in:
@@ -237,6 +237,13 @@ pub(crate) struct LocalOAuthInvalidationEffect<'a> {
|
||||
pub(crate) response_text: Option<&'a str>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub(crate) struct LocalOAuthSuccessEffect<'a> {
|
||||
pub(crate) status_code: u16,
|
||||
pub(crate) request_started_at_unix_ms: Option<u64>,
|
||||
pub(crate) request_order_id: Option<&'a str>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub(crate) enum LocalExecutionEffect<'a> {
|
||||
AttemptFailure(LocalAttemptFailureEffect),
|
||||
@@ -245,6 +252,7 @@ pub(crate) enum LocalExecutionEffect<'a> {
|
||||
HealthSuccess(LocalHealthSuccessEffect),
|
||||
AdaptiveSuccess(LocalAdaptiveSuccessEffect),
|
||||
OauthInvalidation(LocalOAuthInvalidationEffect<'a>),
|
||||
OauthSuccess(LocalOAuthSuccessEffect<'a>),
|
||||
PoolSuccessSync {
|
||||
payload: &'a GatewaySyncReportRequest,
|
||||
},
|
||||
@@ -255,6 +263,68 @@ pub(crate) enum LocalExecutionEffect<'a> {
|
||||
PoolStreamTimeout,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct OwnedLocalOAuthSuccessEffect {
|
||||
status_code: u16,
|
||||
provider_id: String,
|
||||
endpoint_id: String,
|
||||
key_id: String,
|
||||
authorization: String,
|
||||
request_started_at_unix_ms: u64,
|
||||
request_order_id: String,
|
||||
observed_credential_generation: Option<String>,
|
||||
}
|
||||
|
||||
fn owned_local_oauth_success_effect(
|
||||
plan: &ExecutionPlan,
|
||||
report_context: Option<&Value>,
|
||||
effect: LocalOAuthSuccessEffect<'_>,
|
||||
) -> Option<OwnedLocalOAuthSuccessEffect> {
|
||||
if !(200..300).contains(&effect.status_code) {
|
||||
return None;
|
||||
}
|
||||
let request_started_at_unix_ms = effect.request_started_at_unix_ms?;
|
||||
let request_order_id = effect
|
||||
.request_order_id
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())?;
|
||||
let authorization = execution_plan_authorization(plan)?.trim().to_string();
|
||||
if authorization.is_empty() || bearer_access_token(&authorization).is_none() {
|
||||
return None;
|
||||
}
|
||||
Some(OwnedLocalOAuthSuccessEffect {
|
||||
status_code: effect.status_code,
|
||||
provider_id: plan.provider_id.clone(),
|
||||
endpoint_id: plan.endpoint_id.clone(),
|
||||
key_id: plan.key_id.clone(),
|
||||
authorization,
|
||||
request_started_at_unix_ms,
|
||||
request_order_id: request_order_id.to_string(),
|
||||
observed_credential_generation: report_context_string_field(
|
||||
report_context,
|
||||
"codex_credential_generation",
|
||||
)
|
||||
.map(ToOwned::to_owned),
|
||||
})
|
||||
}
|
||||
|
||||
/// Schedule a fenced Codex OAuth-success observation after provider headers are available.
|
||||
/// Only the small identity/observation tuple is moved into the task; request bodies and plans
|
||||
/// remain owned by the caller.
|
||||
pub(crate) fn spawn_local_oauth_success_effect(
|
||||
state: AppState,
|
||||
plan: &ExecutionPlan,
|
||||
report_context: Option<&Value>,
|
||||
effect: LocalOAuthSuccessEffect<'_>,
|
||||
) {
|
||||
let Some(effect) = owned_local_oauth_success_effect(plan, report_context, effect) else {
|
||||
return;
|
||||
};
|
||||
tokio::spawn(async move {
|
||||
record_oauth_success_effect_owned(&state, effect).await;
|
||||
});
|
||||
}
|
||||
|
||||
struct PoolFeedbackContext {
|
||||
pool_config: AdminProviderPoolConfig,
|
||||
sticky_session_token: Option<String>,
|
||||
@@ -302,6 +372,9 @@ pub(crate) async fn apply_local_execution_effect(
|
||||
LocalExecutionEffect::OauthInvalidation(effect) => {
|
||||
record_oauth_invalidation_effect(state, context, effect).await;
|
||||
}
|
||||
LocalExecutionEffect::OauthSuccess(effect) => {
|
||||
record_oauth_success_effect(state, context, effect).await;
|
||||
}
|
||||
LocalExecutionEffect::PoolSuccessSync { payload } => {
|
||||
record_sync_pool_success_effect(state, context, payload).await;
|
||||
release_pool_key_lease_effect(state, context).await;
|
||||
@@ -349,6 +422,16 @@ fn report_context_string_field<'a>(
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
fn report_context_u64_field(report_context: Option<&Value>, field: &str) -> Option<u64> {
|
||||
report_context
|
||||
.and_then(|context| context.get(field))
|
||||
.and_then(|value| {
|
||||
value
|
||||
.as_u64()
|
||||
.or_else(|| value.as_str().and_then(|value| value.parse::<u64>().ok()))
|
||||
})
|
||||
}
|
||||
|
||||
fn local_scheduler_affinity_cache_key(report_context: Option<&Value>) -> Option<String> {
|
||||
let client_session_affinity = local_client_session_affinity(report_context);
|
||||
let policy_context = scheduler_affinity_policy_context_from_report_context(report_context);
|
||||
@@ -1463,9 +1546,23 @@ async fn record_oauth_invalidation_effect(
|
||||
) else {
|
||||
return;
|
||||
};
|
||||
let request_started_at_unix_ms = report_context_u64_field(
|
||||
context.report_context,
|
||||
"provider_request_started_at_unix_ms",
|
||||
);
|
||||
let request_order_id =
|
||||
report_context_string_field(context.report_context, "provider_request_order_id");
|
||||
let observed_credential_generation =
|
||||
report_context_string_field(context.report_context, "codex_credential_generation");
|
||||
|
||||
if let Err(err) = state
|
||||
.mark_provider_transport_oauth_invalid_fenced(&transport, invalid_reason.as_str())
|
||||
.mark_provider_transport_oauth_invalid_fenced(
|
||||
&transport,
|
||||
invalid_reason.as_str(),
|
||||
request_started_at_unix_ms,
|
||||
request_order_id,
|
||||
observed_credential_generation,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
@@ -1475,6 +1572,70 @@ async fn record_oauth_invalidation_effect(
|
||||
}
|
||||
}
|
||||
|
||||
async fn record_oauth_success_effect(
|
||||
state: &AppState,
|
||||
context: LocalExecutionEffectContext<'_>,
|
||||
effect: LocalOAuthSuccessEffect<'_>,
|
||||
) {
|
||||
let Some(effect) =
|
||||
owned_local_oauth_success_effect(context.plan, context.report_context, effect)
|
||||
else {
|
||||
return;
|
||||
};
|
||||
record_oauth_success_effect_owned(state, effect).await;
|
||||
}
|
||||
|
||||
async fn record_oauth_success_effect_owned(state: &AppState, effect: OwnedLocalOAuthSuccessEffect) {
|
||||
if !(200..300).contains(&effect.status_code) {
|
||||
return;
|
||||
}
|
||||
let transport = match state
|
||||
.read_provider_transport_snapshot(&effect.provider_id, &effect.endpoint_id, &effect.key_id)
|
||||
.await
|
||||
{
|
||||
Ok(Some(transport)) => transport,
|
||||
Ok(None) => return,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
"gateway orchestration effects: failed to read transport snapshot for oauth success provider {} endpoint {} key {}: {:?}",
|
||||
effect.provider_id, effect.endpoint_id, effect.key_id, err
|
||||
);
|
||||
return;
|
||||
}
|
||||
};
|
||||
if !transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("codex")
|
||||
|| crate::provider_transport::is_codex_agent_identity_transport(&transport)
|
||||
|| !transport.key.auth_type.trim().eq_ignore_ascii_case("oauth")
|
||||
|| crate::provider_transport::resolve_local_generic_oauth_transport_authorization(
|
||||
&transport,
|
||||
)
|
||||
.as_deref()
|
||||
.and_then(bearer_access_token)
|
||||
!= bearer_access_token(effect.authorization.as_str())
|
||||
{
|
||||
return;
|
||||
}
|
||||
|
||||
if let Err(err) = state
|
||||
.mark_provider_transport_oauth_success_fenced(
|
||||
&transport,
|
||||
Some(effect.request_started_at_unix_ms),
|
||||
Some(effect.request_order_id.as_str()),
|
||||
effect.observed_credential_generation.as_deref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
"gateway orchestration effects: failed to persist oauth success for provider {} endpoint {} key {}: {:?}",
|
||||
effect.provider_id, effect.endpoint_id, effect.key_id, err
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
fn execution_plan_uses_codex_agent_identity(plan: &ExecutionPlan) -> bool {
|
||||
execution_plan_authorization(plan)
|
||||
.is_some_and(crate::provider_transport::is_codex_agent_identity_authorization)
|
||||
@@ -1802,6 +1963,7 @@ mod tests {
|
||||
StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_test_support::ManagedRedisServer;
|
||||
use aether_usage_runtime::GatewaySyncReportRequest;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
use super::{
|
||||
@@ -1811,11 +1973,13 @@ mod tests {
|
||||
pool_score_hard_state_for_status, resolve_pool_feedback_context,
|
||||
LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect,
|
||||
LocalExecutionEffect, LocalExecutionEffectContext, LocalHealthFailureEffect,
|
||||
LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalPoolErrorEffect,
|
||||
ProviderKeyEffectLockPool,
|
||||
LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalOAuthSuccessEffect,
|
||||
LocalPoolErrorEffect, ProviderKeyEffectLockPool,
|
||||
};
|
||||
use crate::data::{GatewayDataConfig, GatewayDataState};
|
||||
use crate::orchestration::LocalFailoverClassification;
|
||||
use crate::orchestration::{
|
||||
apply_local_report_effect, LocalFailoverClassification, LocalReportEffect,
|
||||
};
|
||||
use crate::scheduler::affinity::SCHEDULER_AFFINITY_TTL;
|
||||
use crate::AppState;
|
||||
use aether_scheduler_core::{
|
||||
@@ -3496,6 +3660,277 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oauth_success_clears_recoverable_codex_invalid_state() {
|
||||
let mut key = sample_codex_key();
|
||||
key.oauth_invalid_at_unix_secs = Some(100);
|
||||
key.oauth_invalid_reason = Some("[OAUTH_EXPIRED] session expired".to_string());
|
||||
key.upstream_metadata = Some(json!({
|
||||
"codex": {
|
||||
"credential_generation": "credential-generation-current",
|
||||
"oauth_state_request_started_at_unix_ms": 100_000u64,
|
||||
"oauth_state_request_id": "00000001-86a0-7000-8000-000000000001"
|
||||
}
|
||||
}));
|
||||
let state = codex_state_with_provider_and_key(sample_codex_provider(), key);
|
||||
let plan = sample_codex_plan();
|
||||
let report_context = json!({
|
||||
"codex_credential_generation": "credential-generation-current",
|
||||
"provider_request_started_at_unix_ms": 200_000u64,
|
||||
"provider_request_order_id": "00000003-0d40-7000-8000-000000000001"
|
||||
});
|
||||
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: Some(&report_context),
|
||||
},
|
||||
LocalExecutionEffect::OauthSuccess(LocalOAuthSuccessEffect {
|
||||
status_code: 200,
|
||||
request_started_at_unix_ms: Some(200_000),
|
||||
request_order_id: Some("00000003-0d40-7000-8000-000000000001"),
|
||||
}),
|
||||
)
|
||||
.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_eq!(stored_key.oauth_invalid_at_unix_secs, None);
|
||||
assert_eq!(stored_key.oauth_invalid_reason, None);
|
||||
assert_eq!(
|
||||
stored_key.upstream_metadata.as_ref().and_then(
|
||||
|metadata| metadata.pointer("/codex/oauth_state_request_started_at_unix_ms")
|
||||
),
|
||||
Some(&json!(200_000u64))
|
||||
);
|
||||
assert_eq!(
|
||||
stored_key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/codex/oauth_state_request_id")),
|
||||
Some(&json!("00000003-0d40-7000-8000-000000000001"))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn same_millisecond_older_oauth_success_does_not_clear_newer_codex_invalidation() {
|
||||
let mut key = sample_codex_key();
|
||||
key.oauth_invalid_at_unix_secs = Some(300);
|
||||
key.oauth_invalid_reason = Some("[OAUTH_EXPIRED] session expired".to_string());
|
||||
key.upstream_metadata = Some(json!({
|
||||
"codex": {
|
||||
"credential_generation": "credential-generation-current",
|
||||
"oauth_state_request_started_at_unix_ms": 300_000u64,
|
||||
"oauth_state_request_id": "00000004-93e0-7000-8000-000000000002"
|
||||
}
|
||||
}));
|
||||
let state = codex_state_with_provider_and_key(sample_codex_provider(), key);
|
||||
let plan = sample_codex_plan();
|
||||
let report_context = json!({
|
||||
"codex_credential_generation": "credential-generation-current",
|
||||
"provider_request_started_at_unix_ms": 300_000u64,
|
||||
"provider_request_order_id": "00000004-93e0-7000-8000-000000000001"
|
||||
});
|
||||
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: Some(&report_context),
|
||||
},
|
||||
LocalExecutionEffect::OauthSuccess(LocalOAuthSuccessEffect {
|
||||
status_code: 200,
|
||||
request_started_at_unix_ms: Some(300_000),
|
||||
request_order_id: Some("00000004-93e0-7000-8000-000000000001"),
|
||||
}),
|
||||
)
|
||||
.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_eq!(stored_key.oauth_invalid_at_unix_secs, Some(300));
|
||||
assert_eq!(
|
||||
stored_key.oauth_invalid_reason.as_deref(),
|
||||
Some("[OAUTH_EXPIRED] session expired")
|
||||
);
|
||||
assert_eq!(
|
||||
stored_key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/codex/oauth_state_request_id")),
|
||||
Some(&json!("00000004-93e0-7000-8000-000000000002"))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oauth_success_preserves_codex_account_block() {
|
||||
let mut key = sample_codex_key();
|
||||
key.oauth_invalid_at_unix_secs = Some(100);
|
||||
key.oauth_invalid_reason = Some("[ACCOUNT_BLOCK] account deactivated".to_string());
|
||||
key.upstream_metadata = Some(json!({
|
||||
"codex": {
|
||||
"credential_generation": "credential-generation-current"
|
||||
}
|
||||
}));
|
||||
let state = codex_state_with_provider_and_key(sample_codex_provider(), key);
|
||||
let plan = sample_codex_plan();
|
||||
let report_context = json!({
|
||||
"codex_credential_generation": "credential-generation-current",
|
||||
"provider_request_started_at_unix_ms": 200_000u64,
|
||||
"provider_request_order_id": "00000003-0d40-7000-8000-000000000001"
|
||||
});
|
||||
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: Some(&report_context),
|
||||
},
|
||||
LocalExecutionEffect::OauthSuccess(LocalOAuthSuccessEffect {
|
||||
status_code: 200,
|
||||
request_started_at_unix_ms: Some(200_000),
|
||||
request_order_id: Some("00000003-0d40-7000-8000-000000000001"),
|
||||
}),
|
||||
)
|
||||
.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_eq!(stored_key.oauth_invalid_at_unix_secs, Some(100));
|
||||
assert_eq!(
|
||||
stored_key.oauth_invalid_reason.as_deref(),
|
||||
Some("[ACCOUNT_BLOCK] account deactivated")
|
||||
);
|
||||
assert_eq!(
|
||||
stored_key.upstream_metadata.as_ref().and_then(
|
||||
|metadata| metadata.pointer("/codex/oauth_state_request_started_at_unix_ms")
|
||||
),
|
||||
Some(&json!(200_000u64))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oauth_success_does_not_clear_codex_invalid_state_for_replaced_bearer() {
|
||||
let mut key = sample_codex_key();
|
||||
key.oauth_invalid_at_unix_secs = Some(100);
|
||||
key.oauth_invalid_reason = Some("[OAUTH_EXPIRED] session expired".to_string());
|
||||
key.upstream_metadata = Some(json!({
|
||||
"codex": {
|
||||
"credential_generation": "credential-generation-current"
|
||||
}
|
||||
}));
|
||||
let state = codex_state_with_provider_and_key(sample_codex_provider(), key);
|
||||
let mut plan = sample_codex_plan();
|
||||
plan.headers.insert(
|
||||
"authorization".to_string(),
|
||||
"Bearer replaced-access-token".to_string(),
|
||||
);
|
||||
let report_context = json!({
|
||||
"codex_credential_generation": "credential-generation-current",
|
||||
"provider_request_started_at_unix_ms": 200_000u64,
|
||||
"provider_request_order_id": "00000003-0d40-7000-8000-000000000001"
|
||||
});
|
||||
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: Some(&report_context),
|
||||
},
|
||||
LocalExecutionEffect::OauthSuccess(LocalOAuthSuccessEffect {
|
||||
status_code: 200,
|
||||
request_started_at_unix_ms: Some(200_000),
|
||||
request_order_id: Some("00000003-0d40-7000-8000-000000000001"),
|
||||
}),
|
||||
)
|
||||
.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_eq!(stored_key.oauth_invalid_at_unix_secs, Some(100));
|
||||
assert_eq!(
|
||||
stored_key.oauth_invalid_reason.as_deref(),
|
||||
Some("[OAUTH_EXPIRED] session expired")
|
||||
);
|
||||
assert!(stored_key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/codex/oauth_state_request_started_at_unix_ms"))
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oauth_success_effect_ignores_non_success_status() {
|
||||
let mut key = sample_codex_key();
|
||||
key.oauth_invalid_at_unix_secs = Some(100);
|
||||
key.oauth_invalid_reason = Some("[OAUTH_EXPIRED] session expired".to_string());
|
||||
key.upstream_metadata = Some(json!({
|
||||
"codex": {
|
||||
"credential_generation": "credential-generation-current"
|
||||
}
|
||||
}));
|
||||
let state = codex_state_with_provider_and_key(sample_codex_provider(), key);
|
||||
let plan = sample_codex_plan();
|
||||
let report_context = json!({
|
||||
"codex_credential_generation": "credential-generation-current",
|
||||
"provider_request_started_at_unix_ms": 200_000u64,
|
||||
"provider_request_order_id": "00000003-0d40-7000-8000-000000000001"
|
||||
});
|
||||
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: Some(&report_context),
|
||||
},
|
||||
LocalExecutionEffect::OauthSuccess(LocalOAuthSuccessEffect {
|
||||
status_code: 304,
|
||||
request_started_at_unix_ms: Some(200_000),
|
||||
request_order_id: Some("00000003-0d40-7000-8000-000000000001"),
|
||||
}),
|
||||
)
|
||||
.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_eq!(stored_key.oauth_invalid_at_unix_secs, Some(100));
|
||||
assert_eq!(
|
||||
stored_key.oauth_invalid_reason.as_deref(),
|
||||
Some("[OAUTH_EXPIRED] session expired")
|
||||
);
|
||||
assert!(stored_key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/codex/oauth_state_request_started_at_unix_ms"))
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oauth_invalidation_marks_generic_codex_403_as_token_invalid() {
|
||||
let state = codex_state();
|
||||
@@ -3581,7 +4016,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oauth_invalidation_auto_removes_inactive_pat_owner() {
|
||||
async fn oauth_invalidation_auto_remove_keeps_inactive_pat_owner() {
|
||||
let state = codex_state_with_auto_remove();
|
||||
let plan = sample_codex_plan();
|
||||
|
||||
@@ -3600,6 +4035,39 @@ mod tests {
|
||||
)
|
||||
.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("recoverable token invalidation should retain the key");
|
||||
assert_eq!(
|
||||
stored_key.oauth_invalid_reason.as_deref(),
|
||||
Some("[OAUTH_EXPIRED] Personal access token owner is inactive.")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oauth_invalidation_auto_removes_account_block() {
|
||||
let state = codex_state_with_auto_remove();
|
||||
let plan = sample_codex_plan();
|
||||
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: None,
|
||||
},
|
||||
LocalExecutionEffect::OauthInvalidation(LocalOAuthInvalidationEffect {
|
||||
status_code: 403,
|
||||
response_text: Some(
|
||||
r#"{"error":{"message":"account has been deactivated"},"status":403}"#,
|
||||
),
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
|
||||
let stored_keys = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
|
||||
.await
|
||||
@@ -3709,6 +4177,155 @@ mod tests {
|
||||
assert_eq!(stored_key.oauth_invalid_reason, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oauth_invalidation_rejects_stale_codex_credential_generation() {
|
||||
let mut key = sample_codex_key();
|
||||
key.upstream_metadata = Some(json!({
|
||||
"codex": {
|
||||
"credential_generation": "credential-generation-current"
|
||||
}
|
||||
}));
|
||||
let state = codex_state_with_provider_and_key(sample_codex_provider(), key);
|
||||
let plan = sample_codex_plan();
|
||||
let report_context = json!({
|
||||
"codex_credential_generation": "credential-generation-stale",
|
||||
"provider_request_started_at_unix_ms": 100_000u64,
|
||||
"provider_request_order_id": "00000001-86a0-7000-8000-000000000001"
|
||||
});
|
||||
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: Some(&report_context),
|
||||
},
|
||||
LocalExecutionEffect::OauthInvalidation(LocalOAuthInvalidationEffect {
|
||||
status_code: 401,
|
||||
response_text: Some(r#"{"error":{"message":"session expired"}}"#),
|
||||
}),
|
||||
)
|
||||
.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_eq!(stored_key.oauth_invalid_at_unix_secs, None);
|
||||
assert_eq!(stored_key.oauth_invalid_reason, None);
|
||||
assert!(stored_key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/codex/oauth_state_request_started_at_unix_ms"))
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn delayed_codex_quota_report_does_not_clear_newer_oauth_invalidation() {
|
||||
let mut key = sample_codex_key();
|
||||
key.upstream_metadata = Some(json!({
|
||||
"codex": {
|
||||
"credential_generation": "credential-generation-current",
|
||||
"account_quota_reset_generation": 0
|
||||
}
|
||||
}));
|
||||
let state = codex_state_with_provider_and_key(sample_codex_provider(), key);
|
||||
let plan = sample_codex_plan();
|
||||
let older_request_id = "00000001-86a0-7000-8000-000000000001";
|
||||
let newer_request_id = "00000001-86a0-7000-8000-000000000002";
|
||||
let older_uuid = uuid::Uuid::parse_str(older_request_id).expect("older id should parse");
|
||||
let newer_uuid = uuid::Uuid::parse_str(newer_request_id).expect("newer id should parse");
|
||||
assert_eq!(older_uuid.get_version_num(), 7);
|
||||
assert_eq!(newer_uuid.get_version_num(), 7);
|
||||
assert!(older_uuid < newer_uuid);
|
||||
|
||||
let invalidation_context = json!({
|
||||
"codex_credential_generation": "credential-generation-current",
|
||||
"provider_request_started_at_unix_ms": 100_000u64,
|
||||
"provider_request_order_id": newer_request_id
|
||||
});
|
||||
apply_local_execution_effect(
|
||||
&state,
|
||||
LocalExecutionEffectContext {
|
||||
plan: &plan,
|
||||
report_context: Some(&invalidation_context),
|
||||
},
|
||||
LocalExecutionEffect::OauthInvalidation(LocalOAuthInvalidationEffect {
|
||||
status_code: 401,
|
||||
response_text: Some(r#"{"error":{"message":"session expired"}}"#),
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
|
||||
let delayed_report_context = json!({
|
||||
"key_id": plan.key_id,
|
||||
"codex_credential_generation": "credential-generation-current",
|
||||
"codex_quota_reset_generation": 0,
|
||||
"provider_request_started_at_unix_ms": 100_000u64,
|
||||
"provider_response_headers_observed_at_unix_ms": 110_000u64,
|
||||
"provider_request_order_id": older_request_id
|
||||
});
|
||||
let delayed_headers = BTreeMap::from([
|
||||
("x-codex-plan-type".to_string(), "free".to_string()),
|
||||
("x-codex-primary-used-percent".to_string(), "25".to_string()),
|
||||
(
|
||||
"x-codex-primary-reset-at".to_string(),
|
||||
"2000000000".to_string(),
|
||||
),
|
||||
(
|
||||
"x-codex-primary-window-minutes".to_string(),
|
||||
"300".to_string(),
|
||||
),
|
||||
]);
|
||||
let delayed_report = GatewaySyncReportRequest {
|
||||
trace_id: "trace-delayed-codex-quota".to_string(),
|
||||
report_kind: "openai_responses_sync_success".to_string(),
|
||||
report_context: Some(delayed_report_context),
|
||||
status_code: 200,
|
||||
headers: delayed_headers,
|
||||
body_json: None,
|
||||
client_body_json: None,
|
||||
body_base64: None,
|
||||
telemetry: None,
|
||||
};
|
||||
apply_local_report_effect(
|
||||
&state,
|
||||
LocalReportEffect::Sync {
|
||||
payload: &delayed_report,
|
||||
},
|
||||
)
|
||||
.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_eq!(
|
||||
stored_key.oauth_invalid_reason.as_deref(),
|
||||
Some("[OAUTH_EXPIRED] session expired")
|
||||
);
|
||||
assert!(stored_key.oauth_invalid_at_unix_secs.is_some());
|
||||
assert_eq!(
|
||||
stored_key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/codex/oauth_state_request_id")),
|
||||
Some(&json!(newer_request_id))
|
||||
);
|
||||
assert_eq!(
|
||||
stored_key
|
||||
.status_snapshot
|
||||
.as_ref()
|
||||
.and_then(|snapshot| snapshot.pointer("/oauth/code")),
|
||||
Some(&json!("expired"))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn health_failure_updates_codex_key_for_current_bearer_request() {
|
||||
let state = codex_state();
|
||||
|
||||
@@ -33,10 +33,10 @@ pub(crate) use self::classifier::{
|
||||
LocalTransportFailoverClassification,
|
||||
};
|
||||
pub(crate) use self::effects::{
|
||||
apply_local_execution_effect, LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect,
|
||||
LocalAttemptFailureEffect, LocalExecutionEffect, LocalExecutionEffectContext,
|
||||
LocalHealthFailureEffect, LocalHealthSuccessEffect, LocalOAuthInvalidationEffect,
|
||||
LocalPoolErrorEffect,
|
||||
apply_local_execution_effect, spawn_local_oauth_success_effect, LocalAdaptiveRateLimitEffect,
|
||||
LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, LocalExecutionEffect,
|
||||
LocalExecutionEffectContext, LocalHealthFailureEffect, LocalHealthSuccessEffect,
|
||||
LocalOAuthInvalidationEffect, LocalOAuthSuccessEffect, LocalPoolErrorEffect,
|
||||
};
|
||||
pub(crate) use self::health::{
|
||||
project_local_failure_health, project_local_key_circuit_closed,
|
||||
|
||||
@@ -227,6 +227,30 @@ pub(crate) fn append_local_failover_policy_to_value(
|
||||
"local_failover_policy".to_string(),
|
||||
local_failover_policy_to_value(&local_failover_policy_from_transport(transport)),
|
||||
);
|
||||
if transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("codex")
|
||||
{
|
||||
let codex = transport
|
||||
.key
|
||||
.upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|metadata| metadata.get("codex"));
|
||||
object.insert(
|
||||
"codex_quota_reset_generation".to_string(),
|
||||
Value::from(aether_admin::provider::quota::codex_quota_account_reset_generation(codex)),
|
||||
);
|
||||
if let Some(generation) = aether_admin::provider::quota::codex_credential_generation(codex)
|
||||
{
|
||||
object.insert(
|
||||
"codex_credential_generation".to_string(),
|
||||
Value::String(generation.to_string()),
|
||||
);
|
||||
}
|
||||
}
|
||||
Value::Object(object)
|
||||
}
|
||||
|
||||
@@ -459,6 +483,26 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn codex_report_context_captures_quota_and_credential_generations() {
|
||||
let mut transport = sample_transport(None, None, None);
|
||||
transport.provider.provider_type = "codex".to_string();
|
||||
transport.key.upstream_metadata = Some(json!({
|
||||
"codex": {
|
||||
"account_quota_reset_generation": 7,
|
||||
"credential_generation": "credential-generation-7"
|
||||
}
|
||||
}));
|
||||
|
||||
let report_context = append_local_failover_policy_to_value(json!({}), &transport);
|
||||
|
||||
assert_eq!(report_context["codex_quota_reset_generation"], json!(7u64));
|
||||
assert_eq!(
|
||||
report_context["codex_credential_generation"],
|
||||
json!("credential-generation-7")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transport_error_failover_defaults_to_continue_and_accepts_explicit_stop() {
|
||||
let default_policy =
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use std::collections::{BTreeMap, HashMap};
|
||||
use std::sync::{Mutex, OnceLock};
|
||||
use std::time::{Duration, Instant};
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::OnceLock;
|
||||
use std::time::Duration;
|
||||
|
||||
use aether_admin::provider::quota as admin_provider_quota_pure;
|
||||
use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyRuntimeMetadataUpdate;
|
||||
@@ -21,13 +21,7 @@ use crate::handlers::shared::sync_provider_key_quota_status_snapshot;
|
||||
use crate::log_ids::short_request_id;
|
||||
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)>>;
|
||||
|
||||
static CODEX_QUOTA_HEADER_FINGERPRINT_CACHE: OnceLock<HeaderFingerprintCache> = OnceLock::new();
|
||||
static GROK_CHINESE_WAIT_DURATION_RE: OnceLock<Regex> = OnceLock::new();
|
||||
static GROK_ENGLISH_WAIT_DURATION_RE: OnceLock<Regex> = OnceLock::new();
|
||||
|
||||
@@ -62,10 +56,6 @@ pub(crate) async fn apply_local_report_effect(state: &AppState, effect: LocalRep
|
||||
}
|
||||
}
|
||||
|
||||
fn codex_quota_header_fingerprint_cache() -> &'static HeaderFingerprintCache {
|
||||
CODEX_QUOTA_HEADER_FINGERPRINT_CACHE.get_or_init(|| Mutex::new(HashMap::new()))
|
||||
}
|
||||
|
||||
fn report_context_key_id(report_context: Option<&Value>) -> Option<String> {
|
||||
report_context
|
||||
.and_then(|context| context.get("key_id"))
|
||||
@@ -75,6 +65,20 @@ fn report_context_key_id(report_context: Option<&Value>) -> Option<String> {
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn report_context_u64(report_context: Option<&Value>, key: &str) -> Option<u64> {
|
||||
report_context
|
||||
.and_then(|context| context.get(key))
|
||||
.and_then(admin_provider_quota_pure::coerce_json_u64)
|
||||
}
|
||||
|
||||
fn report_context_string<'a>(report_context: Option<&'a Value>, key: &str) -> Option<&'a str> {
|
||||
report_context
|
||||
.and_then(|context| context.get(key))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
fn report_context_provider_response_headers(
|
||||
report_context: Option<&Value>,
|
||||
) -> Option<BTreeMap<String, String>> {
|
||||
@@ -91,86 +95,6 @@ fn report_context_provider_response_headers(
|
||||
(!out.is_empty()).then_some(out)
|
||||
}
|
||||
|
||||
fn is_volatile_compare_field(key: &str) -> bool {
|
||||
key == "updated_at" || key.ends_with("_reset_seconds") || key.ends_with("_reset_after_seconds")
|
||||
}
|
||||
|
||||
fn canonicalize_value(value: &Value) -> Value {
|
||||
match value {
|
||||
Value::Array(items) => Value::Array(items.iter().map(canonicalize_value).collect()),
|
||||
Value::Object(object) => {
|
||||
let mut entries = object.iter().collect::<Vec<_>>();
|
||||
entries.sort_by(|left, right| left.0.cmp(right.0));
|
||||
let mut normalized = serde_json::Map::new();
|
||||
for (key, value) in entries {
|
||||
normalized.insert(key.clone(), canonicalize_value(value));
|
||||
}
|
||||
Value::Object(normalized)
|
||||
}
|
||||
_ => value.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn fingerprint_codex_payload(value: &Value) -> Option<String> {
|
||||
let object = value.as_object()?;
|
||||
let mut entries = object
|
||||
.iter()
|
||||
.filter(|(key, _)| !is_volatile_compare_field(key))
|
||||
.collect::<Vec<_>>();
|
||||
entries.sort_by(|left, right| left.0.cmp(right.0));
|
||||
|
||||
let mut normalized = serde_json::Map::new();
|
||||
for (key, value) in entries {
|
||||
normalized.insert(key.clone(), canonicalize_value(value));
|
||||
}
|
||||
serde_json::to_string(&Value::Object(normalized)).ok()
|
||||
}
|
||||
|
||||
fn get_cached_codex_quota_fingerprint(key_id: &str, now: Instant) -> Option<String> {
|
||||
let mut cache = codex_quota_header_fingerprint_cache()
|
||||
.lock()
|
||||
.expect("codex realtime quota cache should lock");
|
||||
match cache.get(key_id) {
|
||||
Some((fingerprint, expires_at)) if *expires_at > now => Some(fingerprint.clone()),
|
||||
Some(_) => {
|
||||
cache.remove(key_id);
|
||||
None
|
||||
}
|
||||
None => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn set_cached_codex_quota_fingerprint(key_id: &str, fingerprint: String, now: Instant) {
|
||||
let mut cache = codex_quota_header_fingerprint_cache()
|
||||
.lock()
|
||||
.expect("codex realtime quota cache should lock");
|
||||
cache.insert(
|
||||
key_id.to_string(),
|
||||
(
|
||||
fingerprint,
|
||||
now.checked_add(Duration::from_secs(CODEX_QUOTA_CACHE_TTL_SECONDS))
|
||||
.unwrap_or(now),
|
||||
),
|
||||
);
|
||||
|
||||
cache.retain(|_, (_, expires_at)| *expires_at > now);
|
||||
if cache.len() <= CODEX_QUOTA_CACHE_MAX_ENTRIES {
|
||||
return;
|
||||
}
|
||||
|
||||
let mut entries = cache
|
||||
.iter()
|
||||
.map(|(key, (_, expires_at))| (key.clone(), *expires_at))
|
||||
.collect::<Vec<_>>();
|
||||
entries.sort_by_key(|entry| entry.1);
|
||||
for (key, _) in entries
|
||||
.into_iter()
|
||||
.take(cache.len() - CODEX_QUOTA_CACHE_MAX_ENTRIES)
|
||||
{
|
||||
cache.remove(&key);
|
||||
}
|
||||
}
|
||||
|
||||
fn merge_metadata_object(
|
||||
current: Option<&Value>,
|
||||
section_key: &str,
|
||||
@@ -593,21 +517,23 @@ async fn sync_grok_quota_from_report_context(
|
||||
|
||||
async fn apply_local_sync_report_effect(state: &AppState, payload: &GatewaySyncReportRequest) {
|
||||
apply_local_gemini_file_mapping_report_effect(state, payload).await;
|
||||
if let Err(err) = sync_codex_quota_from_response_headers(
|
||||
state,
|
||||
payload.report_context.as_ref(),
|
||||
&payload.headers,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
event_name = "codex_realtime_quota_sync_failed",
|
||||
log_type = "ops",
|
||||
report_kind = %payload.report_kind,
|
||||
report_request_id = %short_request_id(report_request_id(payload.report_context.as_ref())),
|
||||
error = ?err,
|
||||
"gateway failed to persist codex realtime quota from sync response headers"
|
||||
);
|
||||
if (200..300).contains(&payload.status_code) {
|
||||
if let Err(err) = sync_codex_quota_from_response_headers(
|
||||
state,
|
||||
payload.report_context.as_ref(),
|
||||
&payload.headers,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
event_name = "codex_realtime_quota_sync_failed",
|
||||
log_type = "ops",
|
||||
report_kind = %payload.report_kind,
|
||||
report_request_id = %short_request_id(report_request_id(payload.report_context.as_ref())),
|
||||
error = ?err,
|
||||
"gateway failed to persist codex realtime quota from sync response headers"
|
||||
);
|
||||
}
|
||||
}
|
||||
if let Err(err) = sync_grok_quota_from_report_context(
|
||||
state,
|
||||
@@ -646,21 +572,23 @@ async fn apply_local_sync_report_effect(state: &AppState, payload: &GatewaySyncR
|
||||
}
|
||||
|
||||
async fn apply_local_stream_report_effect(state: &AppState, payload: &GatewayStreamReportRequest) {
|
||||
if let Err(err) = sync_codex_quota_from_response_headers(
|
||||
state,
|
||||
payload.report_context.as_ref(),
|
||||
&payload.headers,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
event_name = "codex_realtime_quota_sync_failed",
|
||||
log_type = "ops",
|
||||
report_kind = %payload.report_kind,
|
||||
report_request_id = %short_request_id(report_request_id(payload.report_context.as_ref())),
|
||||
error = ?err,
|
||||
"gateway failed to persist codex realtime quota from stream response headers"
|
||||
);
|
||||
if (200..300).contains(&payload.status_code) {
|
||||
if let Err(err) = sync_codex_quota_from_response_headers(
|
||||
state,
|
||||
payload.report_context.as_ref(),
|
||||
&payload.headers,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
event_name = "codex_realtime_quota_sync_failed",
|
||||
log_type = "ops",
|
||||
report_kind = %payload.report_kind,
|
||||
report_request_id = %short_request_id(report_request_id(payload.report_context.as_ref())),
|
||||
error = ?err,
|
||||
"gateway failed to persist codex realtime quota from stream response headers"
|
||||
);
|
||||
}
|
||||
}
|
||||
if let Err(err) = sync_grok_quota_from_report_context(
|
||||
state,
|
||||
@@ -840,103 +768,112 @@ async fn sync_codex_quota_from_response_headers(
|
||||
};
|
||||
|
||||
let now_unix_secs = current_unix_secs();
|
||||
let observed_at_unix_secs = report_context_u64(
|
||||
report_context,
|
||||
"provider_response_headers_observed_at_unix_ms",
|
||||
)
|
||||
.map(|value| value / 1_000)
|
||||
.filter(|value| *value > 0)
|
||||
.unwrap_or(now_unix_secs);
|
||||
let request_started_at_unix_ms =
|
||||
report_context_u64(report_context, "provider_request_started_at_unix_ms");
|
||||
let request_order_id = report_context_string(report_context, "provider_request_order_id");
|
||||
let observed_reset_generation =
|
||||
report_context_u64(report_context, "codex_quota_reset_generation");
|
||||
let observed_credential_generation =
|
||||
report_context_string(report_context, "codex_credential_generation");
|
||||
let provider_headers = report_context_provider_response_headers(report_context);
|
||||
let parsed_from_provider_headers = provider_headers.as_ref().and_then(|headers| {
|
||||
admin_provider_quota_pure::parse_codex_usage_headers(headers, now_unix_secs)
|
||||
admin_provider_quota_pure::parse_codex_usage_headers(headers, observed_at_unix_secs)
|
||||
});
|
||||
let Some(parsed) = parsed_from_provider_headers
|
||||
.or_else(|| admin_provider_quota_pure::parse_codex_usage_headers(headers, now_unix_secs))
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
let Some(incoming_fingerprint) = fingerprint_codex_payload(&parsed) else {
|
||||
let Some(parsed) = parsed_from_provider_headers.or_else(|| {
|
||||
admin_provider_quota_pure::parse_codex_usage_headers(headers, observed_at_unix_secs)
|
||||
}) else {
|
||||
return Ok(false);
|
||||
};
|
||||
// Runtime headers can be partial (for example only the primary window),
|
||||
// so absence never authoritatively removes another stored window.
|
||||
let coverage = admin_provider_quota_pure::CodexQuotaWindowCoverage::Patch;
|
||||
|
||||
let now = Instant::now();
|
||||
if get_cached_codex_quota_fingerprint(&key_id, now).as_deref()
|
||||
== Some(incoming_fingerprint.as_str())
|
||||
{
|
||||
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("codex") {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
let expected_namespace_value =
|
||||
upstream_metadata_namespace_value(key.upstream_metadata.as_ref(), "codex");
|
||||
let Some(outcome) = admin_provider_quota_pure::merge_codex_quota_metadata_snapshot(
|
||||
expected_namespace_value.as_ref(),
|
||||
&parsed,
|
||||
admin_provider_quota_pure::CodexQuotaMergeContext {
|
||||
observed_at_unix_secs,
|
||||
request_started_at_unix_ms,
|
||||
request_order_id,
|
||||
observed_reset_generation,
|
||||
authoritative_reset_generation: None,
|
||||
observed_credential_generation,
|
||||
account_reset_fence_id: None,
|
||||
coverage,
|
||||
},
|
||||
) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if !outcome.changed {
|
||||
return Ok(false);
|
||||
}
|
||||
let next_codex = outcome.metadata;
|
||||
let updated_upstream_metadata =
|
||||
merge_metadata_object(key.upstream_metadata.as_ref(), "codex", next_codex.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 updated = state
|
||||
.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: next_codex,
|
||||
status_snapshot_patch: quota_status_snapshot_patch(
|
||||
updated_status_snapshot.as_ref(),
|
||||
),
|
||||
updated_at_unix_secs: Some(observed_at_unix_secs),
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
if updated {
|
||||
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;
|
||||
}
|
||||
}
|
||||
|
||||
let Some(key) = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
set_cached_codex_quota_fingerprint(&key_id, incoming_fingerprint.clone(), now);
|
||||
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 {
|
||||
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.clone(), now);
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
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_default();
|
||||
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.clone(), now);
|
||||
return Ok(false);
|
||||
};
|
||||
if current_fingerprint == incoming_fingerprint {
|
||||
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.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 updated = state
|
||||
.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.clone(), now);
|
||||
return Ok(true);
|
||||
}
|
||||
// Response headers describe an authoritative quota snapshot. A CAS
|
||||
// conflict means a newer snapshot/delta won, so avoid replaying stale
|
||||
// data over it.
|
||||
Ok(false)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn clear_local_report_effect_caches_for_tests() {
|
||||
if let Some(cache) = CODEX_QUOTA_HEADER_FINGERPRINT_CACHE.get() {
|
||||
cache
|
||||
.lock()
|
||||
.expect("codex realtime quota cache should lock")
|
||||
.clear();
|
||||
}
|
||||
}
|
||||
pub(crate) fn clear_local_report_effect_caches_for_tests() {}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
@@ -950,6 +887,151 @@ mod tests {
|
||||
|
||||
use crate::data::GatewayDataState;
|
||||
|
||||
fn codex_headers(used_percent: f64, reset_at: u64) -> BTreeMap<String, String> {
|
||||
BTreeMap::from([
|
||||
("x-codex-plan-type".to_string(), "free".to_string()),
|
||||
(
|
||||
"x-codex-primary-used-percent".to_string(),
|
||||
used_percent.to_string(),
|
||||
),
|
||||
("x-codex-primary-reset-at".to_string(), reset_at.to_string()),
|
||||
(
|
||||
"x-codex-primary-window-minutes".to_string(),
|
||||
"300".to_string(),
|
||||
),
|
||||
])
|
||||
}
|
||||
|
||||
fn codex_test_state(key_id: &str) -> AppState {
|
||||
let provider = StoredProviderCatalogProvider::new(
|
||||
"codex-provider".to_string(),
|
||||
"Codex".to_string(),
|
||||
None,
|
||||
"codex".to_string(),
|
||||
)
|
||||
.expect("provider should build");
|
||||
let key = StoredProviderCatalogKey::new(
|
||||
key_id.to_string(),
|
||||
provider.id.clone(),
|
||||
"Codex Key".to_string(),
|
||||
"oauth".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build");
|
||||
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
vec![],
|
||||
vec![key],
|
||||
));
|
||||
AppState::new()
|
||||
.expect("gateway state should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_repository_for_tests(repository),
|
||||
)
|
||||
}
|
||||
|
||||
async fn stored_codex_metadata(state: &AppState, key_id: &str) -> Value {
|
||||
state
|
||||
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
|
||||
.await
|
||||
.expect("key should reload")
|
||||
.pop()
|
||||
.and_then(|key| key.upstream_metadata)
|
||||
.and_then(|metadata| metadata.get("codex").cloned())
|
||||
.expect("codex metadata should exist")
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn out_of_order_codex_quota_reports_keep_max_usage_and_allow_next_reset() {
|
||||
clear_local_report_effect_caches_for_tests();
|
||||
let key_id = "codex-out-of-order-quota-key";
|
||||
let state = codex_test_state(key_id);
|
||||
let report_context = json!({"key_id": key_id});
|
||||
let reset_at = 2_000_000_000u64;
|
||||
let higher = codex_headers(60.0, reset_at);
|
||||
let lower = codex_headers(50.0, reset_at);
|
||||
|
||||
let (higher_result, lower_result) = tokio::join!(
|
||||
sync_codex_quota_from_response_headers(&state, Some(&report_context), &higher),
|
||||
sync_codex_quota_from_response_headers(&state, Some(&report_context), &lower),
|
||||
);
|
||||
higher_result.expect("higher quota report should complete");
|
||||
lower_result.expect("lower quota report should complete");
|
||||
|
||||
let stored = stored_codex_metadata(&state, key_id).await;
|
||||
assert_eq!(stored["primary_used_percent"], json!(60.0));
|
||||
assert_eq!(stored["primary_reset_at"], json!(reset_at));
|
||||
|
||||
let next_reset_at = reset_at + 18_000;
|
||||
assert!(sync_codex_quota_from_response_headers(
|
||||
&state,
|
||||
Some(&report_context),
|
||||
&codex_headers(2.0, next_reset_at),
|
||||
)
|
||||
.await
|
||||
.expect("new reset quota report should complete"));
|
||||
let reset = stored_codex_metadata(&state, key_id).await;
|
||||
assert_eq!(reset["primary_used_percent"], json!(2.0));
|
||||
assert_eq!(reset["primary_reset_at"], json!(next_reset_at));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn non_success_sync_report_does_not_persist_codex_quota_headers() {
|
||||
let key_id = "codex-non-success-sync-report";
|
||||
let state = codex_test_state(key_id);
|
||||
let payload = GatewaySyncReportRequest {
|
||||
trace_id: "trace-non-success-sync".to_string(),
|
||||
report_kind: "openai_responses_sync_error".to_string(),
|
||||
report_context: Some(json!({"key_id": key_id})),
|
||||
status_code: 401,
|
||||
headers: codex_headers(85.0, 2_000_000_000),
|
||||
body_json: None,
|
||||
client_body_json: None,
|
||||
body_base64: None,
|
||||
telemetry: None,
|
||||
};
|
||||
|
||||
apply_local_report_effect(&state, LocalReportEffect::Sync { payload: &payload }).await;
|
||||
|
||||
let stored = state
|
||||
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
|
||||
.await
|
||||
.expect("key should reload")
|
||||
.pop()
|
||||
.expect("key should exist");
|
||||
assert!(stored.upstream_metadata.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn non_success_stream_report_does_not_persist_codex_quota_headers() {
|
||||
let key_id = "codex-non-success-stream-report";
|
||||
let state = codex_test_state(key_id);
|
||||
let payload = GatewayStreamReportRequest {
|
||||
trace_id: "trace-non-success-stream".to_string(),
|
||||
report_kind: "openai_responses_stream_error".to_string(),
|
||||
report_context: Some(json!({"key_id": key_id})),
|
||||
status_code: 429,
|
||||
headers: codex_headers(95.0, 2_000_000_000),
|
||||
provider_body_base64: None,
|
||||
provider_body_state: None,
|
||||
client_body_base64: None,
|
||||
client_body_state: None,
|
||||
terminal_summary: None,
|
||||
telemetry: None,
|
||||
};
|
||||
|
||||
apply_local_report_effect(&state, LocalReportEffect::Stream { payload: &payload }).await;
|
||||
|
||||
let stored = state
|
||||
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
|
||||
.await
|
||||
.expect("key should reload")
|
||||
.pop()
|
||||
.expect("key should exist");
|
||||
assert!(stored.upstream_metadata.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gemini_report_metadata_write_preserves_adaptive_and_other_provider_state() {
|
||||
let provider = StoredProviderCatalogProvider::new(
|
||||
|
||||
Reference in New Issue
Block a user