fix(codex): fence concurrent quota updates

This commit is contained in:
elky
2026-08-14 09:28:07 +08:00
parent f3a12c1008
commit 5b0c763086
65 changed files with 13009 additions and 738 deletions
+22 -8
View File
@@ -1,14 +1,14 @@
use super::{ use super::{
ApiKeyLastUsedDelta, DataLayerError, GatewayDataState, GeminiFileMappingListQuery, ApiKeyLastUsedDelta, DataLayerError, GatewayDataState, GeminiFileMappingListQuery,
GeminiFileMappingStats, ProviderCatalogKeyAdaptiveStateUpdate, GeminiFileMappingStats, ProviderCatalogKeyAdaptiveStateUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery, ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate,
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
PublicHealthStatusCount, PublicHealthTimelineBucket, StoredGeminiFileMapping, ProviderCatalogKeyStatusSnapshotUpdate, PublicHealthStatusCount, PublicHealthTimelineBucket,
StoredGeminiFileMappingListPage, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredGeminiFileMapping, StoredGeminiFileMappingListPage, StoredProviderCatalogEndpoint,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage, StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary,
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, StoredRequestCandidate, StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord, StoredRequestCandidate, UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord,
}; };
impl GatewayDataState { impl GatewayDataState {
@@ -561,6 +561,20 @@ impl GatewayDataState {
Ok(updated) Ok(updated)
} }
pub(crate) async fn compare_and_update_provider_catalog_key_admin_state(
&self,
update: &ProviderCatalogKeyAdminCasUpdate,
) -> Result<bool, DataLayerError> {
let updated = match &self.provider_catalog_writer {
Some(repository) => repository.compare_and_update_key_admin_state(update).await,
None => Ok(false),
}?;
// Clear on both success and conflict so a retry cannot reuse the stale
// credential snapshot that lost the CAS.
self.clear_provider_catalog_cache();
Ok(updated)
}
pub(crate) async fn update_provider_catalog_keys( pub(crate) async fn update_provider_catalog_keys(
&self, &self,
keys: &[StoredProviderCatalogKey], keys: &[StoredProviderCatalogKey],
+7 -7
View File
@@ -121,13 +121,13 @@ use aether_data_contracts::repository::pool_scores::{
UpsertPoolMemberScore, UpsertPoolMemberScore,
}; };
use aether_data_contracts::repository::provider_catalog::{ use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdminCasUpdate,
ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery,
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, ProviderCatalogReadRepository, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage, StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary,
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
}; };
use aether_data_contracts::repository::quota::{ use aether_data_contracts::repository::quota::{
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot, ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
@@ -2573,6 +2573,7 @@ fn json_execution_result(
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code, status_code,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]), headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: Some(body), json_body: Some(body),
body_bytes_b64: None, body_bytes_b64: None,
@@ -2615,6 +2616,7 @@ fn bytes_execution_result(
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code, status_code,
headers, headers,
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: None, json_body: None,
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)), body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
@@ -2637,6 +2639,7 @@ fn execution_result_frame_stream(
payload: StreamFramePayload::Headers { payload: StreamFramePayload::Headers {
status_code: result.status_code, status_code: result.status_code,
headers: result.headers.clone(), headers: result.headers.clone(),
response_observation: result.response_observation.clone(),
}, },
}, },
StreamFrame { StreamFrame {
@@ -480,6 +480,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 502, status_code: 502,
headers: Default::default(), headers: Default::default(),
response_observation: None,
body: None, body: None,
telemetry: None, telemetry: None,
error: None, error: None,
@@ -519,6 +520,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 502, status_code: 502,
headers: Default::default(), headers: Default::default(),
response_observation: None,
body: None, body: None,
telemetry: None, telemetry: None,
error: None, error: None,
@@ -583,6 +585,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 429, status_code: 429,
headers: Default::default(), headers: Default::default(),
response_observation: None,
body: None, body: None,
telemetry: None, telemetry: None,
error: None, error: None,
@@ -614,6 +617,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 401, status_code: 401,
headers: Default::default(), headers: Default::default(),
response_observation: None,
body: None, body: None,
telemetry: None, telemetry: None,
error: None, error: None,
@@ -645,6 +649,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 502, status_code: 502,
headers: Default::default(), headers: Default::default(),
response_observation: None,
body: None, body: None,
telemetry: None, telemetry: None,
error: Some(ExecutionError { error: Some(ExecutionError {
@@ -708,6 +713,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 404, status_code: 404,
headers: Default::default(), headers: Default::default(),
response_observation: None,
body: None, body: None,
telemetry: None, telemetry: None,
error: None, error: None,
@@ -900,6 +906,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 200, status_code: 200,
headers: Default::default(), headers: Default::default(),
response_observation: None,
body: None, body: None,
telemetry: None, telemetry: None,
error: None, error: None,
@@ -1022,6 +1029,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 429, status_code: 429,
headers: Default::default(), headers: Default::default(),
response_observation: None,
body: None, body: None,
telemetry: None, telemetry: None,
error: None, error: None,
@@ -1068,6 +1076,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 429, status_code: 429,
headers: Default::default(), headers: Default::default(),
response_observation: None,
body: None, body: None,
telemetry: None, telemetry: None,
error: None, error: None,
@@ -1174,6 +1183,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 200, status_code: 200,
headers: Default::default(), headers: Default::default(),
response_observation: None,
body: None, body: None,
telemetry: None, telemetry: None,
error: None, error: None,
@@ -1211,6 +1221,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 400, status_code: 400,
headers: Default::default(), headers: Default::default(),
response_observation: None,
body: None, body: None,
telemetry: None, telemetry: None,
error: None, error: None,
@@ -1259,6 +1270,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 429, status_code: 429,
headers: Default::default(), headers: Default::default(),
response_observation: None,
body: None, body: None,
telemetry: None, telemetry: None,
error: None, error: None,
@@ -841,6 +841,7 @@ fn encode_grok_headers_frame(
payload: StreamFramePayload::Headers { payload: StreamFramePayload::Headers {
status_code, status_code,
headers, headers,
response_observation: None,
}, },
}) })
} }
@@ -2157,6 +2158,7 @@ fn grok_execution_result(
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code, status_code,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]), headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: Some(body_json), json_body: Some(body_json),
body_bytes_b64: None, body_bytes_b64: None,
@@ -2220,6 +2222,7 @@ fn grok_collected_frame_stream(
"application/json".to_string() "application/json".to_string()
}, },
)]), )]),
response_observation: None,
}, },
}, },
StreamFrame { StreamFrame {
@@ -279,6 +279,7 @@ fn raw_response_frame_stream(
payload: StreamFramePayload::Headers { payload: StreamFramePayload::Headers {
status_code, status_code,
headers, headers,
response_observation: None,
}, },
}, },
StreamFrame { StreamFrame {
@@ -1449,6 +1450,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 200, status_code: 200,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"jsonrpc": "2.0", "jsonrpc": "2.0",
@@ -2,10 +2,10 @@ use aether_contracts::ExecutionPlan;
use tracing::warn; use tracing::warn;
use crate::orchestration::{ use crate::orchestration::{
oauth_status_may_be_invalid as status_may_be_oauth_invalid, local_failover_error_message, oauth_status_may_be_invalid as status_may_be_oauth_invalid,
oauth_status_proves_access_token_invalid as status_proves_access_token_invalid, oauth_status_proves_access_token_invalid as status_proves_access_token_invalid,
}; };
use crate::state::AgentIdentityAuthConfigFence; use crate::state::{AgentIdentityAuthConfigFence, CodexRuntimeOAuthObservation};
use crate::{provider_transport::LocalOAuthRefreshError, AppState}; use crate::{provider_transport::LocalOAuthRefreshError, AppState};
pub(crate) async fn refresh_oauth_plan_auth_for_retry( pub(crate) async fn refresh_oauth_plan_auth_for_retry(
@@ -14,6 +14,9 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
status_code: u16, status_code: u16,
response_text: Option<&str>, response_text: Option<&str>,
trace_id: &str, trace_id: &str,
report_context: Option<&serde_json::Value>,
request_started_at_unix_ms: Option<u64>,
request_order_id: Option<&str>,
) -> bool { ) -> bool {
if !status_may_be_oauth_invalid(status_code, response_text) { if !status_may_be_oauth_invalid(status_code, response_text) {
return false; return false;
@@ -109,15 +112,49 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
body_excerpt, body_excerpt,
.. ..
}) if matches!(refresh_status_code, 400 | 401 | 403) => { }) if matches!(refresh_status_code, 400 | 401 | 403) => {
if let Err(err) = state let observed_credential_generation =
.persist_local_oauth_refresh_failure_state( report_context_string(report_context, "codex_credential_generation");
&transport, let runtime_invalid_message = local_failover_error_message(response_text);
refresh_status_code, let runtime_invalid_reason =
body_excerpt.as_str(), aether_admin::provider::quota::codex_runtime_invalid_reason(
access_token_invalid_proven, status_code,
) runtime_invalid_message.as_deref(),
.await );
{ let persist_result = match (request_started_at_unix_ms, request_order_id) {
(Some(request_started_at_unix_ms), Some(request_order_id))
if transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("codex") =>
{
state
.persist_local_oauth_refresh_failure_state_observed(
&transport,
refresh_status_code,
body_excerpt.as_str(),
access_token_invalid_proven,
CodexRuntimeOAuthObservation {
request_started_at_unix_ms,
request_order_id,
observed_credential_generation,
runtime_invalid_reason: runtime_invalid_reason.as_deref(),
},
)
.await
}
_ => {
state
.persist_local_oauth_refresh_failure_state(
&transport,
refresh_status_code,
body_excerpt.as_str(),
access_token_invalid_proven,
)
.await
}
};
if let Err(err) = persist_result {
warn!( warn!(
event_name = "local_oauth_retry_refresh_failure_persist_failed", event_name = "local_oauth_retry_refresh_failure_persist_failed",
log_type = "ops", log_type = "ops",
@@ -161,6 +198,17 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
} }
} }
fn report_context_string<'a>(
report_context: Option<&'a serde_json::Value>,
field: &str,
) -> Option<&'a str> {
report_context
.and_then(|context| context.get(field))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
}
fn execution_plan_authorization(plan: &ExecutionPlan) -> Option<&str> { fn execution_plan_authorization(plan: &ExecutionPlan) -> Option<&str> {
plan.headers plan.headers
.iter() .iter()
@@ -209,6 +257,7 @@ mod tests {
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data_contracts::repository::provider_catalog::{ use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyOAuthCredentialFence,
ProviderCatalogReadRepository, ProviderCatalogWriteRepository, ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
}; };
@@ -312,7 +361,7 @@ mod tests {
} }
#[tokio::test] #[tokio::test]
async fn auto_removes_codex_key_after_request_proven_terminal_refresh_failure() { async fn retains_codex_key_after_request_proven_terminal_refresh_failure() {
let token_hits = Arc::new(Mutex::new(0usize)); let token_hits = Arc::new(Mutex::new(0usize));
let token_hits_clone = Arc::clone(&token_hits); let token_hits_clone = Arc::clone(&token_hits);
let token_server = Router::new().route( let token_server = Router::new().route(
@@ -458,16 +507,34 @@ mod tests {
401, 401,
Some(r#"{"error":"oauth_token_invalid"}"#), Some(r#"{"error":"oauth_token_invalid"}"#),
"trace-oauth-retry", "trace-oauth-retry",
None,
Some(1_000),
Some("01900000-0000-7000-8000-000000000010"),
) )
.await; .await;
assert!(!retried); assert!(!retried);
assert_eq!(*token_hits.lock().expect("mutex should lock"), 1); assert_eq!(*token_hits.lock().expect("mutex should lock"), 1);
let keys = provider_catalog_repository let stored_key = provider_catalog_repository
.list_keys_by_ids(&["key-codex-oauth-retry".to_string()]) .list_keys_by_ids(&["key-codex-oauth-retry".to_string()])
.await .await
.expect("keys should read"); .expect("keys should read")
assert!(keys.is_empty()); .into_iter()
.next()
.expect("request-scoped refresh failure should retain the key");
let invalid_reason = stored_key
.oauth_invalid_reason
.as_deref()
.expect("combined invalid reason should persist");
assert!(invalid_reason.contains("[OAUTH_EXPIRED]"));
assert!(invalid_reason.contains("[REFRESH_FAILED]"));
assert_eq!(
stored_key
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/codex/oauth_state_request_id")),
Some(&json!("01900000-0000-7000-8000-000000000010"))
);
token_handle.abort(); token_handle.abort();
} }
@@ -619,6 +686,9 @@ mod tests {
401, 401,
Some(r#"{"error":"invalid_token"}"#), Some(r#"{"error":"invalid_token"}"#),
"trace-claude-oauth-fence-first", "trace-claude-oauth-fence-first",
None,
None,
None,
) )
.await .await
); );
@@ -647,6 +717,9 @@ mod tests {
401, 401,
Some(r#"{"error":"invalid_token"}"#), Some(r#"{"error":"invalid_token"}"#),
"trace-claude-oauth-fence-stale", "trace-claude-oauth-fence-stale",
None,
None,
None,
) )
.await .await
); );
@@ -665,6 +738,7 @@ mod tests {
.expect("Claude key should load") .expect("Claude key should load")
.pop() .pop()
.expect("Claude key should exist"); .expect("Claude key should exist");
let expected_admin_replacement = admin_replacement.clone();
admin_replacement.encrypted_api_key = Some( admin_replacement.encrypted_api_key = Some(
encrypt_python_fernet_plaintext( encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY, DEVELOPMENT_ENCRYPTION_KEY,
@@ -673,10 +747,23 @@ mod tests {
.expect("admin access token should encrypt"), .expect("admin access token should encrypt"),
); );
admin_replacement.expires_at_unix_secs = Some(4_102_444_800); admin_replacement.expires_at_unix_secs = Some(4_102_444_800);
provider_catalog_repository assert!(provider_catalog_repository
.update_key(&admin_replacement) .compare_and_update_key_admin_state(&ProviderCatalogKeyAdminCasUpdate {
expected_encrypted_auth_config: expected_admin_replacement
.encrypted_auth_config
.clone(),
expected_credential: ProviderCatalogKeyOAuthCredentialFence {
encrypted_api_key: expected_admin_replacement.encrypted_api_key.clone(),
auth_type: expected_admin_replacement.auth_type.clone(),
provider_id: expected_admin_replacement.provider_id.clone(),
provider_type: "claude_code".to_string(),
},
key: admin_replacement,
codex_rotation: None,
reset_oauth_runtime: true,
})
.await .await
.expect("admin replacement should persist"); .expect("admin replacement CAS should run"));
let admin_result = state let admin_result = state
.force_local_oauth_refresh_entry(&stale_transport) .force_local_oauth_refresh_entry(&stale_transport)
@@ -10,6 +10,10 @@ use crate::{AppState, GatewayError};
const RESPONSE_HEADER_RULES_KEY: &str = "response_header_rules"; const RESPONSE_HEADER_RULES_KEY: &str = "response_header_rules";
const RESPONSE_HEADER_RULES_CAMEL_KEY: &str = "responseHeaderRules"; const RESPONSE_HEADER_RULES_CAMEL_KEY: &str = "responseHeaderRules";
const PROVIDER_RESPONSE_HEADERS_CONTEXT_KEY: &str = "provider_response_headers"; const PROVIDER_RESPONSE_HEADERS_CONTEXT_KEY: &str = "provider_response_headers";
const PROVIDER_REQUEST_STARTED_AT_UNIX_MS_CONTEXT_KEY: &str = "provider_request_started_at_unix_ms";
const PROVIDER_REQUEST_ORDER_ID_CONTEXT_KEY: &str = "provider_request_order_id";
const PROVIDER_RESPONSE_HEADERS_OBSERVED_AT_UNIX_MS_CONTEXT_KEY: &str =
"provider_response_headers_observed_at_unix_ms";
const RESPONSE_HEADER_RULE_PROTECTED_KEYS: &[&str] = &["content-length"]; const RESPONSE_HEADER_RULE_PROTECTED_KEYS: &[&str] = &["content-length"];
const RESPONSE_HEADER_RULES_CACHE_TTL: Duration = Duration::from_secs(5); const RESPONSE_HEADER_RULES_CACHE_TTL: Duration = Duration::from_secs(5);
@@ -98,6 +102,9 @@ pub(crate) async fn apply_endpoint_response_header_rules(
pub(crate) fn attach_provider_response_headers_to_report_context( pub(crate) fn attach_provider_response_headers_to_report_context(
report_context: Option<Value>, report_context: Option<Value>,
provider_headers: &BTreeMap<String, String>, provider_headers: &BTreeMap<String, String>,
provider_request_started_at_unix_ms: u64,
provider_response_headers_observed_at_unix_ms: u64,
provider_request_order_id: &str,
) -> Option<Value> { ) -> Option<Value> {
let provider_headers = serde_json::to_value(provider_headers).ok()?; let provider_headers = serde_json::to_value(provider_headers).ok()?;
let mut object = match report_context { let mut object = match report_context {
@@ -105,9 +112,99 @@ pub(crate) fn attach_provider_response_headers_to_report_context(
Some(other) => Map::from_iter([("seed".to_string(), other)]), Some(other) => Map::from_iter([("seed".to_string(), other)]),
None => Map::new(), None => Map::new(),
}; };
object.insert( let observation_is_absent = !object.contains_key(PROVIDER_RESPONSE_HEADERS_CONTEXT_KEY)
PROVIDER_RESPONSE_HEADERS_CONTEXT_KEY.to_string(), && !object.contains_key(PROVIDER_REQUEST_STARTED_AT_UNIX_MS_CONTEXT_KEY)
provider_headers, && !object.contains_key(PROVIDER_RESPONSE_HEADERS_OBSERVED_AT_UNIX_MS_CONTEXT_KEY)
); && !object.contains_key(PROVIDER_REQUEST_ORDER_ID_CONTEXT_KEY);
if observation_is_absent {
object.insert(
PROVIDER_RESPONSE_HEADERS_CONTEXT_KEY.to_string(),
provider_headers,
);
object.insert(
PROVIDER_REQUEST_STARTED_AT_UNIX_MS_CONTEXT_KEY.to_string(),
Value::from(provider_request_started_at_unix_ms),
);
object.insert(
PROVIDER_RESPONSE_HEADERS_OBSERVED_AT_UNIX_MS_CONTEXT_KEY.to_string(),
Value::from(provider_response_headers_observed_at_unix_ms),
);
object.insert(
PROVIDER_REQUEST_ORDER_ID_CONTEXT_KEY.to_string(),
Value::from(provider_request_order_id),
);
}
Some(Value::Object(object)) Some(Value::Object(object))
} }
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn provider_response_observation_is_first_write_wins() {
let first_headers =
BTreeMap::from([("x-codex-primary-used-percent".to_string(), "10".to_string())]);
let second_headers =
BTreeMap::from([("x-codex-primary-used-percent".to_string(), "20".to_string())]);
let report_context = attach_provider_response_headers_to_report_context(
Some(json!("seed-value")),
&first_headers,
100,
200,
"observation-1",
);
let report_context = attach_provider_response_headers_to_report_context(
report_context,
&second_headers,
300,
400,
"observation-2",
)
.expect("report context should exist");
assert_eq!(report_context["seed"], json!("seed-value"));
assert_eq!(
report_context["provider_response_headers"]["x-codex-primary-used-percent"],
json!("10")
);
assert_eq!(
report_context["provider_request_started_at_unix_ms"],
json!(100)
);
assert_eq!(
report_context["provider_response_headers_observed_at_unix_ms"],
json!(200)
);
assert_eq!(
report_context["provider_request_order_id"],
json!("observation-1")
);
}
#[test]
fn provider_response_observation_does_not_complete_a_partial_triplet() {
let report_context = attach_provider_response_headers_to_report_context(
Some(json!({"provider_response_headers": {"x-existing": "1"}})),
&BTreeMap::from([("x-new".to_string(), "2".to_string())]),
300,
400,
"observation-2",
)
.expect("report context should exist");
assert_eq!(
report_context["provider_response_headers"]["x-existing"],
json!("1")
);
assert!(report_context
.get("provider_request_started_at_unix_ms")
.is_none());
assert!(report_context
.get("provider_response_headers_observed_at_unix_ms")
.is_none());
assert!(report_context.get("provider_request_order_id").is_none());
}
}
@@ -11,8 +11,8 @@ use std::time::{Duration, Instant};
use aether_ai_serving::{AiAttemptExecutionOutcome, AiAttemptRetryScope}; use aether_ai_serving::{AiAttemptExecutionOutcome, AiAttemptRetryScope};
use aether_contracts::{ use aether_contracts::{
ExecutionPlan, ExecutionStreamTerminalSummary, ExecutionTelemetry, StandardizedUsage, ExecutionPlan, ExecutionResponseObservation, ExecutionStreamTerminalSummary,
StreamFrame, StreamFramePayload, ExecutionTelemetry, StandardizedUsage, StreamFrame, StreamFramePayload,
}; };
use aether_data_contracts::repository::candidates::{ use aether_data_contracts::repository::candidates::{
RequestCandidateStatus, UpsertRequestCandidateRecord, RequestCandidateStatus, UpsertRequestCandidateRecord,
@@ -112,12 +112,13 @@ use crate::execution_runtime::{
use crate::log_ids::short_request_id; use crate::log_ids::short_request_id;
use crate::orchestration::{ use crate::orchestration::{
apply_local_execution_effect, build_local_error_flow_metadata, classify_failure_disposition, apply_local_execution_effect, build_local_error_flow_metadata, classify_failure_disposition,
cyber_continue_failover_enabled, trace_upstream_response_body, with_error_flow_report_context, cyber_continue_failover_enabled, spawn_local_oauth_success_effect,
trace_upstream_response_body, with_error_flow_report_context,
with_upstream_response_report_context, FailureDisposition, FailureTokenAction, with_upstream_response_report_context, FailureDisposition, FailureTokenAction,
LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect,
LocalExecutionEffect, LocalExecutionEffectContext, LocalFailoverAnalysis, LocalExecutionEffect, LocalExecutionEffectContext, LocalFailoverAnalysis,
LocalHealthFailureEffect, LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalHealthFailureEffect, LocalHealthSuccessEffect, LocalOAuthInvalidationEffect,
LocalPoolErrorEffect, LocalOAuthSuccessEffect, LocalPoolErrorEffect,
}; };
use crate::provider_pool_demand::{ use crate::provider_pool_demand::{
acquire_provider_pool_in_flight_guard, ProviderPoolInFlightGuard, acquire_provider_pool_in_flight_guard, ProviderPoolInFlightGuard,
@@ -1249,6 +1250,9 @@ async fn execute_in_process_stream_with_oauth_retry(
retry_status_code, retry_status_code,
response_text.as_deref(), response_text.as_deref(),
trace_id, trace_id,
report_context,
Some(execution.response_observation.request_started_at_unix_ms),
Some(&execution.response_observation.request_order_id),
) )
.await .await
{ {
@@ -2818,6 +2822,7 @@ async fn execute_stream_from_direct_passthrough(
stream_precommit_committed: _, stream_precommit_committed: _,
response, response,
started_at: upstream_started_at, started_at: upstream_started_at,
response_observation,
stream_first_byte_timeout, stream_first_byte_timeout,
upstream_target_permit, upstream_target_permit,
} = execution; } = execution;
@@ -2834,8 +2839,23 @@ async fn execute_stream_from_direct_passthrough(
let request_id = plan.request_id.clone(); let request_id = plan.request_id.clone();
let candidate_id = plan.candidate_id.clone(); let candidate_id = plan.candidate_id.clone();
let request_id_for_log = short_request_id(request_id.as_str()); let request_id_for_log = short_request_id(request_id.as_str());
let mut report_context = let mut report_context = attach_provider_response_headers_to_report_context(
attach_provider_response_headers_to_report_context(report_context, &headers); report_context,
&headers,
response_observation.request_started_at_unix_ms,
response_observation.response_headers_observed_at_unix_ms,
&response_observation.request_order_id,
);
spawn_local_oauth_success_effect(
state.clone(),
&plan,
report_context.as_ref(),
LocalOAuthSuccessEffect {
status_code,
request_started_at_unix_ms: Some(response_observation.request_started_at_unix_ms),
request_order_id: Some(&response_observation.request_order_id),
},
);
if status_code == 200 { if status_code == 200 {
seed_kiro_simulated_cache_enabled(state, &plan, &mut report_context).await; seed_kiro_simulated_cache_enabled(state, &plan, &mut report_context).await;
if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) { if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) {
@@ -3819,6 +3839,7 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(), provider_pool_in_flight_guard.take(),
retry_scope_out.as_deref_mut(), retry_scope_out.as_deref_mut(),
retry_fallback_out.as_deref_mut(), retry_fallback_out.as_deref_mut(),
None,
) )
.await; .await;
} }
@@ -3891,6 +3912,7 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(), provider_pool_in_flight_guard.take(),
retry_scope_out.as_deref_mut(), retry_scope_out.as_deref_mut(),
retry_fallback_out.as_deref_mut(), retry_fallback_out.as_deref_mut(),
None,
) )
.await; .await;
} }
@@ -3963,6 +3985,7 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(), provider_pool_in_flight_guard.take(),
retry_scope_out.as_deref_mut(), retry_scope_out.as_deref_mut(),
retry_fallback_out.as_deref_mut(), retry_fallback_out.as_deref_mut(),
None,
) )
.await; .await;
} }
@@ -4035,6 +4058,7 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(), provider_pool_in_flight_guard.take(),
retry_scope_out.as_deref_mut(), retry_scope_out.as_deref_mut(),
retry_fallback_out.as_deref_mut(), retry_fallback_out.as_deref_mut(),
None,
) )
.await; .await;
} }
@@ -4192,6 +4216,15 @@ async fn execute_execution_runtime_stream_inner(
record_stream_pending_lifecycle(state, seed, &mut stage_trace).await; record_stream_pending_lifecycle(state, seed, &mut stage_trace).await;
lifecycle_pending_recorded = true; lifecycle_pending_recorded = true;
} }
let report_context = attach_provider_response_headers_to_report_context(
report_context,
&execution.headers,
execution.response_observation.request_started_at_unix_ms,
execution
.response_observation
.response_headers_observed_at_unix_ms,
&execution.response_observation.request_order_id,
);
let stream_precommit_committed = execution.stream_precommit_committed; let stream_precommit_committed = execution.stream_precommit_committed;
let frame_stream = build_direct_execution_frame_stream(execution).boxed(); let frame_stream = build_direct_execution_frame_stream(execution).boxed();
return execute_stream_from_frame_stream_with_retry_scope( return execute_stream_from_frame_stream_with_retry_scope(
@@ -4211,6 +4244,7 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(), provider_pool_in_flight_guard.take(),
retry_scope_out, retry_scope_out,
retry_fallback_out, retry_fallback_out,
None,
) )
.await; .await;
} }
@@ -4327,6 +4361,15 @@ async fn execute_execution_runtime_stream_inner(
record_stream_pending_lifecycle(state, seed, &mut stage_trace).await; record_stream_pending_lifecycle(state, seed, &mut stage_trace).await;
lifecycle_pending_recorded = true; lifecycle_pending_recorded = true;
} }
let report_context = attach_provider_response_headers_to_report_context(
report_context,
&execution.headers,
execution.response_observation.request_started_at_unix_ms,
execution
.response_observation
.response_headers_observed_at_unix_ms,
&execution.response_observation.request_order_id,
);
let stream_precommit_committed = execution.stream_precommit_committed; let stream_precommit_committed = execution.stream_precommit_committed;
let frame_stream = build_direct_execution_frame_stream(execution).boxed(); let frame_stream = build_direct_execution_frame_stream(execution).boxed();
return execute_stream_from_frame_stream_with_retry_scope( return execute_stream_from_frame_stream_with_retry_scope(
@@ -4346,10 +4389,13 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(), provider_pool_in_flight_guard.take(),
retry_scope_out.as_deref_mut(), retry_scope_out.as_deref_mut(),
retry_fallback_out.as_deref_mut(), retry_fallback_out.as_deref_mut(),
None,
) )
.await; .await;
} }
let remote_request_started_at_unix_ms = current_request_candidate_unix_ms();
let remote_request_order_id = uuid::Uuid::now_v7().to_string();
let response = match post_stream_plan_to_remote_execution_runtime( let response = match post_stream_plan_to_remote_execution_runtime(
state, state,
remote_execution_runtime_base_url, remote_execution_runtime_base_url,
@@ -4431,6 +4477,12 @@ async fn execute_execution_runtime_stream_inner(
)?)); )?));
} }
let remote_response_observed_at_unix_ms = current_request_candidate_unix_ms();
let remote_fallback_observation = ExecutionResponseObservation {
request_started_at_unix_ms: remote_request_started_at_unix_ms,
response_headers_observed_at_unix_ms: remote_response_observed_at_unix_ms,
request_order_id: remote_request_order_id,
};
let frame_stream = response let frame_stream = response
.bytes_stream() .bytes_stream()
.map_err(|err| IoError::other(err.to_string())) .map_err(|err| IoError::other(err.to_string()))
@@ -4452,6 +4504,7 @@ async fn execute_execution_runtime_stream_inner(
provider_pool_in_flight_guard.take(), provider_pool_in_flight_guard.take(),
retry_scope_out.as_deref_mut(), retry_scope_out.as_deref_mut(),
retry_fallback_out.as_deref_mut(), retry_fallback_out.as_deref_mut(),
Some(remote_fallback_observation),
) )
.await; .await;
} }
@@ -5481,6 +5534,7 @@ async fn execute_stream_from_frame_stream(
in_flight_guard, in_flight_guard,
None, None,
None, None,
None,
) )
.await .await
} }
@@ -5503,6 +5557,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
in_flight_guard: Option<ProviderPoolInFlightGuard>, in_flight_guard: Option<ProviderPoolInFlightGuard>,
mut retry_scope_out: Option<&mut AiAttemptRetryScope>, mut retry_scope_out: Option<&mut AiAttemptRetryScope>,
mut retry_fallback_out: Option<&mut Option<Response<Body>>>, mut retry_fallback_out: Option<&mut Option<Response<Body>>>,
fallback_response_observation: Option<ExecutionResponseObservation>,
) -> Result<Option<Response<Body>>, GatewayError> { ) -> Result<Option<Response<Body>>, GatewayError> {
let request_id = plan.request_id.as_str(); let request_id = plan.request_id.as_str();
let request_id_for_log = short_request_id(request_id); let request_id_for_log = short_request_id(request_id);
@@ -5535,14 +5590,37 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
let StreamFramePayload::Headers { let StreamFramePayload::Headers {
status_code, status_code,
mut headers, mut headers,
response_observation,
} = first_frame.payload } = first_frame.payload
else { else {
return Err(GatewayError::Internal( return Err(GatewayError::Internal(
"execution runtime stream must start with headers frame".to_string(), "execution runtime stream must start with headers frame".to_string(),
)); ));
}; };
let mut report_context = let response_observation = response_observation
attach_provider_response_headers_to_report_context(report_context, &headers); .or(fallback_response_observation)
.unwrap_or(ExecutionResponseObservation {
request_started_at_unix_ms: candidate_started_unix_secs,
response_headers_observed_at_unix_ms: current_request_candidate_unix_ms(),
request_order_id: uuid::Uuid::now_v7().to_string(),
});
let mut report_context = attach_provider_response_headers_to_report_context(
report_context,
&headers,
response_observation.request_started_at_unix_ms,
response_observation.response_headers_observed_at_unix_ms,
&response_observation.request_order_id,
);
spawn_local_oauth_success_effect(
state.clone(),
&plan,
report_context.as_ref(),
LocalOAuthSuccessEffect {
status_code,
request_started_at_unix_ms: Some(response_observation.request_started_at_unix_ms),
request_order_id: Some(&response_observation.request_order_id),
},
);
if status_code == 200 { if status_code == 200 {
seed_kiro_simulated_cache_enabled(state, &plan, &mut report_context).await; seed_kiro_simulated_cache_enabled(state, &plan, &mut report_context).await;
if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) { if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) {
@@ -8310,6 +8388,7 @@ mod tests {
"content-type".to_string(), "content-type".to_string(),
"text/event-stream".to_string(), "text/event-stream".to_string(),
)]), )]),
response_observation: None,
}, },
})); }));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame { yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -8389,6 +8468,7 @@ mod tests {
"content-type".to_string(), "content-type".to_string(),
"text/event-stream".to_string(), "text/event-stream".to_string(),
)]), )]),
response_observation: None,
}, },
})); }));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame { yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -8431,6 +8511,7 @@ mod tests {
None, None,
Some(&mut retry_scope), Some(&mut retry_scope),
None, None,
None,
) )
.await .await
.expect("prefetch transport execution should resolve"); .expect("prefetch transport execution should resolve");
@@ -8480,6 +8561,7 @@ mod tests {
"content-type".to_string(), "content-type".to_string(),
"text/event-stream".to_string(), "text/event-stream".to_string(),
)]), )]),
response_observation: None,
}, },
})); }));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame { yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -8522,6 +8604,7 @@ mod tests {
None, None,
Some(&mut retry_scope), Some(&mut retry_scope),
None, None,
None,
) )
.await .await
.expect("prefetch HTTP status execution should resolve"); .expect("prefetch HTTP status execution should resolve");
@@ -8680,6 +8763,7 @@ mod tests {
"content-type".to_string(), "content-type".to_string(),
"text/event-stream".to_string(), "text/event-stream".to_string(),
)]), )]),
response_observation: None,
}, },
})); }));
for chunk in chunks { for chunk in chunks {
@@ -8725,6 +8809,7 @@ mod tests {
None, None,
Some(&mut retry_scope), Some(&mut retry_scope),
Some(&mut fallback_response), Some(&mut fallback_response),
None,
) )
.await .await
.expect("native Anthropic stream execution should succeed"); .expect("native Anthropic stream execution should succeed");
@@ -9364,6 +9449,7 @@ mod tests {
"content-type".to_string(), "content-type".to_string(),
"text/event-stream".to_string(), "text/event-stream".to_string(),
)]), )]),
response_observation: None,
}, },
})); }));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame { yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -9413,6 +9499,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
), ),
) )
.await .await
@@ -9850,6 +9937,7 @@ mod tests {
"content-type".to_string(), "content-type".to_string(),
"text/event-stream".to_string(), "text/event-stream".to_string(),
)]), )]),
response_observation: None,
}, },
})); }));
} }
@@ -11532,6 +11620,7 @@ mod tests {
"content-type".to_string(), "content-type".to_string(),
"text/event-stream".to_string(), "text/event-stream".to_string(),
)]), )]),
response_observation: None,
}, },
})); }));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame { yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -11660,6 +11749,7 @@ mod tests {
"content-type".to_string(), "content-type".to_string(),
"text/event-stream".to_string(), "text/event-stream".to_string(),
)]), )]),
response_observation: None,
}, },
})); }));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame { yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -12384,6 +12474,7 @@ mod tests {
"content-type".to_string(), "content-type".to_string(),
"text/event-stream".to_string(), "text/event-stream".to_string(),
)]), )]),
response_observation: None,
}, },
})); }));
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame { yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
@@ -4,8 +4,9 @@ use std::io::Error as IoError;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use aether_contracts::{ use aether_contracts::{
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionStreamTerminalSummary, ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionResponseObservation,
ExecutionTelemetry, StreamFrame, StreamFramePayload, StreamFrameType, ExecutionStreamTerminalSummary, ExecutionTelemetry, StreamFrame, StreamFramePayload,
StreamFrameType,
}; };
use async_stream::stream; use async_stream::stream;
use axum::body::Bytes; use axum::body::Bytes;
@@ -44,6 +45,7 @@ pub(crate) fn build_direct_execution_frame_stream(
stream_precommit_committed: _, stream_precommit_committed: _,
response, response,
started_at, started_at,
response_observation,
stream_first_byte_timeout, stream_first_byte_timeout,
upstream_target_permit, upstream_target_permit,
} = execution; } = execution;
@@ -108,7 +110,11 @@ pub(crate) fn build_direct_execution_frame_stream(
} }
} }
match encode_headers_frame(status_code, response_headers) { match encode_headers_frame(
status_code,
response_headers,
&response_observation,
) {
Ok(frame) => yield Ok(frame), Ok(frame) => yield Ok(frame),
Err(err) => { Err(err) => {
yield Err(err); yield Err(err);
@@ -153,7 +159,11 @@ pub(crate) fn build_direct_execution_frame_stream(
upstream_bytes, upstream_bytes,
first_byte_timeout, first_byte_timeout,
}) => { }) => {
match encode_headers_frame(status_code, original_headers) { match encode_headers_frame(
status_code,
original_headers,
&response_observation,
) {
Ok(frame) => yield Ok(frame), Ok(frame) => yield Ok(frame),
Err(err) => { Err(err) => {
yield Err(err); yield Err(err);
@@ -192,7 +202,11 @@ pub(crate) fn build_direct_execution_frame_stream(
return; return;
} }
match encode_headers_frame(status_code, headers) { match encode_headers_frame(
status_code,
headers,
&response_observation,
) {
Ok(frame) => yield Ok(frame), Ok(frame) => yield Ok(frame),
Err(err) => { Err(err) => {
yield Err(err); yield Err(err);
@@ -611,12 +625,14 @@ pub(crate) fn build_direct_execution_frame_stream(
fn encode_headers_frame( fn encode_headers_frame(
status_code: u16, status_code: u16,
headers: BTreeMap<String, String>, headers: BTreeMap<String, String>,
response_observation: &ExecutionResponseObservation,
) -> Result<Bytes, IoError> { ) -> Result<Bytes, IoError> {
encode_stream_frame_ndjson(&StreamFrame { encode_stream_frame_ndjson(&StreamFrame {
frame_type: StreamFrameType::Headers, frame_type: StreamFrameType::Headers,
payload: StreamFramePayload::Headers { payload: StreamFramePayload::Headers {
status_code, status_code,
headers, headers,
response_observation: Some(response_observation.clone()),
}, },
}) })
} }
@@ -1606,43 +1622,47 @@ mod tests {
.expect("listener should bind"); .expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve"); let addr = listener.local_addr().expect("local addr should resolve");
let server = tokio::spawn(async move { let server = tokio::spawn(async move {
let app = Router::new().route( let (mut socket, _) = listener.accept().await.expect("client should connect");
"/responses", let mut request = [0_u8; 4096];
post(|| async { let _ = socket
let body = serde_json::json!({ .read(&mut request)
"id": "resp_sync_bridge_123",
"object": "response",
"model": "gpt-5.4",
"status": "completed",
"output": [{
"type": "message",
"id": "msg_sync_bridge_123",
"role": "assistant",
"content": [{
"type": "output_text",
"text": "Hello from buffered JSON stream",
"annotations": []
}]
}],
"usage": {
"input_tokens": 1,
"output_tokens": 2,
"total_tokens": 3
}
});
let mut response = axum::http::Response::new(Body::from(
serde_json::to_vec(&body).expect("json should encode"),
));
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/json"),
);
response
}),
);
axum::serve(listener, app)
.await .await
.expect("server should start"); .expect("request should read");
let body = serde_json::to_vec(&serde_json::json!({
"id": "resp_sync_bridge_123",
"object": "response",
"model": "gpt-5.4",
"status": "completed",
"output": [{
"type": "message",
"id": "msg_sync_bridge_123",
"role": "assistant",
"content": [{
"type": "output_text",
"text": "Hello from buffered JSON stream",
"annotations": []
}]
}],
"usage": {
"input_tokens": 1,
"output_tokens": 2,
"total_tokens": 3
}
}))
.expect("json should encode");
socket
.write_all(
format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\n\r\n",
body.len()
)
.as_bytes(),
)
.await
.expect("headers should write");
socket.flush().await.expect("headers should flush");
tokio::time::sleep(Duration::from_millis(75)).await;
socket.write_all(&body).await.expect("body should write");
}); });
let runtime = DirectSyncExecutionRuntime::new(); let runtime = DirectSyncExecutionRuntime::new();
@@ -1678,6 +1698,12 @@ mod tests {
}) })
.await .await
.expect("stream execution should succeed"); .expect("stream execution should succeed");
let expected_observation = execution.response_observation.clone();
assert!(
expected_observation.response_headers_observed_at_unix_ms
>= expected_observation.request_started_at_unix_ms
);
assert!(!expected_observation.request_order_id.is_empty());
let frames = build_direct_execution_frame_stream(execution) let frames = build_direct_execution_frame_stream(execution)
.map(|item| item.expect("frame should encode")) .map(|item| item.expect("frame should encode"))
@@ -1691,6 +1717,10 @@ mod tests {
let header_frame: Value = let header_frame: Value =
serde_json::from_str(&frames[0]).expect("headers frame should parse"); serde_json::from_str(&frames[0]).expect("headers frame should parse");
let encoded_observation: aether_contracts::ExecutionResponseObservation =
serde_json::from_value(header_frame["payload"]["response_observation"].clone())
.expect("headers frame should retain the response observation");
assert_eq!(encoded_observation, expected_observation);
assert_eq!( assert_eq!(
header_frame header_frame
.get("payload") .get("payload")
@@ -5,8 +5,8 @@ use std::time::{Duration, Instant};
use aether_ai_serving::{AiAttemptExecutionOutcome, AiAttemptRetryScope, UPSTREAM_IS_STREAM_KEY}; use aether_ai_serving::{AiAttemptExecutionOutcome, AiAttemptRetryScope, UPSTREAM_IS_STREAM_KEY};
use aether_contracts::{ use aether_contracts::{
ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionPlan, ExecutionResult, ExecutionError, ExecutionErrorKind, ExecutionPhase, ExecutionPlan,
ExecutionTelemetry, ExecutionResponseObservation, ExecutionResult, ExecutionTelemetry,
}; };
use aether_data_contracts::repository::candidates::RequestCandidateStatus; use aether_data_contracts::repository::candidates::RequestCandidateStatus;
use aether_scheduler_core::{ use aether_scheduler_core::{
@@ -70,11 +70,12 @@ use crate::execution_runtime::{
}; };
use crate::log_ids::short_request_id; use crate::log_ids::short_request_id;
use crate::orchestration::{ use crate::orchestration::{
apply_local_execution_effect, build_local_error_flow_metadata, trace_upstream_response_body, apply_local_execution_effect, build_local_error_flow_metadata,
with_error_flow_report_context, with_upstream_response_report_context, spawn_local_oauth_success_effect, trace_upstream_response_body, with_error_flow_report_context,
LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, with_upstream_response_report_context, LocalAdaptiveRateLimitEffect,
LocalExecutionEffect, LocalExecutionEffectContext, LocalHealthFailureEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, LocalExecutionEffect,
LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalPoolErrorEffect, LocalExecutionEffectContext, LocalHealthFailureEffect, LocalHealthSuccessEffect,
LocalOAuthInvalidationEffect, LocalOAuthSuccessEffect, LocalPoolErrorEffect,
}; };
use crate::provider_pool_demand::acquire_provider_pool_in_flight_guard; use crate::provider_pool_demand::acquire_provider_pool_in_flight_guard;
use crate::request_candidate_runtime::{ use crate::request_candidate_runtime::{
@@ -1379,7 +1380,19 @@ async fn execute_direct_sync_runtime_candidate(
candidate_started_unix_ms, candidate_started_unix_ms,
event.status_code, event.status_code,
event.ttfb_ms, event.ttfb_ms,
) );
spawn_local_oauth_success_effect(
state_for_response_started.clone(),
plan,
report_context,
LocalOAuthSuccessEffect {
status_code: event.status_code,
request_started_at_unix_ms: Some(
event.response_observation.request_started_at_unix_ms,
),
request_order_id: Some(&event.response_observation.request_order_id),
},
);
}) })
.await .await
.map_err(SyncExecutionFailure::from_transport); .map_err(SyncExecutionFailure::from_transport);
@@ -1483,12 +1496,25 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
OpenAiImageSyncProgressRecorder::new(state, plan, report_context, progress_snapshot); OpenAiImageSyncProgressRecorder::new(state, plan, report_context, progress_snapshot);
progress.record_connecting().await; progress.record_connecting().await;
let request_started_at_unix_ms = current_request_candidate_unix_ms();
let request_order_id = uuid::Uuid::now_v7().to_string();
let response = send_request(plan, request_body) let response = send_request(plan, request_body)
.await .await
.map_err(SyncExecutionFailure::from_transport)?; .map_err(SyncExecutionFailure::from_transport)?;
let ttfb_ms = started_at.elapsed().as_millis() as u64; let ttfb_ms = started_at.elapsed().as_millis() as u64;
let response_headers_observed_at_unix_ms = current_request_candidate_unix_ms();
let status_code = response.status_code(); let status_code = response.status_code();
let headers = response.headers(); let headers = response.headers();
spawn_local_oauth_success_effect(
state.clone(),
plan,
report_context,
LocalOAuthSuccessEffect {
status_code,
request_started_at_unix_ms: Some(request_started_at_unix_ms),
request_order_id: Some(&request_order_id),
},
);
progress.record_response_started(status_code, ttfb_ms).await; progress.record_response_started(status_code, ttfb_ms).await;
let mut body_bytes = Vec::new(); let mut body_bytes = Vec::new();
@@ -1569,6 +1595,11 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code, status_code,
headers, headers,
response_observation: Some(ExecutionResponseObservation {
request_started_at_unix_ms,
response_headers_observed_at_unix_ms,
request_order_id,
}),
body, body,
telemetry: Some(ExecutionTelemetry { telemetry: Some(ExecutionTelemetry {
ttfb_ms: Some(ttfb_ms), ttfb_ms: Some(ttfb_ms),
@@ -2461,6 +2492,16 @@ async fn execute_execution_runtime_sync_impl(
}; };
let mut candidate_first_byte_elapsed_ms = let mut candidate_first_byte_elapsed_ms =
calibrated_sync_candidate_first_byte_elapsed_ms(candidate_started_at, &result); calibrated_sync_candidate_first_byte_elapsed_ms(candidate_started_at, &result);
let initial_response_observed_at_unix_ms = current_request_candidate_unix_ms();
let mut provider_response_observation =
result
.response_observation
.clone()
.unwrap_or(ExecutionResponseObservation {
request_started_at_unix_ms: candidate_started_unix_secs,
response_headers_observed_at_unix_ms: initial_response_observed_at_unix_ms,
request_order_id: uuid::Uuid::now_v7().to_string(),
});
let mut oauth_retry_attempted = false; let mut oauth_retry_attempted = false;
let ( let (
result_error_type, result_error_type,
@@ -2473,6 +2514,18 @@ async fn execute_execution_runtime_sync_impl(
local_failover_response_text, local_failover_response_text,
local_failover_analysis, local_failover_analysis,
) = loop { ) = loop {
spawn_local_oauth_success_effect(
state.clone(),
&plan,
report_context.as_ref(),
LocalOAuthSuccessEffect {
status_code: result.status_code,
request_started_at_unix_ms: Some(
provider_response_observation.request_started_at_unix_ms,
),
request_order_id: Some(&provider_response_observation.request_order_id),
},
);
let result_latency_ms = result let result_latency_ms = result
.telemetry .telemetry
.as_ref() .as_ref()
@@ -2534,10 +2587,15 @@ async fn execute_execution_runtime_sync_impl(
result.status_code, result.status_code,
local_failover_response_text.as_deref(), local_failover_response_text.as_deref(),
trace_id, trace_id,
report_context.as_ref(),
Some(provider_response_observation.request_started_at_unix_ms),
Some(&provider_response_observation.request_order_id),
) )
.await .await
{ {
oauth_retry_attempted = true; oauth_retry_attempted = true;
let retry_started_at_unix_ms = current_request_candidate_unix_ms();
let retry_request_order_id = uuid::Uuid::now_v7().to_string();
match crate::execution_runtime::execute_execution_runtime_sync_plan( match crate::execution_runtime::execute_execution_runtime_sync_plan(
state, state,
Some(trace_id), Some(trace_id),
@@ -2546,6 +2604,16 @@ async fn execute_execution_runtime_sync_impl(
.await .await
{ {
Ok(retry_result) => { Ok(retry_result) => {
let retry_response_observed_at_unix_ms = current_request_candidate_unix_ms();
provider_response_observation = retry_result
.response_observation
.clone()
.unwrap_or(ExecutionResponseObservation {
request_started_at_unix_ms: retry_started_at_unix_ms,
response_headers_observed_at_unix_ms:
retry_response_observed_at_unix_ms,
request_order_id: retry_request_order_id,
});
candidate_first_byte_elapsed_ms = candidate_first_byte_elapsed_ms =
calibrated_sync_candidate_first_byte_elapsed_ms( calibrated_sync_candidate_first_byte_elapsed_ms(
candidate_started_at, candidate_started_at,
@@ -2594,6 +2662,13 @@ async fn execute_execution_runtime_sync_impl(
local_failover_analysis, local_failover_analysis,
); );
}; };
let mut report_context = attach_provider_response_headers_to_report_context(
report_context,
&headers,
provider_response_observation.request_started_at_unix_ms,
provider_response_observation.response_headers_observed_at_unix_ms,
&provider_response_observation.request_order_id,
);
if result.status_code >= 400 { if result.status_code >= 400 {
apply_local_execution_effect( apply_local_execution_effect(
state, state,
@@ -2739,8 +2814,6 @@ async fn execute_execution_runtime_sync_impl(
} }
let status_code = result.status_code; let status_code = result.status_code;
let has_body_bytes = body_base64.is_some(); let has_body_bytes = body_base64.is_some();
let mut report_context =
attach_provider_response_headers_to_report_context(report_context, &headers);
if (200..300).contains(&status_code) { if (200..300).contains(&status_code) {
seed_kiro_sync_simulated_cache_enabled(state, &plan, &mut report_context).await; seed_kiro_sync_simulated_cache_enabled(state, &plan, &mut report_context).await;
if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) { if kiro_simulated_cache_enabled_from_report_context(report_context.as_ref()) {
@@ -3231,6 +3304,8 @@ async fn execute_sync_via_remote_execution_runtime(
candidate_started_unix_secs: u64, candidate_started_unix_secs: u64,
candidate_started_at: Instant, candidate_started_at: Instant,
) -> Result<RemoteSyncFallbackOutcome, GatewayError> { ) -> Result<RemoteSyncFallbackOutcome, GatewayError> {
let remote_request_started_at_unix_ms = current_request_candidate_unix_ms();
let remote_request_order_id = uuid::Uuid::now_v7().to_string();
let response = match post_sync_plan_to_remote_execution_runtime( let response = match post_sync_plan_to_remote_execution_runtime(
state, state,
remote_execution_runtime_base_url, remote_execution_runtime_base_url,
@@ -3299,11 +3374,19 @@ async fn execute_sync_via_remote_execution_runtime(
)); ));
} }
response let remote_response_observed_at_unix_ms = current_request_candidate_unix_ms();
.json() let mut result = response
.json::<ExecutionResult>()
.await .await
.map(RemoteSyncFallbackOutcome::Executed) .map_err(|err| GatewayError::Internal(err.to_string()))?;
.map_err(|err| GatewayError::Internal(err.to_string())) result
.response_observation
.get_or_insert(ExecutionResponseObservation {
request_started_at_unix_ms: remote_request_started_at_unix_ms,
response_headers_observed_at_unix_ms: remote_response_observed_at_unix_ms,
request_order_id: remote_request_order_id,
});
Ok(RemoteSyncFallbackOutcome::Executed(result))
} }
#[cfg(test)] #[cfg(test)]
@@ -9,12 +9,12 @@ use std::sync::{Arc, LazyLock, Mutex as StdMutex, OnceLock, RwLock as StdRwLock}
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use aether_contracts::{ use aether_contracts::{
ExecutionPlan, ExecutionResponseBodyMode, ExecutionResult, ExecutionTelemetry, ProxySnapshot, ExecutionPlan, ExecutionResponseBodyMode, ExecutionResponseObservation, ExecutionResult,
ResolvedTransportProfile, ResponseBody, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, ExecutionTelemetry, ProxySnapshot, ResolvedTransportProfile, ResponseBody,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
EXECUTION_RESPONSE_BODY_MODE_HEADER, TRANSPORT_BACKEND_BROWSER_WREQ, EXECUTION_REQUEST_HTTP1_ONLY_HEADER, EXECUTION_RESPONSE_BODY_MODE_HEADER,
TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_REQWEST_RUSTLS,
TRANSPORT_HTTP_MODE_HTTP1_ONLY, TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY,
}; };
use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation; use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation;
use aether_http::{apply_http_client_config, HttpClientConfig}; use aether_http::{apply_http_client_config, HttpClientConfig};
@@ -691,6 +691,7 @@ pub(crate) struct DirectUpstreamStreamExecution {
pub(crate) stream_precommit_committed: bool, pub(crate) stream_precommit_committed: bool,
pub(crate) response: DirectUpstreamResponse, pub(crate) response: DirectUpstreamResponse,
pub(crate) started_at: Instant, pub(crate) started_at: Instant,
pub(crate) response_observation: ExecutionResponseObservation,
pub(crate) stream_first_byte_timeout: Option<Duration>, pub(crate) stream_first_byte_timeout: Option<Duration>,
pub(crate) upstream_target_permit: Option<UpstreamTargetAdmissionPermit>, pub(crate) upstream_target_permit: Option<UpstreamTargetAdmissionPermit>,
} }
@@ -699,6 +700,7 @@ pub(crate) struct DirectUpstreamStreamExecution {
pub(crate) struct DirectSyncResponseStarted { pub(crate) struct DirectSyncResponseStarted {
pub(crate) status_code: u16, pub(crate) status_code: u16,
pub(crate) ttfb_ms: u64, pub(crate) ttfb_ms: u64,
pub(crate) response_observation: ExecutionResponseObservation,
} }
impl DirectSyncExecutionRuntime { impl DirectSyncExecutionRuntime {
@@ -724,14 +726,23 @@ impl DirectSyncExecutionRuntime {
let body_bytes = build_request_body(plan)?; let body_bytes = build_request_body(plan)?;
let started_at = Instant::now(); let started_at = Instant::now();
let request_started_at_unix_ms = crate::clock::current_unix_ms();
let request_order_id = uuid::Uuid::now_v7().to_string();
with_non_stream_total_timeout(plan, async move { with_non_stream_total_timeout(plan, async move {
let response = send_request_inner(plan, body_bytes, false).await?; let response = send_request_inner(plan, body_bytes, false).await?;
let ttfb_ms = started_at.elapsed().as_millis() as u64; let ttfb_ms = started_at.elapsed().as_millis() as u64;
let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms();
let status_code = response.status_code(); let status_code = response.status_code();
let headers = response.headers(); let headers = response.headers();
let response_observation = ExecutionResponseObservation {
request_started_at_unix_ms,
response_headers_observed_at_unix_ms,
request_order_id,
};
on_response_started(DirectSyncResponseStarted { on_response_started(DirectSyncResponseStarted {
status_code, status_code,
ttfb_ms, ttfb_ms,
response_observation: response_observation.clone(),
}); });
let (body_bytes, stream_ttfb_ms) = let (body_bytes, stream_ttfb_ms) =
response.bytes_with_stream_timeout(plan, started_at).await?; response.bytes_with_stream_timeout(plan, started_at).await?;
@@ -752,6 +763,7 @@ impl DirectSyncExecutionRuntime {
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code, status_code,
headers, headers,
response_observation: Some(response_observation),
body, body,
telemetry: Some(ExecutionTelemetry { telemetry: Some(ExecutionTelemetry {
ttfb_ms: stream_ttfb_ms.or(Some(ttfb_ms)), ttfb_ms: stream_ttfb_ms.or(Some(ttfb_ms)),
@@ -776,6 +788,8 @@ impl DirectSyncExecutionRuntime {
); );
let started_at = Instant::now(); let started_at = Instant::now();
let request_started_at_unix_ms = crate::clock::current_unix_ms();
let request_order_id = uuid::Uuid::now_v7().to_string();
let response = send_request(plan, body_bytes).await?; let response = send_request(plan, body_bytes).await?;
observe_gateway_stage_ms( observe_gateway_stage_ms(
"direct_send_headers", "direct_send_headers",
@@ -783,6 +797,7 @@ impl DirectSyncExecutionRuntime {
); );
let status_code = response.status_code(); let status_code = response.status_code();
let headers = response.headers(); let headers = response.headers();
let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms();
let stream_summary_report_context = build_stream_summary_report_context(plan); let stream_summary_report_context = build_stream_summary_report_context(plan);
@@ -797,6 +812,11 @@ impl DirectSyncExecutionRuntime {
stream_precommit_committed: false, stream_precommit_committed: false,
response: response.into_direct_upstream_response(), response: response.into_direct_upstream_response(),
started_at, started_at,
response_observation: ExecutionResponseObservation {
request_started_at_unix_ms,
response_headers_observed_at_unix_ms,
request_order_id,
},
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan), stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
upstream_target_permit: None, upstream_target_permit: None,
}) })
@@ -834,7 +854,7 @@ pub(crate) async fn execute_sync_plan_with_report_context(
} }
if resolve_local_tunnel_node_id(state, plan.proxy.as_ref()).is_some() { if resolve_local_tunnel_node_id(state, plan.proxy.as_ref()).is_some() {
return execute_sync_plan_via_local_tunnel(state, plan) return execute_sync_plan_via_local_tunnel(state, plan, report_context)
.await .await
.map_err(|err| GatewayError::Internal(err.to_string())); .map_err(|err| GatewayError::Internal(err.to_string()));
} }
@@ -857,7 +877,24 @@ pub(crate) async fn execute_sync_plan_with_report_context(
Ok(None) => {} Ok(None) => {}
Err(err) => return Err(GatewayError::Internal(err.to_string())), Err(err) => return Err(GatewayError::Internal(err.to_string())),
} }
match DirectSyncExecutionRuntime::new().execute_sync(plan).await { let state_for_response_started = state.clone();
match DirectSyncExecutionRuntime::new()
.execute_sync_with_response_started(plan, move |event| {
crate::orchestration::spawn_local_oauth_success_effect(
state_for_response_started,
plan,
report_context,
crate::orchestration::LocalOAuthSuccessEffect {
status_code: event.status_code,
request_started_at_unix_ms: Some(
event.response_observation.request_started_at_unix_ms,
),
request_order_id: Some(&event.response_observation.request_order_id),
},
);
})
.await
{
Ok(result) => { Ok(result) => {
record_manual_proxy_request_outcome(state, plan, result.status_code).await; record_manual_proxy_request_outcome(state, plan, result.status_code).await;
Ok(result) Ok(result)
@@ -889,6 +926,8 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
plan.body.body_bytes_b64.is_some(), plan.body.body_bytes_b64.is_some(),
)?; )?;
let started_at = Instant::now(); let started_at = Instant::now();
let request_started_at_unix_ms = crate::clock::current_unix_ms();
let request_order_id = uuid::Uuid::now_v7().to_string();
let response = state let response = state
.tunnel .tunnel
.open_direct_relay_stream( .open_direct_relay_stream(
@@ -900,6 +939,7 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
.map_err(ExecutionRuntimeTransportError::RelayError)?; .map_err(ExecutionRuntimeTransportError::RelayError)?;
let status_code = response.status(); let status_code = response.status();
let headers = collect_tunnel_response_headers(response.headers()); let headers = collect_tunnel_response_headers(response.headers());
let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms();
Ok(Some(DirectUpstreamStreamExecution { Ok(Some(DirectUpstreamStreamExecution {
request_id: plan.request_id.clone(), request_id: plan.request_id.clone(),
@@ -912,6 +952,11 @@ pub(crate) async fn execute_stream_plan_via_local_tunnel(
stream_precommit_committed: false, stream_precommit_committed: false,
response: DirectUpstreamResponse::LocalTunnel(response), response: DirectUpstreamResponse::LocalTunnel(response),
started_at, started_at,
response_observation: ExecutionResponseObservation {
request_started_at_unix_ms,
response_headers_observed_at_unix_ms,
request_order_id,
},
stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan), stream_first_byte_timeout: resolve_stream_first_byte_timeout(plan),
upstream_target_permit: None, upstream_target_permit: None,
})) }))
@@ -991,13 +1036,19 @@ fn manual_proxy_node_id(proxy: Option<&ProxySnapshot>) -> Option<String> {
async fn execute_sync_plan_via_local_tunnel( async fn execute_sync_plan_via_local_tunnel(
state: &AppState, state: &AppState,
plan: &ExecutionPlan, plan: &ExecutionPlan,
report_context: Option<&serde_json::Value>,
) -> Result<ExecutionResult, ExecutionRuntimeTransportError> { ) -> Result<ExecutionResult, ExecutionRuntimeTransportError> {
with_non_stream_total_timeout(plan, execute_sync_plan_via_local_tunnel_inner(state, plan)).await with_non_stream_total_timeout(
plan,
execute_sync_plan_via_local_tunnel_inner(state, plan, report_context),
)
.await
} }
async fn execute_sync_plan_via_local_tunnel_inner( async fn execute_sync_plan_via_local_tunnel_inner(
state: &AppState, state: &AppState,
plan: &ExecutionPlan, plan: &ExecutionPlan,
report_context: Option<&serde_json::Value>,
) -> Result<ExecutionResult, ExecutionRuntimeTransportError> { ) -> Result<ExecutionResult, ExecutionRuntimeTransportError> {
let node_id = resolve_local_tunnel_node_id(state, plan.proxy.as_ref()).ok_or_else(|| { let node_id = resolve_local_tunnel_node_id(state, plan.proxy.as_ref()).ok_or_else(|| {
ExecutionRuntimeTransportError::RelayError("local tunnel node unavailable".to_string()) ExecutionRuntimeTransportError::RelayError("local tunnel node unavailable".to_string())
@@ -1030,6 +1081,8 @@ async fn execute_sync_plan_via_local_tunnel_inner(
"gateway execution runtime local tunnel request prepared" "gateway execution runtime local tunnel request prepared"
); );
let started_at = Instant::now(); let started_at = Instant::now();
let request_started_at_unix_ms = crate::clock::current_unix_ms();
let request_order_id = uuid::Uuid::now_v7().to_string();
let mut response = state let mut response = state
.tunnel .tunnel
.open_direct_relay_stream( .open_direct_relay_stream(
@@ -1040,8 +1093,24 @@ async fn execute_sync_plan_via_local_tunnel_inner(
.await .await
.map_err(ExecutionRuntimeTransportError::RelayError)?; .map_err(ExecutionRuntimeTransportError::RelayError)?;
let ttfb_ms = started_at.elapsed().as_millis() as u64; let ttfb_ms = started_at.elapsed().as_millis() as u64;
let response_headers_observed_at_unix_ms = crate::clock::current_unix_ms();
let status_code = response.status(); let status_code = response.status();
let headers = collect_tunnel_response_headers(response.headers()); let headers = collect_tunnel_response_headers(response.headers());
let response_observation = ExecutionResponseObservation {
request_started_at_unix_ms,
response_headers_observed_at_unix_ms,
request_order_id,
};
crate::orchestration::spawn_local_oauth_success_effect(
state.clone(),
plan,
report_context,
crate::orchestration::LocalOAuthSuccessEffect {
status_code,
request_started_at_unix_ms: Some(response_observation.request_started_at_unix_ms),
request_order_id: Some(&response_observation.request_order_id),
},
);
let proxy_timing = execution_header_for_log(&headers, "x-proxy-timing").unwrap_or("-"); let proxy_timing = execution_header_for_log(&headers, "x-proxy-timing").unwrap_or("-");
let (body_bytes, stream_ttfb_ms) = let (body_bytes, stream_ttfb_ms) =
collect_local_tunnel_response_body(response, plan, started_at).await?; collect_local_tunnel_response_body(response, plan, started_at).await?;
@@ -1095,6 +1164,7 @@ async fn execute_sync_plan_via_local_tunnel_inner(
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code, status_code,
headers, headers,
response_observation: Some(response_observation),
body, body,
telemetry: Some(ExecutionTelemetry { telemetry: Some(ExecutionTelemetry {
ttfb_ms: stream_ttfb_ms.or(Some(ttfb_ms)), ttfb_ms: stream_ttfb_ms.or(Some(ttfb_ms)),
@@ -5605,6 +5675,8 @@ mod tests {
) )
.await .await
.expect("headers should write"); .expect("headers should write");
socket.flush().await.expect("headers should flush");
tokio::time::sleep(std::time::Duration::from_millis(40)).await;
socket socket
.write_all(b"b\r\ndata: one\n\n\r\n") .write_all(b"b\r\ndata: one\n\n\r\n")
.await .await
@@ -5634,12 +5706,34 @@ mod tests {
let body = result let body = result
.body .body
.clone()
.and_then(|body| body.body_bytes_b64) .and_then(|body| body.body_bytes_b64)
.and_then(|body| base64::engine::general_purpose::STANDARD.decode(body).ok()) .and_then(|body| base64::engine::general_purpose::STANDARD.decode(body).ok())
.expect("stream body should be captured as bytes"); .expect("stream body should be captured as bytes");
let body = String::from_utf8(body).expect("stream body should be utf8"); let body = String::from_utf8(body).expect("stream body should be utf8");
assert!(body.contains("data: one")); assert!(body.contains("data: one"));
assert!(body.contains("data: two")); assert!(body.contains("data: two"));
let observation = result
.response_observation
.expect("stream sync execution should preserve header observation");
let telemetry = result
.telemetry
.expect("stream sync execution should include telemetry");
let ttfb_ms = telemetry
.ttfb_ms
.expect("stream sync execution should measure the first body byte");
assert!(
observation.response_headers_observed_at_unix_ms
>= observation.request_started_at_unix_ms
);
assert!(
observation
.response_headers_observed_at_unix_ms
.saturating_sub(observation.request_started_at_unix_ms)
< ttfb_ms,
"header observation must not be derived from body-byte ttfb"
);
assert!(!observation.request_order_id.is_empty());
} }
#[tokio::test] #[tokio::test]
@@ -266,6 +266,7 @@ pub(crate) async fn maybe_execute_windsurf_sync(
candidate_id: prepared.candidate_id, candidate_id: prepared.candidate_id,
status_code: 200, status_code: 200,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]), headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: Some(body_json), json_body: Some(body_json),
body_bytes_b64: None, body_bytes_b64: None,
@@ -527,6 +528,7 @@ fn build_windsurf_stream_frame_stream(
("cache-control".to_string(), "no-cache".to_string()), ("cache-control".to_string(), "no-cache".to_string()),
("content-type".to_string(), "text/event-stream".to_string()), ("content-type".to_string(), "text/event-stream".to_string()),
]), ]),
response_observation: None,
}, },
}); });
@@ -1687,6 +1687,7 @@ mod tests {
CONTENT_TYPE.as_str().to_string(), CONTENT_TYPE.as_str().to_string(),
"application/json".to_string(), "application/json".to_string(),
)]), )]),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(body_json), json_body: Some(body_json),
body_bytes_b64: None, body_bytes_b64: None,
@@ -55,6 +55,33 @@ pub(super) async fn maybe_handle(
if idempotency_key.is_empty() { if idempotency_key.is_empty() {
return Ok(Some(bad_request_response("idempotency_key 不能为空"))); return Ok(Some(bad_request_response("idempotency_key 不能为空")));
} }
if idempotency_key.len() > 256 {
return Ok(Some(bad_request_response(
"idempotency_key 不能超过 256 个字节",
)));
}
let expected_credential_generation = match payload.expected_credential_generation {
serde_json::Value::Null => None,
serde_json::Value::String(value) => {
let value = value.trim().to_string();
if value.is_empty() {
return Ok(Some(bad_request_response(
"expected_credential_generation 不能为空字符串",
)));
}
if value.len() > 256 {
return Ok(Some(bad_request_response(
"expected_credential_generation 不能超过 256 个字节",
)));
}
Some(value)
}
_ => {
return Ok(Some(bad_request_response(
"expected_credential_generation 必须是字符串或 null",
)));
}
};
let Some(key) = state let Some(key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id)) .read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
@@ -93,9 +120,15 @@ pub(super) async fn maybe_handle(
))); )));
}; };
let (status, payload) = let (status, payload) = consume_codex_reset_credit_locally(
consume_codex_reset_credit_locally(state, &provider, &endpoint, key, &idempotency_key) state,
.await?; &provider,
&endpoint,
key,
&idempotency_key,
expected_credential_generation.as_deref(),
)
.await?;
Ok(Some((status, Json(payload)).into_response())) Ok(Some((status, Json(payload)).into_response()))
} }
@@ -1,7 +1,10 @@
use crate::handlers::admin::admin_provider_pool_config; use crate::handlers::admin::admin_provider_pool_config;
use crate::handlers::admin::provider::shared::paths::admin_update_key_id; use crate::handlers::admin::provider::shared::paths::admin_update_key_id;
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdatePatch; use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdatePatch;
use crate::handlers::admin::provider::write::keys::admin_provider_key_update_requires_immediate_model_fetch; use crate::handlers::admin::provider::write::keys::{
admin_provider_key_update_requires_immediate_model_fetch,
build_provider_catalog_key_admin_cas_update,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::maintenance::ensure_provider_key_pool_scores_for_keys; use crate::maintenance::ensure_provider_key_pool_scores_for_keys;
use crate::provider_key_auth::provider_key_effective_api_formats; use crate::provider_key_auth::provider_key_effective_api_formats;
@@ -82,7 +85,25 @@ pub(super) async fn maybe_handle(
Ok(record) => record, Ok(record) => record,
Err(detail) => return Ok(Some(bad_request_response(detail))), Err(detail) => return Ok(Some(bad_request_response(detail))),
}; };
let Some(mut updated) = state.update_provider_catalog_key(&updated_record).await? else { let admin_update = build_provider_catalog_key_admin_cas_update(
&existing_key,
updated_record.clone(),
&provider.provider_type,
);
if !state
.compare_and_update_provider_catalog_key_admin_state(&admin_update)
.await?
{
return Ok(Some(conflict_response(
"Key 凭据或配置已被其他请求更新,请刷新后重试",
)));
}
let Some(mut updated) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
else {
return Ok(None); return Ok(None);
}; };
if updated_record.learned_rpm_limit != existing_key.learned_rpm_limit { if updated_record.learned_rpm_limit != existing_key.learned_rpm_limit {
@@ -183,3 +204,11 @@ fn not_found_response(detail: impl Into<String>) -> Response<Body> {
) )
.into_response() .into_response()
} }
fn conflict_response(detail: impl Into<String>) -> Response<Body> {
(
http::StatusCode::CONFLICT,
Json(json!({ "detail": detail.into() })),
)
.into_response()
}
@@ -22,16 +22,60 @@ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::shared::sync_provider_key_oauth_status_snapshot; use crate::handlers::shared::sync_provider_key_oauth_status_snapshot;
use crate::provider_key_auth::provider_key_is_oauth_managed; use crate::provider_key_auth::provider_key_is_oauth_managed;
use crate::GatewayError; use crate::GatewayError;
use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyOAuthRuntimeStateCasUpdate; use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogUpstreamMetadataNamespaceExpectation,
};
use axum::{ use axum::{
body::{Body, Bytes}, body::{Body, Bytes},
http, http,
response::{IntoResponse, Response}, response::{IntoResponse, Response},
Json, Json,
}; };
use serde_json::json; use serde_json::{json, Value};
use std::time::{SystemTime, UNIX_EPOCH}; use std::time::{SystemTime, UNIX_EPOCH};
const CODEX_OAUTH_COMPLETE_NAMESPACE_CAS_MAX_RETRIES: usize = 3;
const CODEX_CREDENTIAL_GENERATION_KEY: &str = "credential_generation";
#[derive(Debug, PartialEq)]
enum CodexOAuthCompleteCasMissAction {
AlreadyCompleted,
RetryNamespace(Option<Value>),
Conflict,
}
fn codex_oauth_complete_cas_miss_action(
latest_encrypted_auth_config: Option<&str>,
latest_upstream_metadata: Option<&Value>,
latest_status_snapshot: Option<&Value>,
expected_encrypted_auth_config: Option<&str>,
persisted_encrypted_auth_config: &str,
expected_codex_metadata_value: Option<&Value>,
replacement_codex_metadata_value: &Value,
) -> CodexOAuthCompleteCasMissAction {
let latest_codex_metadata_value = latest_upstream_metadata
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.cloned();
let quota_is_cleared = latest_status_snapshot
.and_then(Value::as_object)
.and_then(|snapshot| snapshot.get("quota"))
== Some(&Value::Null);
if latest_encrypted_auth_config == Some(persisted_encrypted_auth_config)
&& latest_codex_metadata_value.as_ref() == Some(replacement_codex_metadata_value)
&& quota_is_cleared
{
return CodexOAuthCompleteCasMissAction::AlreadyCompleted;
}
if latest_encrypted_auth_config != expected_encrypted_auth_config
|| latest_codex_metadata_value.as_ref() == expected_codex_metadata_value
{
return CodexOAuthCompleteCasMissAction::Conflict;
}
CodexOAuthCompleteCasMissAction::RetryNamespace(latest_codex_metadata_value)
}
pub(super) async fn handle_admin_provider_oauth_complete_key( pub(super) async fn handle_admin_provider_oauth_complete_key(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>, request_context: &AdminRequestContext<'_>,
@@ -277,29 +321,98 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
.and_then(|snapshot| snapshot.get("oauth")) .and_then(|snapshot| snapshot.get("oauth"))
.cloned() .cloned()
.unwrap_or(serde_json::Value::Null); .unwrap_or(serde_json::Value::Null);
let mut expected_codex_metadata_value = key
.upstream_metadata
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.cloned();
let mut status_snapshot_patch =
serde_json::Map::from_iter([("oauth".to_string(), oauth_status)]);
if provider_type == "codex" {
status_snapshot_patch.insert("quota".to_string(), serde_json::Value::Null);
}
let persisted_encrypted_auth_config = recovered_key let persisted_encrypted_auth_config = recovered_key
.encrypted_auth_config .encrypted_auth_config
.clone() .clone()
.expect("recovered auth config should be present"); .expect("recovered auth config should be present");
let updated_result = state let replacement_codex_metadata_value = json!({
.app() CODEX_CREDENTIAL_GENERATION_KEY: uuid::Uuid::now_v7().to_string()
.compare_and_update_provider_catalog_key_oauth_runtime_state( });
&ProviderCatalogKeyOAuthRuntimeStateCasUpdate { let expected_encrypted_auth_config = state_data.expected_encrypted_auth_config.clone();
key_id: key_id.clone(), let updated_result: Result<bool, GatewayError> = async {
expected_encrypted_auth_config: state_data.expected_encrypted_auth_config, let max_namespace_retries = if provider_type == "codex" {
expected_credential: None, CODEX_OAUTH_COMPLETE_NAMESPACE_CAS_MAX_RETRIES
encrypted_auth_config: persisted_encrypted_auth_config.clone(), } else {
encrypted_api_key_update: Some(encrypted_api_key), 0
expires_at_unix_secs_update: Some(expires_at), };
oauth_invalid_at_unix_secs: None, for retry in 0..=max_namespace_retries {
oauth_invalid_reason: None, let updated = state
reset_error_count: true, .app()
upstream_metadata_patch: None, .compare_and_update_provider_catalog_key_oauth_runtime_state(
status_snapshot_patch: json!({ "oauth": oauth_status }), &ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
updated_at_unix_secs: Some(now_unix_secs), key_id: key_id.clone(),
}, expected_encrypted_auth_config: expected_encrypted_auth_config.clone(),
) expected_credential: None,
.await; expected_upstream_metadata_namespace: (provider_type == "codex").then(
|| ProviderCatalogUpstreamMetadataNamespaceExpectation {
namespace: "codex".to_string(),
expected_value: expected_codex_metadata_value.clone(),
},
),
encrypted_auth_config: persisted_encrypted_auth_config.clone(),
encrypted_api_key_update: Some(encrypted_api_key.clone()),
expires_at_unix_secs_update: Some(expires_at),
oauth_invalid_at_unix_secs: None,
oauth_invalid_reason: None,
reset_error_count: true,
upstream_metadata_patch: (provider_type == "codex")
.then(|| json!({"codex": replacement_codex_metadata_value.clone()})),
upstream_metadata_namespace_to_remove: None,
status_snapshot_patch: serde_json::Value::Object(
status_snapshot_patch.clone(),
),
updated_at_unix_secs: Some(now_unix_secs),
},
)
.await?;
if updated {
return Ok(true);
}
if provider_type != "codex" {
return Ok(false);
}
let Some(latest_key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
else {
return Ok(false);
};
match codex_oauth_complete_cas_miss_action(
latest_key.encrypted_auth_config.as_deref(),
latest_key.upstream_metadata.as_ref(),
latest_key.status_snapshot.as_ref(),
expected_encrypted_auth_config.as_deref(),
&persisted_encrypted_auth_config,
expected_codex_metadata_value.as_ref(),
&replacement_codex_metadata_value,
) {
CodexOAuthCompleteCasMissAction::AlreadyCompleted => return Ok(true),
CodexOAuthCompleteCasMissAction::Conflict => return Ok(false),
CodexOAuthCompleteCasMissAction::RetryNamespace(latest_codex_metadata_value) => {
if retry == max_namespace_retries {
return Ok(false);
}
expected_codex_metadata_value = latest_codex_metadata_value;
}
}
}
Ok(false)
}
.await;
let _ = state let _ = state
.app() .app()
.invalidate_local_oauth_refresh_entry(&key_id) .invalidate_local_oauth_refresh_entry(&key_id)
@@ -397,3 +510,78 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
})) }))
.into_response()) .into_response())
} }
#[cfg(test)]
mod tests {
use super::{codex_oauth_complete_cas_miss_action, CodexOAuthCompleteCasMissAction};
use serde_json::json;
#[test]
fn codex_oauth_complete_retries_only_when_namespace_changed() {
let expected_codex = json!({"request_id": "old"});
let replacement_codex = json!({"credential_generation": "generation-new"});
let latest_metadata = json!({
"codex": {"request_id": "new"},
"unrelated": {"preserved": true}
});
assert_eq!(
codex_oauth_complete_cas_miss_action(
Some("old-auth"),
Some(&latest_metadata),
None,
Some("old-auth"),
"new-auth",
Some(&expected_codex),
&replacement_codex,
),
CodexOAuthCompleteCasMissAction::RetryNamespace(Some(json!({
"request_id": "new"
})))
);
assert_eq!(
codex_oauth_complete_cas_miss_action(
Some("old-auth"),
Some(&json!({"codex": expected_codex.clone()})),
None,
Some("old-auth"),
"new-auth",
Some(&expected_codex),
&replacement_codex,
),
CodexOAuthCompleteCasMissAction::Conflict
);
}
#[test]
fn codex_oauth_complete_accepts_an_ambiguous_success_but_rejects_auth_rotation() {
let replacement_codex = json!({"credential_generation": "generation-new"});
assert_eq!(
codex_oauth_complete_cas_miss_action(
Some("new-auth"),
Some(&json!({
"codex": replacement_codex.clone(),
"unrelated": {"preserved": true}
})),
Some(&json!({"quota": null})),
Some("old-auth"),
"new-auth",
Some(&json!({"request_id": "old"})),
&replacement_codex,
),
CodexOAuthCompleteCasMissAction::AlreadyCompleted
);
assert_eq!(
codex_oauth_complete_cas_miss_action(
Some("other-auth"),
Some(&json!({"codex": {"request_id": "new"}})),
Some(&json!({"quota": null})),
Some("old-auth"),
"new-auth",
Some(&json!({"request_id": "old"})),
&replacement_codex,
),
CodexOAuthCompleteCasMissAction::Conflict
);
}
}
@@ -12,6 +12,7 @@ use crate::ai_serving::{
build_provider_key_pool_score_upsert, provider_key_pool_score_id, provider_key_pool_score_scope, build_provider_key_pool_score_upsert, provider_key_pool_score_id, provider_key_pool_score_scope,
}; };
use crate::handlers::admin::admin_provider_pool_config; use crate::handlers::admin::admin_provider_pool_config;
use crate::handlers::admin::provider::write::keys::build_provider_catalog_key_admin_cas_update;
use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::request::AdminAppState;
use crate::provider_key_auth::provider_active_api_formats; use crate::provider_key_auth::provider_active_api_formats;
use crate::GatewayError; use crate::GatewayError;
@@ -286,6 +287,61 @@ fn grok_oauth_catalog_key_fingerprint(
grok_browser_transport_fingerprint_from_auth_config(auth_config) grok_browser_transport_fingerprint_from_auth_config(auth_config)
} }
pub(crate) fn rotate_codex_credential_generation(
key: &mut StoredProviderCatalogKey,
provider_type: &str,
) {
if !provider_type.trim().eq_ignore_ascii_case("codex") {
return;
}
let mut upstream_metadata = key
.upstream_metadata
.as_ref()
.and_then(Value::as_object)
.cloned()
.unwrap_or_default();
upstream_metadata.insert(
"codex".to_string(),
json!({
aether_admin::provider::quota::CODEX_CREDENTIAL_GENERATION_KEY:
Uuid::now_v7().to_string(),
}),
);
key.upstream_metadata = Some(Value::Object(upstream_metadata));
if let Some(mut status_snapshot) = key
.status_snapshot
.as_ref()
.and_then(Value::as_object)
.cloned()
{
status_snapshot.insert("quota".to_string(), Value::Null);
key.status_snapshot = Some(Value::Object(status_snapshot));
}
}
pub(crate) fn ensure_codex_credential_generation_rotated(
key: &mut StoredProviderCatalogKey,
provider_type: &str,
previous_generation: Option<&str>,
) {
if !provider_type.trim().eq_ignore_ascii_case("codex") {
return;
}
let current_generation = key
.upstream_metadata
.as_ref()
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.and_then(|codex| aether_admin::provider::quota::codex_credential_generation(Some(codex)));
let already_rotated = current_generation.is_some() && current_generation != previous_generation;
if !already_rotated {
rotate_codex_credential_generation(key, provider_type);
}
}
pub(crate) async fn create_provider_oauth_catalog_key( pub(crate) async fn create_provider_oauth_catalog_key(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
provider_id: &str, provider_id: &str,
@@ -344,6 +400,7 @@ pub(crate) async fn create_provider_oauth_catalog_key(
record.circuit_breaker_by_format = Some(json!({})); record.circuit_breaker_by_format = Some(json!({}));
record.created_at_unix_ms = Some(now_unix_secs); record.created_at_unix_ms = Some(now_unix_secs);
record.updated_at_unix_secs = Some(now_unix_secs); record.updated_at_unix_secs = Some(now_unix_secs);
rotate_codex_credential_generation(&mut record, provider_type);
let created = state.create_provider_catalog_key(&record).await?; let created = state.create_provider_catalog_key(&record).await?;
if let Some(key) = created.as_ref() { if let Some(key) = created.as_ref() {
let _ = state let _ = state
@@ -395,17 +452,23 @@ pub(crate) async fn update_existing_provider_oauth_catalog_key(
updated.proxy = Some(proxy); updated.proxy = Some(proxy);
} }
updated.updated_at_unix_secs = Some(now_unix_secs); updated.updated_at_unix_secs = Some(now_unix_secs);
if state.update_provider_catalog_key(&updated).await?.is_none() { rotate_codex_credential_generation(&mut updated, provider_type);
return Ok(None); let admin_update =
} build_provider_catalog_key_admin_cas_update(existing_key, updated.clone(), provider_type);
if !state if !state
.clear_provider_catalog_key_oauth_invalid_marker(&updated.id) .compare_and_update_provider_catalog_key_admin_state(&admin_update)
.await? .await?
{ {
return Ok(None); return Ok(None);
} }
let persisted = state let persisted = state
.reset_provider_catalog_key_recovery_state(&updated.id) .reset_provider_catalog_key_recovery_state_fenced(
&updated.id,
updated
.encrypted_auth_config
.as_deref()
.expect("OAuth update always supplies encrypted auth_config"),
)
.await?; .await?;
if let Some(key) = persisted.as_ref() { if let Some(key) = persisted.as_ref() {
let _ = state let _ = state
@@ -502,10 +565,12 @@ fn provider_oauth_catalog_key_api_formats(
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::{ use super::{
grok_oauth_catalog_key_fingerprint, provider_oauth_token_payload_expires_at_unix_secs, ensure_codex_credential_generation_rotated, grok_oauth_catalog_key_fingerprint,
provider_oauth_token_payload_expires_at_unix_secs, rotate_codex_credential_generation,
}; };
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _}; use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use serde_json::json; use serde_json::{json, Value};
fn sample_unsigned_jwt(payload: serde_json::Value) -> String { fn sample_unsigned_jwt(payload: serde_json::Value) -> String {
let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"none","typ":"JWT"}"#); let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"none","typ":"JWT"}"#);
@@ -609,4 +674,101 @@ mod tests {
assert!(grok_oauth_catalog_key_fingerprint("openai", auth_config).is_none()); assert!(grok_oauth_catalog_key_fingerprint("openai", auth_config).is_none());
} }
#[test]
fn codex_credential_rotation_replaces_quota_namespace_and_preserves_unrelated_state() {
let mut key = StoredProviderCatalogKey::new(
"key".to_string(),
"provider".to_string(),
"Codex".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build");
key.upstream_metadata = Some(json!({
"codex": {
"credential_generation": "old-generation",
"primary_used_percent": 75.0,
},
"unrelated": {"preserved": true},
}));
key.status_snapshot = Some(json!({
"oauth": {"status": "valid"},
"quota": {"used_ratio": 0.75},
}));
rotate_codex_credential_generation(&mut key, "codex");
let codex = key
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.get("codex"))
.and_then(Value::as_object)
.expect("codex namespace should exist");
assert_eq!(codex.len(), 1);
assert_ne!(
codex
.get(aether_admin::provider::quota::CODEX_CREDENTIAL_GENERATION_KEY)
.and_then(Value::as_str),
Some("old-generation")
);
assert_eq!(
key.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/unrelated/preserved")),
Some(&json!(true))
);
assert_eq!(
key.status_snapshot
.as_ref()
.and_then(|snapshot| snapshot.get("quota")),
Some(&Value::Null)
);
assert_eq!(
key.status_snapshot
.as_ref()
.and_then(|snapshot| snapshot.pointer("/oauth/status")),
Some(&json!("valid"))
);
}
#[test]
fn codex_credential_rotation_ensure_does_not_rotate_twice_in_one_write() {
let mut key = StoredProviderCatalogKey::new(
"key".to_string(),
"provider".to_string(),
"Codex".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build");
key.upstream_metadata = Some(json!({
"codex": {"credential_generation": "generation-before-write"}
}));
rotate_codex_credential_generation(&mut key, "codex");
let builder_generation = key
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/codex/credential_generation"))
.and_then(Value::as_str)
.expect("builder should rotate the generation")
.to_string();
ensure_codex_credential_generation_rotated(
&mut key,
"codex",
Some("generation-before-write"),
);
assert_eq!(
key.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/codex/credential_generation"))
.and_then(Value::as_str),
Some(builder_generation.as_str())
);
}
} }
@@ -470,6 +470,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 403, status_code: 403,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: None, json_body: None,
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)), body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
@@ -18,24 +18,229 @@ use self::plan::{
execute_codex_reset_credit_plan, execute_codex_reset_credit_plan,
}; };
use super::shared::{ use super::shared::{
build_quota_snapshot_payload, extract_execution_error_message, build_quota_snapshot_payload, complete_codex_account_reset, extract_execution_error_message,
oauth_refresh_auto_removed_result, persist_fenced_provider_quota_refresh_state, oauth_refresh_auto_removed_result, persist_codex_provider_quota_refresh_state,
persist_provider_quota_refresh_state, provider_auto_remove_banned_keys, persist_fenced_provider_quota_refresh_state, provider_auto_remove_banned_keys,
provider_auto_remove_quota_exhausted_keys, quota_key_auto_removed, provider_auto_remove_quota_exhausted_keys, quota_key_auto_removed,
quota_refresh_success_invalid_state, should_auto_remove_oauth_invalid_key, quota_refresh_success_invalid_state, reserve_codex_account_reset,
ProviderQuotaExecutionOutcome, should_auto_remove_oauth_invalid_key, CodexAccountResetCompleteResult,
CodexAccountResetReserveResult, CodexAccountResetTerminal, ProviderQuotaExecutionOutcome,
}; };
use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::provider_key_auth::provider_key_is_oauth_managed; use crate::provider_key_auth::provider_key_is_oauth_managed;
use crate::state::ProviderTransportCredentialFence;
use crate::GatewayError; use crate::GatewayError;
use aether_contracts::ProxySnapshot; use aether_contracts::ProxySnapshot;
use aether_data_contracts::repository::provider_catalog::{ use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthCredentialFence, ProviderCatalogKeyOAuthCredentialCasDelete,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, ProviderCatalogUpstreamMetadataNamespaceExpectation, StoredProviderCatalogEndpoint,
StoredProviderCatalogKey, StoredProviderCatalogProvider,
}; };
use axum::http::StatusCode; use axum::http::StatusCode;
use serde_json::{json, Map, Value}; use serde_json::{json, Map, Value};
use std::time::{SystemTime, UNIX_EPOCH};
const CODEX_OAUTH_CREDENTIAL_STABILIZATION_ATTEMPTS: usize = 3;
const CODEX_RESET_QUOTA_RECONCILIATION_DELAYS_MS: [u64; 4] = [1_000, 2_000, 4_000, 8_000];
enum CodexOAuthRequestPreparation {
Ready {
transport: AdminGatewayProviderTransportSnapshot,
auth: (String, String),
credential_fence: ProviderTransportCredentialFence,
},
MissingAuth,
Conflict,
}
async fn prepare_codex_oauth_request(
state: &AdminAppState<'_>,
initial_transport: &AdminGatewayProviderTransportSnapshot,
) -> Result<CodexOAuthRequestPreparation, GatewayError> {
for _ in 0..CODEX_OAUTH_CREDENTIAL_STABILIZATION_ATTEMPTS {
let Some(transport) = state
.read_provider_transport_snapshot_uncached(
&initial_transport.provider.id,
&initial_transport.endpoint.id,
&initial_transport.key.id,
)
.await?
else {
return Ok(CodexOAuthRequestPreparation::Conflict);
};
if !crate::state::provider_transport_context_allows_credential_rotation(
initial_transport,
&transport,
) {
return Ok(CodexOAuthRequestPreparation::Conflict);
}
let Some(before_fence) = state
.app()
.capture_provider_transport_credential_fence(&transport)
.await?
else {
continue;
};
let resolved_auth = state.resolve_local_oauth_header_auth(&transport).await?;
let Some(current_transport) = state
.read_provider_transport_snapshot_uncached(
&initial_transport.provider.id,
&initial_transport.endpoint.id,
&initial_transport.key.id,
)
.await?
else {
return Ok(CodexOAuthRequestPreparation::Conflict);
};
if !crate::state::provider_transport_context_allows_credential_rotation(
initial_transport,
&current_transport,
) {
return Ok(CodexOAuthRequestPreparation::Conflict);
}
let Some(after_fence) = state
.app()
.capture_provider_transport_credential_fence(&current_transport)
.await?
else {
continue;
};
if before_fence != after_fence {
continue;
}
return Ok(match resolved_auth {
Some(auth) => CodexOAuthRequestPreparation::Ready {
transport: current_transport,
auth,
credential_fence: after_fence,
},
None => CodexOAuthRequestPreparation::MissingAuth,
});
}
Ok(CodexOAuthRequestPreparation::Conflict)
}
fn codex_reset_refresh_succeeded(payload: Option<&Value>, key_id: &str) -> bool {
payload
.and_then(|payload| payload.get("results"))
.and_then(Value::as_array)
.into_iter()
.flatten()
.filter_map(Value::as_object)
.find(|item| item.get("key_id").and_then(Value::as_str) == Some(key_id))
.and_then(|item| item.get("status"))
.and_then(Value::as_str)
.is_some_and(|status| status.eq_ignore_ascii_case("success"))
}
async fn codex_reset_fence_is_still_pending(
state: &AdminAppState<'_>,
key_id: &str,
expected_credential: &ProviderTransportCredentialFence,
reset_fence: &super::shared::CodexAccountResetFence,
) -> Result<bool, GatewayError> {
let Some(key) = state
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
.await?
.into_iter()
.next()
else {
return Ok(false);
};
if key.encrypted_auth_config.as_deref()
!= Some(expected_credential.encrypted_auth_config.as_str())
|| key.encrypted_api_key != expected_credential.credential.encrypted_api_key
|| key.auth_type != expected_credential.credential.auth_type
|| key.provider_id != expected_credential.credential.provider_id
{
return Ok(false);
}
let provider_type_matches = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&key.provider_id))
.await?
.into_iter()
.next()
.is_some_and(|provider| {
provider.provider_type == expected_credential.credential.provider_type
});
if !provider_type_matches {
return Ok(false);
}
let codex = key
.upstream_metadata
.as_ref()
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.and_then(Value::as_object);
Ok(codex.is_some_and(|codex| {
codex
.get(aether_admin::provider::quota::CODEX_QUOTA_ACCOUNT_RESET_FENCE_ID_KEY)
.and_then(Value::as_str)
== Some(reset_fence.id.as_str())
&& codex
.get(aether_admin::provider::quota::CODEX_QUOTA_ACCOUNT_RESET_GENERATION_KEY)
.and_then(aether_admin::provider::quota::coerce_json_u64)
== Some(reset_fence.generation)
&& codex
.get(
aether_admin::provider::quota::CODEX_QUOTA_ACCOUNT_RESET_PENDING_GENERATION_KEY,
)
.and_then(aether_admin::provider::quota::coerce_json_u64)
== Some(reset_fence.generation)
&& codex
.get(aether_admin::provider::quota::CODEX_QUOTA_ACCOUNT_RESET_PENDING_KEY)
.and_then(Value::as_bool)
== Some(true)
}))
}
async fn refresh_codex_quota_after_reset_until_settled(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
key: &StoredProviderCatalogKey,
reset_fence: &super::shared::CodexAccountResetFence,
expected_credential: &ProviderTransportCredentialFence,
) -> Result<Option<Value>, GatewayError> {
let mut latest_payload = None;
for attempt in 0..=CODEX_RESET_QUOTA_RECONCILIATION_DELAYS_MS.len() {
if attempt > 0 {
tokio::time::sleep(std::time::Duration::from_millis(
CODEX_RESET_QUOTA_RECONCILIATION_DELAYS_MS[attempt - 1],
))
.await;
}
if !codex_reset_fence_is_still_pending(state, &key.id, expected_credential, reset_fence)
.await?
{
break;
}
let payload = refresh_codex_provider_quota_locally_with_reset_fence(
state,
provider,
endpoint,
vec![key.clone()],
None,
Some(reset_fence.id.as_str()),
Some(reset_fence.generation),
Some(expected_credential),
)
.await?;
let refresh_succeeded = codex_reset_refresh_succeeded(payload.as_ref(), &key.id);
latest_payload = payload;
if !refresh_succeeded
|| !codex_reset_fence_is_still_pending(state, &key.id, expected_credential, reset_fence)
.await?
{
break;
}
}
Ok(latest_payload)
}
fn merge_codex_quota_metadata( fn merge_codex_quota_metadata(
header_metadata: Option<&serde_json::Value>, header_metadata: Option<&serde_json::Value>,
@@ -53,6 +258,26 @@ fn merge_codex_quota_metadata(
serde_json::Value::Object(merged) serde_json::Value::Object(merged)
} }
fn codex_quota_window_coverage(
body_json: Option<&Value>,
) -> aether_admin::provider::quota::CodexQuotaWindowCoverage {
let body = body_json.and_then(Value::as_object);
let has_account_snapshot = body
.and_then(|body| body.get("rate_limit"))
.and_then(Value::as_object)
.is_some();
let has_spark_snapshot = body
.and_then(|body| body.get("additional_rate_limits"))
.and_then(Value::as_array)
.is_some();
match (has_account_snapshot, has_spark_snapshot) {
(true, true) => aether_admin::provider::quota::CodexQuotaWindowCoverage::FullSnapshot,
(true, false) => aether_admin::provider::quota::CodexQuotaWindowCoverage::AccountSnapshot,
_ => aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch,
}
}
fn truncate_codex_reset_credit_detail_error(message: impl Into<String>) -> String { fn truncate_codex_reset_credit_detail_error(message: impl Into<String>) -> String {
let message = message.into(); let message = message.into();
let mut sanitized = message.replace('\n', " "); let mut sanitized = message.replace('\n', " ");
@@ -198,6 +423,10 @@ fn codex_consume_success_status(outcome: &str) -> &'static str {
} }
} }
fn codex_reset_credit_outcome_allows_usage_drop(outcome: &str) -> bool {
matches!(outcome, "reset" | "already_redeemed")
}
fn codex_extract_refresh_result_fields( fn codex_extract_refresh_result_fields(
refresh_payload: Option<&Value>, refresh_payload: Option<&Value>,
key_id: &str, key_id: &str,
@@ -247,12 +476,74 @@ fn codex_extract_refresh_result_fields(
) )
} }
async fn finish_codex_reset_replay(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
key: &StoredProviderCatalogKey,
credential: &ProviderTransportCredentialFence,
terminal: CodexAccountResetTerminal,
) -> Result<(StatusCode, Value), GatewayError> {
let mut refresh_status = "skipped".to_string();
let mut refresh_error = None;
let mut metadata = None;
let mut quota_snapshot = None;
if codex_reset_credit_outcome_allows_usage_drop(&terminal.outcome) {
let fence = super::shared::CodexAccountResetFence {
unix_ms: crate::clock::current_unix_ms(),
id: format!("reset:{}", terminal.idempotency_key),
generation: terminal.generation,
};
if codex_reset_fence_is_still_pending(state, &key.id, credential, &fence).await? {
match refresh_codex_quota_after_reset_until_settled(
state, provider, endpoint, key, &fence, credential,
)
.await
{
Ok(payload) => {
(refresh_status, refresh_error, metadata, quota_snapshot) =
codex_extract_refresh_result_fields(payload.as_ref(), &key.id);
}
Err(err) => {
refresh_status = "failed".to_string();
refresh_error =
Some(truncate_codex_reset_credit_detail_error(err.into_message()));
}
}
}
}
let mut payload = Map::new();
payload.insert("key_id".to_string(), json!(key.id));
payload.insert(
"status".to_string(),
json!(codex_consume_success_status(&terminal.outcome)),
);
payload.insert("outcome".to_string(), json!(terminal.outcome));
payload.insert(
"idempotency_key".to_string(),
json!(terminal.idempotency_key),
);
payload.insert("replay".to_string(), json!(true));
payload.insert("refresh_status".to_string(), json!(refresh_status));
if let Some(refresh_error) = refresh_error {
payload.insert("refresh_error".to_string(), json!(refresh_error));
}
if let Some(metadata) = metadata {
payload.insert("metadata".to_string(), metadata);
}
if let Some(quota_snapshot) = quota_snapshot {
payload.insert("quota_snapshot".to_string(), quota_snapshot);
}
Ok((StatusCode::OK, Value::Object(payload)))
}
pub(crate) async fn consume_codex_reset_credit_locally( pub(crate) async fn consume_codex_reset_credit_locally(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider, provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint, endpoint: &StoredProviderCatalogEndpoint,
key: StoredProviderCatalogKey, key: StoredProviderCatalogKey,
idempotency_key: &str, idempotency_key: &str,
expected_credential_generation: Option<&str>,
) -> Result<(StatusCode, Value), GatewayError> { ) -> Result<(StatusCode, Value), GatewayError> {
let transport = match state let transport = match state
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id) .read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
@@ -273,22 +564,48 @@ pub(crate) async fn consume_codex_reset_credit_locally(
}; };
let is_oauth_managed = provider_key_is_oauth_managed(&key, provider.provider_type.as_str()); let is_oauth_managed = provider_key_is_oauth_managed(&key, provider.provider_type.as_str());
let resolved_oauth_auth = if is_oauth_managed { if !is_oauth_managed {
state.resolve_local_oauth_header_auth(&transport).await?
} else {
None
};
if is_oauth_managed && resolved_oauth_auth.is_none() {
return Ok(( return Ok((
StatusCode::BAD_REQUEST, StatusCode::BAD_REQUEST,
json!({ json!({
"key_id": key.id, "key_id": key.id,
"status": "error", "status": "error",
"outcome": "error", "outcome": "error",
"message": "缺少 Codex OAuth 认证信息,请先重新授权/刷新 Token", "message": "Codex reset credit 仅支持 OAuth 托管账号",
}), }),
)); ));
} }
let (transport, resolved_oauth_auth, reset_credential_fence) =
match prepare_codex_oauth_request(state, &transport).await? {
CodexOAuthRequestPreparation::Ready {
transport,
auth,
credential_fence,
} => (transport, Some(auth), credential_fence),
CodexOAuthRequestPreparation::MissingAuth => {
return Ok((
StatusCode::BAD_REQUEST,
json!({
"key_id": key.id,
"status": "error",
"outcome": "error",
"message": "缺少 Codex OAuth 认证信息,请先重新授权/刷新 Token",
}),
));
}
CodexOAuthRequestPreparation::Conflict => {
return Ok((
StatusCode::CONFLICT,
json!({
"key_id": key.id,
"status": "error",
"outcome": "error",
"idempotency_key": idempotency_key,
"message": "Codex credential changed before reset credit could be consumed",
}),
));
}
};
let request_spec = match build_codex_reset_credit_consume_request_spec( let request_spec = match build_codex_reset_credit_consume_request_spec(
&transport, &transport,
@@ -309,6 +626,79 @@ pub(crate) async fn consume_codex_reset_credit_locally(
} }
}; };
let reservation = match reserve_codex_account_reset(
state,
&key.id,
reset_credential_fence.encrypted_auth_config.as_str(),
&reset_credential_fence.credential,
expected_credential_generation,
idempotency_key,
)
.await?
{
Some(CodexAccountResetReserveResult::Reserved(reservation)) => reservation,
Some(CodexAccountResetReserveResult::Replay(terminal)) => {
return finish_codex_reset_replay(
state,
provider,
endpoint,
&key,
&reset_credential_fence,
terminal,
)
.await;
}
Some(CodexAccountResetReserveResult::LegacyReplay) => {
return Ok((
StatusCode::OK,
json!({
"key_id": key.id,
"status": "success",
"outcome": "historical_replay",
"idempotency_key": idempotency_key,
"refresh_status": "skipped",
}),
));
}
Some(CodexAccountResetReserveResult::Busy(active)) => {
return Ok((
StatusCode::CONFLICT,
json!({
"key_id": key.id,
"status": "error",
"outcome": "busy",
"idempotency_key": idempotency_key,
"active_idempotency_key": active.idempotency_key,
"message": "Another Codex reset credit operation is unresolved",
}),
));
}
Some(CodexAccountResetReserveResult::CredentialGenerationMismatch) => {
return Ok((
StatusCode::CONFLICT,
json!({
"key_id": key.id,
"status": "error",
"outcome": "credential_changed",
"idempotency_key": idempotency_key,
"message": "Codex credential changed since this reset request was prepared",
}),
));
}
None => {
return Ok((
StatusCode::CONFLICT,
json!({
"key_id": key.id,
"status": "error",
"outcome": "error",
"idempotency_key": idempotency_key,
"message": "Codex reset reservation could not be persisted",
}),
));
}
};
let result = let result =
match execute_codex_reset_credit_plan(state, &transport, request_spec, None).await? { match execute_codex_reset_credit_plan(state, &transport, request_spec, None).await? {
ProviderQuotaExecutionOutcome::Response(result) => result, ProviderQuotaExecutionOutcome::Response(result) => result,
@@ -331,11 +721,11 @@ pub(crate) async fn consume_codex_reset_credit_locally(
.and_then(|body| body.json_body.as_ref()); .and_then(|body| body.json_body.as_ref());
let outcome = normalize_codex_reset_credit_consume_outcome(body_json) let outcome = normalize_codex_reset_credit_consume_outcome(body_json)
.unwrap_or_else(|| "unknown".to_string()); .unwrap_or_else(|| "unknown".to_string());
let known_non_error_outcome = matches!( let known_terminal_outcome = matches!(
outcome.as_str(), outcome.as_str(),
"reset" | "already_redeemed" | "nothing_to_reset" | "no_credit" "reset" | "already_redeemed" | "nothing_to_reset" | "no_credit"
); );
if result.status_code >= 400 && !known_non_error_outcome { if !known_terminal_outcome {
let detail = extract_execution_error_message(&result) let detail = extract_execution_error_message(&result)
.unwrap_or_else(|| format!("HTTP {}", result.status_code)); .unwrap_or_else(|| format!("HTTP {}", result.status_code));
return Ok(( return Ok((
@@ -345,19 +735,88 @@ pub(crate) async fn consume_codex_reset_credit_locally(
"status": "error", "status": "error",
"outcome": "error", "outcome": "error",
"idempotency_key": idempotency_key, "idempotency_key": idempotency_key,
"message": format!("reset credit consume 返回状态码 {}: {detail}", result.status_code), "message": format!("reset credit consume outcome is ambiguous: {detail}"),
"status_code": result.status_code, "status_code": result.status_code,
}), }),
)); ));
} }
let (refresh_status, refresh_error, metadata, quota_snapshot) = let fence_unix_ms = result
match refresh_codex_provider_quota_locally( .response_observation
.as_ref()
.map(|observation| observation.response_headers_observed_at_unix_ms)
.unwrap_or_else(crate::clock::current_unix_ms);
let Some(completed) = complete_codex_account_reset(
state,
&key.id,
reset_credential_fence.encrypted_auth_config.as_str(),
&reset_credential_fence.credential,
&reservation,
&outcome,
fence_unix_ms,
)
.await?
else {
return Ok((
StatusCode::CONFLICT,
json!({
"key_id": key.id,
"status": "error",
"outcome": "error",
"idempotency_key": idempotency_key,
"message": "Codex reset completion could not be persisted",
}),
));
};
let (effective_outcome, reset_fence) = match completed {
CodexAccountResetCompleteResult::Activated(fence) => (outcome.clone(), Some(fence)),
CodexAccountResetCompleteResult::Noop(terminal)
| CodexAccountResetCompleteResult::Replay(terminal) => {
let fence =
codex_reset_credit_outcome_allows_usage_drop(&terminal.outcome).then(|| {
super::shared::CodexAccountResetFence {
unix_ms: fence_unix_ms,
id: format!("reset:{}", terminal.idempotency_key),
generation: terminal.generation,
}
});
(terminal.outcome, fence)
}
};
let (refresh_status, refresh_error, metadata, quota_snapshot) = match reset_fence.as_ref() {
Some(reset_fence) => {
match refresh_codex_quota_after_reset_until_settled(
state,
provider,
endpoint,
&key,
reset_fence,
&reset_credential_fence,
)
.await
{
Ok(refresh_payload) => {
codex_extract_refresh_result_fields(refresh_payload.as_ref(), &key.id)
}
Err(err) => (
"failed".to_string(),
Some(truncate_codex_reset_credit_detail_error(err.into_message())),
None,
None,
),
}
}
None => match refresh_codex_provider_quota_locally_with_reset_fence(
state, state,
provider, provider,
endpoint, endpoint,
vec![key.clone()], vec![key.clone()],
None, None,
None,
None,
Some(&reset_credential_fence),
) )
.await .await
{ {
@@ -370,15 +829,16 @@ pub(crate) async fn consume_codex_reset_credit_locally(
None, None,
None, None,
), ),
}; },
};
let mut payload = Map::new(); let mut payload = Map::new();
payload.insert("key_id".to_string(), json!(key.id)); payload.insert("key_id".to_string(), json!(key.id));
payload.insert( payload.insert(
"status".to_string(), "status".to_string(),
json!(codex_consume_success_status(&outcome)), json!(codex_consume_success_status(&effective_outcome)),
); );
payload.insert("outcome".to_string(), json!(outcome)); payload.insert("outcome".to_string(), json!(effective_outcome));
payload.insert("idempotency_key".to_string(), json!(idempotency_key)); payload.insert("idempotency_key".to_string(), json!(idempotency_key));
payload.insert("refresh_status".to_string(), json!(refresh_status)); payload.insert("refresh_status".to_string(), json!(refresh_status));
if let Some(refresh_error) = refresh_error { if let Some(refresh_error) = refresh_error {
@@ -400,6 +860,29 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
endpoint: &StoredProviderCatalogEndpoint, endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>, keys: Vec<StoredProviderCatalogKey>,
proxy_override: Option<ProxySnapshot>, proxy_override: Option<ProxySnapshot>,
) -> Result<Option<serde_json::Value>, GatewayError> {
refresh_codex_provider_quota_locally_with_reset_fence(
state,
provider,
endpoint,
keys,
proxy_override,
None,
None,
None,
)
.await
}
async fn refresh_codex_provider_quota_locally_with_reset_fence(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
proxy_override: Option<ProxySnapshot>,
account_reset_fence_id: Option<&str>,
authoritative_reset_generation: Option<u64>,
expected_reset_credential: Option<&crate::state::ProviderTransportCredentialFence>,
) -> Result<Option<serde_json::Value>, GatewayError> { ) -> Result<Option<serde_json::Value>, GatewayError> {
let mut results = Vec::new(); let mut results = Vec::new();
let mut success_count = 0usize; let mut success_count = 0usize;
@@ -412,7 +895,7 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
for key in keys { for key in keys {
let had_oauth_refresh_issue = let had_oauth_refresh_issue =
codex_oauth_refresh_issue_reason(key.oauth_invalid_reason.as_deref()); codex_oauth_refresh_issue_reason(key.oauth_invalid_reason.as_deref());
let transport = match state let initial_transport = match state
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id) .read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
.await? .await?
{ {
@@ -429,48 +912,70 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
} }
}; };
let is_oauth_managed = provider_key_is_oauth_managed(&key, provider.provider_type.as_str()); let is_oauth_managed = provider_key_is_oauth_managed(&key, provider.provider_type.as_str());
let quota_auth_config_fence = if is_oauth_managed { let (transport, resolved_oauth_auth, quota_credential_fence) = if is_oauth_managed {
match state match prepare_codex_oauth_request(state, &initial_transport).await? {
.app() CodexOAuthRequestPreparation::Ready {
.capture_provider_transport_auth_config_fence(&transport) transport,
.await? auth,
{ credential_fence,
Some(ciphertext) => Some(ciphertext), } => (transport, Some(auth), Some(credential_fence)),
None => { CodexOAuthRequestPreparation::MissingAuth => {
failed_count += 1; failed_count += 1;
results.push(json!({ results.push(json!({
"key_id": key.id, "key_id": key.id,
"key_name": key.name, "key_name": key.name,
"status": "error", "status": "error",
"message": "OAuth credential changed before quota refresh", "message": "缺少 Codex OAuth 认证信息,请先重新授权/刷新 Token",
})); }));
continue; continue;
} }
CodexOAuthRequestPreparation::Conflict => {
if quota_key_auto_removed(state, &key.id).await? {
auto_removed_count += 1;
results.push(oauth_refresh_auto_removed_result(&key));
} else {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "OAuth credential changed before quota refresh",
}));
}
continue;
}
} }
} else { } else {
None (initial_transport, None, None)
}; };
if let Some(expected_reset_credential) = expected_reset_credential {
let resolved_oauth_auth = if is_oauth_managed { if quota_credential_fence.as_ref() != Some(expected_reset_credential) {
state.resolve_local_oauth_header_auth(&transport).await? failed_count += 1;
} else { results.push(json!({
None "key_id": key.id,
}; "key_name": key.name,
if is_oauth_managed && quota_key_auto_removed(state, &key.id).await? { "status": "error",
auto_removed_count += 1; "message": "Codex credential changed after reset credit was consumed",
results.push(oauth_refresh_auto_removed_result(&key)); }));
continue; continue;
} }
if is_oauth_managed && resolved_oauth_auth.is_none() {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 Codex OAuth 认证信息,请先重新授权/刷新 Token",
}));
continue;
} }
let transport_codex_metadata = transport
.key
.upstream_metadata
.as_ref()
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("codex"));
let observed_reset_generation = authoritative_reset_generation.or_else(|| {
Some(
aether_admin::provider::quota::codex_quota_account_reset_generation(
transport_codex_metadata,
),
)
});
let observed_credential_generation =
aether_admin::provider::quota::codex_credential_generation(transport_codex_metadata)
.map(ToOwned::to_owned);
let request_spec = let request_spec =
match build_codex_quota_request_spec(&transport, resolved_oauth_auth.clone()) { match build_codex_quota_request_spec(&transport, resolved_oauth_auth.clone()) {
@@ -487,6 +992,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
} }
}; };
let quota_request_fallback_started_at_unix_ms = crate::clock::current_unix_ms();
let quota_request_fallback_order_id = uuid::Uuid::now_v7().to_string();
let result = match execute_codex_quota_plan( let result = match execute_codex_quota_plan(
state, state,
&transport, &transport,
@@ -508,16 +1015,25 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
continue; continue;
} }
}; };
let now_unix_secs = SystemTime::now() let quota_response_fallback_observed_at_unix_ms = crate::clock::current_unix_ms();
.duration_since(UNIX_EPOCH) let quota_response_observation = result.response_observation.as_ref();
.ok() let quota_request_started_at_unix_ms = quota_response_observation
.map(|duration| duration.as_secs()) .map(|observation| observation.request_started_at_unix_ms)
.unwrap_or(0); .unwrap_or(quota_request_fallback_started_at_unix_ms);
let quota_response_observed_at_unix_ms = quota_response_observation
.map(|observation| observation.response_headers_observed_at_unix_ms)
.unwrap_or(quota_response_fallback_observed_at_unix_ms);
let quota_request_order_id = quota_response_observation
.map(|observation| observation.request_order_id.as_str())
.unwrap_or(quota_request_fallback_order_id.as_str());
let now_unix_secs = quota_response_observed_at_unix_ms / 1_000;
let header_metadata = parse_codex_usage_headers(&result.headers, now_unix_secs); let header_metadata = parse_codex_usage_headers(&result.headers, now_unix_secs);
let mut metadata_update = header_metadata let mut metadata_update = header_metadata
.as_ref() .as_ref()
.map(|metadata| json!({ "codex": metadata })); .map(|metadata| json!({ "codex": metadata }));
let mut quota_window_coverage =
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch;
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) = (None, None); let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) = (None, None);
let mut status = "error".to_string(); let mut status = "error".to_string();
let mut message = None::<String>; let mut message = None::<String>;
@@ -544,6 +1060,7 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
now_unix_secs, now_unix_secs,
) )
.await?; .await?;
quota_window_coverage = codex_quota_window_coverage(Some(body_json));
metadata_update = Some(json!({ metadata_update = Some(json!({
"codex": codex_metadata "codex": codex_metadata
})); }));
@@ -586,6 +1103,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
} }
402 => { 402 => {
if codex_looks_like_workspace_deactivated(err_msg.as_deref()) { if codex_looks_like_workspace_deactivated(err_msg.as_deref()) {
quota_window_coverage =
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch;
let mut codex_meta = metadata_update let mut codex_meta = metadata_update
.as_ref() .as_ref()
.and_then(|value| value.get("codex")) .and_then(|value| value.get("codex"))
@@ -627,6 +1146,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
oauth_invalid_reason = reason; oauth_invalid_reason = reason;
status = "workspace_deactivated".to_string(); status = "workspace_deactivated".to_string();
} else { } else {
quota_window_coverage =
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch;
let plan_type = transport let plan_type = transport
.key .key
.decrypted_auth_config .decrypted_auth_config
@@ -667,24 +1188,45 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
} }
} }
let persisted = if let Some(expected_auth_config) = quota_auth_config_fence.as_deref() { let persisted = if let Some(expected_credential) = quota_credential_fence.as_ref() {
persist_fenced_provider_quota_refresh_state( persist_fenced_provider_quota_refresh_state(
state, state,
&key.id, &key.id,
expected_auth_config, expected_credential.encrypted_auth_config.as_str(),
metadata_update.as_ref(), metadata_update.as_ref(),
oauth_invalid_at_unix_secs, oauth_invalid_at_unix_secs,
oauth_invalid_reason.clone(), oauth_invalid_reason.clone(),
aether_admin::provider::quota::CodexQuotaMergeContext {
observed_at_unix_secs: now_unix_secs,
request_started_at_unix_ms: Some(quota_request_started_at_unix_ms),
request_order_id: Some(quota_request_order_id),
observed_reset_generation,
authoritative_reset_generation,
observed_credential_generation: observed_credential_generation.as_deref(),
account_reset_fence_id,
coverage: quota_window_coverage,
},
Some(&expected_credential.credential),
) )
.await? .await?
} else { } else {
persist_provider_quota_refresh_state( persist_codex_provider_quota_refresh_state(
state, state,
&key.id, &key.id,
metadata_update.as_ref(), metadata_update.as_ref(),
oauth_invalid_at_unix_secs, oauth_invalid_at_unix_secs,
oauth_invalid_reason.clone(), oauth_invalid_reason.clone(),
None, None,
aether_admin::provider::quota::CodexQuotaMergeContext {
observed_at_unix_secs: now_unix_secs,
request_started_at_unix_ms: Some(quota_request_started_at_unix_ms),
request_order_id: Some(quota_request_order_id),
observed_reset_generation,
authoritative_reset_generation,
observed_credential_generation: observed_credential_generation.as_deref(),
account_reset_fence_id,
coverage: quota_window_coverage,
},
) )
.await? .await?
}; };
@@ -698,26 +1240,57 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
})); }));
continue; continue;
} }
let credential_cas_delete = quota_auth_config_fence.as_ref().map(|auth_config| { let persisted_key = state
.read_provider_catalog_keys_by_ids(&[key.id.clone()])
.await?
.into_iter()
.next();
let persisted_codex_metadata = persisted_key
.as_ref()
.and_then(|key| key.upstream_metadata.as_ref())
.and_then(serde_json::Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.cloned();
if let Some(codex_metadata) = persisted_codex_metadata.as_ref() {
metadata_update = Some(json!({"codex": codex_metadata}));
}
let persisted_codex_object = persisted_codex_metadata
.as_ref()
.and_then(serde_json::Value::as_object);
let request_owns_persisted_oauth_state = quota_credential_fence.is_none()
|| (persisted_codex_object
.and_then(|codex| codex.get("oauth_state_request_started_at_unix_ms"))
.and_then(aether_admin::provider::quota::coerce_json_u64)
== Some(quota_request_started_at_unix_ms)
&& persisted_codex_object
.and_then(|codex| codex.get("oauth_state_request_id"))
.and_then(serde_json::Value::as_str)
== Some(quota_request_order_id));
let credential_cas_delete = quota_credential_fence.as_ref().map(|credential_fence| {
ProviderCatalogKeyOAuthCredentialCasDelete { ProviderCatalogKeyOAuthCredentialCasDelete {
key_id: key.id.clone(), key_id: key.id.clone(),
expected_encrypted_auth_config: Some(auth_config.clone()), expected_encrypted_auth_config: Some(
expected_credential: ProviderCatalogKeyOAuthCredentialFence { credential_fence.encrypted_auth_config.clone(),
encrypted_api_key: key.encrypted_api_key.clone(), ),
auth_type: key.auth_type.clone(), expected_credential: credential_fence.credential.clone(),
provider_id: key.provider_id.clone(), expected_upstream_metadata_namespace: Some(
provider_type: provider.provider_type.clone(), ProviderCatalogUpstreamMetadataNamespaceExpectation {
}, namespace: "codex".to_string(),
expected_value: persisted_codex_metadata.clone(),
},
),
} }
}); });
let should_auto_remove_hard_banned = let should_auto_remove_hard_banned = request_owns_persisted_oauth_state
provider_auto_remove_banned_keys(provider.config.as_ref()) && provider_auto_remove_banned_keys(provider.config.as_ref())
&& should_auto_remove_oauth_invalid_key( && should_auto_remove_oauth_invalid_key(
&key, persisted_key.as_ref().unwrap_or(&key),
oauth_invalid_reason.as_deref(), persisted_key
matches!(status_code, Some(401 | 403)), .as_ref()
now_unix_secs, .and_then(|key| key.oauth_invalid_reason.as_deref()),
); matches!(status_code, Some(401 | 403)),
now_unix_secs,
);
let auto_removed_hard_banned = if should_auto_remove_hard_banned { let auto_removed_hard_banned = if should_auto_remove_hard_banned {
match credential_cas_delete.as_ref() { match credential_cas_delete.as_ref() {
Some(delete) => { Some(delete) => {
@@ -735,6 +1308,7 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
auto_removed_hard_banned_count += 1; auto_removed_hard_banned_count += 1;
} }
let auto_removed_quota_exhausted = if !auto_removed_hard_banned let auto_removed_quota_exhausted = if !auto_removed_hard_banned
&& request_owns_persisted_oauth_state
&& status == "quota_exhausted" && status == "quota_exhausted"
&& provider_auto_remove_quota_exhausted_keys(provider.config.as_ref()) && provider_auto_remove_quota_exhausted_keys(provider.config.as_ref())
{ {
@@ -798,7 +1372,9 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
} }
if let Some(quota_snapshot) = build_quota_snapshot_payload( if let Some(quota_snapshot) = build_quota_snapshot_payload(
"codex", "codex",
key.status_snapshot.as_ref(), persisted_key
.as_ref()
.and_then(|key| key.status_snapshot.as_ref()),
metadata_update.as_ref(), metadata_update.as_ref(),
) { ) {
payload.insert("quota_snapshot".to_string(), quota_snapshot); payload.insert("quota_snapshot".to_string(), quota_snapshot);
@@ -866,6 +1442,19 @@ mod tests {
); );
} }
#[test]
fn codex_reset_credit_only_allows_usage_drop_after_confirmed_redemption() {
assert!(codex_reset_credit_outcome_allows_usage_drop("reset"));
assert!(codex_reset_credit_outcome_allows_usage_drop(
"already_redeemed"
));
assert!(!codex_reset_credit_outcome_allows_usage_drop(
"nothing_to_reset"
));
assert!(!codex_reset_credit_outcome_allows_usage_drop("no_credit"));
assert!(!codex_reset_credit_outcome_allows_usage_drop("unknown"));
}
#[test] #[test]
fn codex_reset_credit_detail_failure_records_attempt_time() { fn codex_reset_credit_detail_failure_records_attempt_time() {
let mut metadata = Map::new(); let mut metadata = Map::new();
@@ -880,4 +1469,29 @@ mod tests {
Some(&json!(1_777_000_000u64)) Some(&json!(1_777_000_000u64))
); );
} }
#[test]
fn codex_quota_coverage_only_replaces_observed_window_families() {
assert_eq!(
codex_quota_window_coverage(Some(&json!({"credits":{"balance":5}}))),
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch
);
assert_eq!(
codex_quota_window_coverage(None),
aether_admin::provider::quota::CodexQuotaWindowCoverage::Patch
);
assert_eq!(
codex_quota_window_coverage(Some(&json!({
"rate_limit":{"primary_window":{}}
}))),
aether_admin::provider::quota::CodexQuotaWindowCoverage::AccountSnapshot
);
assert_eq!(
codex_quota_window_coverage(Some(&json!({
"rate_limit":{"primary_window":{}},
"additional_rate_limits":[]
}))),
aether_admin::provider::quota::CodexQuotaWindowCoverage::FullSnapshot
);
}
} }
@@ -796,6 +796,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 403, status_code: 403,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: None, json_body: None,
body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)), body_bytes_b64: Some(base64::engine::general_purpose::STANDARD.encode(body)),
File diff suppressed because it is too large Load Diff
@@ -178,6 +178,7 @@ fn provider_query_execution_json_body_decodes_stream_encoded_json_response() {
"content-type".to_string(), "content-type".to_string(),
"application/json".to_string(), "application/json".to_string(),
)]), )]),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: None, json_body: None,
body_bytes_b64: Some(encoded_body), body_bytes_b64: Some(encoded_body),
@@ -436,6 +437,7 @@ fn provider_query_standard_test_aggregates_responses_stream_body() {
candidate_id: Some("candidate-0".to_string()), candidate_id: Some("candidate-0".to_string()),
status_code: 200, status_code: 200,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: None, json_body: None,
body_bytes_b64: Some( body_bytes_b64: Some(
@@ -468,6 +470,7 @@ fn provider_query_standard_test_aggregates_responses_image_generation_call() {
candidate_id: Some("candidate-0".to_string()), candidate_id: Some("candidate-0".to_string()),
status_code: 200, status_code: 200,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: None, json_body: None,
body_bytes_b64: Some( body_bytes_b64: Some(
@@ -601,6 +604,7 @@ fn provider_query_search_success_requires_non_empty_output() {
candidate_id: Some("candidate-0".to_string()), candidate_id: Some("candidate-0".to_string()),
status_code: 200, status_code: 200,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(body), json_body: Some(body),
body_bytes_b64: None, body_bytes_b64: None,
@@ -751,6 +755,7 @@ fn provider_query_standard_test_rejects_gemini_success_without_visible_output()
candidate_id: Some("candidate-0".to_string()), candidate_id: Some("candidate-0".to_string()),
status_code: 200, status_code: 200,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"candidates": [{ "candidates": [{
@@ -122,6 +122,7 @@ pub(crate) struct AdminProviderQuotaRefreshRequest {
#[derive(Debug, Deserialize)] #[derive(Debug, Deserialize)]
pub(crate) struct AdminCodexResetCreditConsumeRequest { pub(crate) struct AdminCodexResetCreditConsumeRequest {
pub(crate) idempotency_key: String, pub(crate) idempotency_key: String,
pub(crate) expected_credential_generation: serde_json::Value,
} }
#[derive(Debug, Deserialize)] #[derive(Debug, Deserialize)]
@@ -339,3 +340,37 @@ pub(crate) struct AdminImportProviderModelsRequest {
)] )]
pub(crate) price_per_request: Option<f64>, pub(crate) price_per_request: Option<f64>,
} }
#[cfg(test)]
mod tests {
use super::AdminCodexResetCreditConsumeRequest;
#[test]
fn codex_reset_credit_consume_requires_an_explicit_credential_generation() {
assert!(
serde_json::from_value::<AdminCodexResetCreditConsumeRequest>(
serde_json::json!({"idempotency_key":"reset-old-client"}),
)
.is_err()
);
let legacy_account =
serde_json::from_value::<AdminCodexResetCreditConsumeRequest>(serde_json::json!({
"idempotency_key":"reset-legacy-account",
"expected_credential_generation":null,
}))
.expect("explicit null should fence an account without a generation");
assert!(legacy_account.expected_credential_generation.is_null());
let generated_account =
serde_json::from_value::<AdminCodexResetCreditConsumeRequest>(serde_json::json!({
"idempotency_key":"reset-generated-account",
"expected_credential_generation":"credential-v2",
}))
.expect("string generation should deserialize");
assert_eq!(
generated_account.expected_credential_generation,
serde_json::json!("credential-v2")
);
}
}
@@ -1,3 +1,4 @@
use crate::handlers::admin::provider::oauth::provisioning::rotate_codex_credential_generation;
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyCreateRequest; use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyCreateRequest;
use crate::handlers::admin::provider::write::normalize::{ use crate::handlers::admin::provider::write::normalize::{
normalize_allow_auth_channel_mismatch_formats, normalize_api_format_json_object_keys, normalize_allow_auth_channel_mismatch_formats, normalize_api_format_json_object_keys,
@@ -216,6 +217,7 @@ pub(crate) async fn build_admin_create_provider_key_record(
)?; )?;
key.created_at_unix_ms = Some(now_unix_secs); key.created_at_unix_ms = Some(now_unix_secs);
key.updated_at_unix_secs = Some(now_unix_secs); key.updated_at_unix_secs = Some(now_unix_secs);
rotate_codex_credential_generation(&mut key, &provider.provider_type);
Ok(key) Ok(key)
} }
@@ -6,6 +6,7 @@ pub(crate) use self::update::build_admin_update_provider_key_record;
pub(crate) use self::update::{ pub(crate) use self::update::{
admin_provider_key_update_requires_immediate_model_fetch, admin_provider_key_update_requires_immediate_model_fetch,
build_admin_update_provider_key_record_with_existing_keys, build_admin_update_provider_key_record_with_existing_keys,
build_provider_catalog_key_admin_cas_update,
}; };
mod batch; mod batch;
@@ -1,3 +1,4 @@
use crate::handlers::admin::provider::oauth::provisioning::rotate_codex_credential_generation;
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdatePatch; use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdatePatch;
use crate::handlers::admin::provider::write::normalize::{ use crate::handlers::admin::provider::write::normalize::{
normalize_allow_auth_channel_mismatch_formats, normalize_api_format_json_object_keys, normalize_allow_auth_channel_mismatch_formats, normalize_api_format_json_object_keys,
@@ -13,6 +14,7 @@ use crate::handlers::admin::shared::{
use crate::handlers::shared::normalize_optional_api_key_concurrent_limit; use crate::handlers::shared::normalize_optional_api_key_concurrent_limit;
use crate::provider_key_auth::provider_key_is_oauth_managed; use crate::provider_key_auth::provider_key_is_oauth_managed;
use aether_data_contracts::repository::provider_catalog::{ use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyOAuthCredentialFence,
StoredProviderCatalogKey, StoredProviderCatalogProvider, StoredProviderCatalogKey, StoredProviderCatalogProvider,
}; };
use aether_provider_transport::provider_types::provider_type_is_fixed; use aether_provider_transport::provider_types::provider_type_is_fixed;
@@ -368,6 +370,12 @@ pub(crate) fn build_admin_update_provider_key_record_with_existing_keys(
.duration_since(UNIX_EPOCH) .duration_since(UNIX_EPOCH)
.ok() .ok()
.map(|duration| duration.as_secs()); .map(|duration| duration.as_secs());
let credential_identity_changed = !updated.auth_type.eq_ignore_ascii_case(&existing.auth_type)
|| updated.encrypted_api_key != existing.encrypted_api_key
|| updated.encrypted_auth_config != existing.encrypted_auth_config;
if credential_identity_changed {
rotate_codex_credential_generation(&mut updated, &provider.provider_type);
}
Ok(updated) Ok(updated)
} }
@@ -382,6 +390,49 @@ pub(crate) fn admin_provider_key_update_requires_immediate_model_fetch(
&& (!existing.auto_fetch_models || filters_changed || locked_models_changed) && (!existing.auto_fetch_models || filters_changed || locked_models_changed)
} }
pub(crate) fn build_provider_catalog_key_admin_cas_update(
existing: &StoredProviderCatalogKey,
updated: StoredProviderCatalogKey,
provider_type: &str,
) -> ProviderCatalogKeyAdminCasUpdate {
let previous_generation = existing
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/codex/credential_generation"))
.and_then(serde_json::Value::as_str);
let next_generation = updated
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/codex/credential_generation"))
.and_then(serde_json::Value::as_str);
let credential_changed = existing.auth_type != updated.auth_type
|| existing.encrypted_api_key != updated.encrypted_api_key
|| existing.encrypted_auth_config != updated.encrypted_auth_config;
let codex_rotation = provider_type
.trim()
.eq_ignore_ascii_case("codex")
.then(|| next_generation.filter(|next| Some(*next) != previous_generation))
.flatten()
.map(|generation| {
json!({
aether_admin::provider::quota::CODEX_CREDENTIAL_GENERATION_KEY: generation,
})
});
ProviderCatalogKeyAdminCasUpdate {
expected_encrypted_auth_config: existing.encrypted_auth_config.clone(),
expected_credential: ProviderCatalogKeyOAuthCredentialFence {
encrypted_api_key: existing.encrypted_api_key.clone(),
auth_type: existing.auth_type.clone(),
provider_id: existing.provider_id.clone(),
provider_type: provider_type.to_string(),
},
key: updated,
codex_rotation,
reset_oauth_runtime: credential_changed,
}
}
fn raw_secret_auth_type(value: &str) -> bool { fn raw_secret_auth_type(value: &str) -> bool {
matches!( matches!(
value.trim().to_ascii_lowercase().as_str(), value.trim().to_ascii_lowercase().as_str(),
@@ -186,6 +186,15 @@ impl<'a> AdminAppState<'a> {
self.app.update_provider_catalog_key(key).await self.app.update_provider_catalog_key(key).await
} }
pub(crate) async fn compare_and_update_provider_catalog_key_admin_state(
&self,
update: &aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyAdminCasUpdate,
) -> Result<bool, GatewayError> {
self.app
.compare_and_update_provider_catalog_key_admin_state(update)
.await
}
pub(crate) async fn compare_and_update_provider_catalog_key_adaptive_state( pub(crate) async fn compare_and_update_provider_catalog_key_adaptive_state(
&self, &self,
update: &aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyAdaptiveStateUpdate, update: &aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyAdaptiveStateUpdate,
@@ -513,6 +513,7 @@ impl<'a> AdminAppState<'a> {
provider_id: key.provider_id.clone(), provider_id: key.provider_id.clone(),
provider_type: provider.provider_type.clone(), provider_type: provider.provider_type.clone(),
}, },
expected_upstream_metadata_namespace: None,
}, },
) )
.await .await
@@ -4,10 +4,12 @@ use crate::api::ai::admin_endpoint_signature_parts;
use crate::handlers::admin::admin_provider_pool_config; use crate::handlers::admin::admin_provider_pool_config;
use crate::handlers::admin::model::ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY; use crate::handlers::admin::model::ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY;
use crate::handlers::admin::provider::endpoints_admin::payloads::AdminProviderEndpointUpdatePatch; use crate::handlers::admin::provider::endpoints_admin::payloads::AdminProviderEndpointUpdatePatch;
use crate::handlers::admin::provider::oauth::provisioning::ensure_codex_credential_generation_rotated;
use crate::handlers::admin::provider::shared::payloads::{ use crate::handlers::admin::provider::shared::payloads::{
AdminProviderCreateRequest, AdminProviderKeyCreateRequest, AdminProviderKeyUpdatePatch, AdminProviderCreateRequest, AdminProviderKeyCreateRequest, AdminProviderKeyUpdatePatch,
AdminProviderUpdatePatch, AdminProviderUpdatePatch,
}; };
use crate::handlers::admin::provider::write::keys::build_provider_catalog_key_admin_cas_update;
use crate::handlers::admin::shared::{ use crate::handlers::admin::shared::{
normalize_json_array, normalize_json_object, normalize_string_list, normalize_json_array, normalize_json_object, normalize_string_list,
}; };
@@ -377,10 +379,14 @@ fn normalize_import_key_raw_payload(
fn apply_imported_oauth_key_credentials( fn apply_imported_oauth_key_credentials(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
provider_type: &str,
previous_codex_credential_generation: Option<&str>,
raw_key: &Map<String, Value>, raw_key: &Map<String, Value>,
normalized_auth_config: Option<&Value>, normalized_auth_config: Option<&Value>,
record: &mut aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey, record: &mut aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey,
) -> Result<bool, String> { ) -> Result<bool, String> {
let previous_encrypted_api_key = record.encrypted_api_key.clone();
let previous_encrypted_auth_config = record.encrypted_auth_config.clone();
let mut credentials_supplied = false; let mut credentials_supplied = false;
let mut api_key_supplied = false; let mut api_key_supplied = false;
if let Some(api_key_value) = raw_key.get("api_key") { if let Some(api_key_value) = raw_key.get("api_key") {
@@ -424,10 +430,19 @@ fn apply_imported_oauth_key_credentials(
api_key_supplied, api_key_supplied,
); );
let credential_material_changed = record.encrypted_api_key != previous_encrypted_api_key
|| record.encrypted_auth_config != previous_encrypted_auth_config;
if credentials_supplied { if credentials_supplied {
record.oauth_invalid_at_unix_secs = None; record.oauth_invalid_at_unix_secs = None;
record.oauth_invalid_reason = None; record.oauth_invalid_reason = None;
} }
if credential_material_changed {
ensure_codex_credential_generation_rotated(
record,
provider_type,
previous_codex_credential_generation,
);
}
Ok(credentials_supplied) Ok(credentials_supplied)
} }
@@ -1861,6 +1876,15 @@ impl<'a> AdminAppState<'a> {
if let Some(existing_index) = existing_key_index { if let Some(existing_index) = existing_key_index {
let existing_key = existing_keys[existing_index].clone(); let existing_key = existing_keys[existing_index].clone();
let previous_codex_credential_generation = existing_key
.upstream_metadata
.as_ref()
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.and_then(|codex| {
aether_admin::provider::quota::codex_credential_generation(Some(codex))
})
.map(ToOwned::to_owned);
match merge_mode { match merge_mode {
AdminImportMergeMode::Skip => { AdminImportMergeMode::Skip => {
stats.keys.skipped += 1; stats.keys.skipped += 1;
@@ -1890,6 +1914,8 @@ impl<'a> AdminAppState<'a> {
let oauth_credentials_supplied = if auth_type == "oauth" { let oauth_credentials_supplied = if auth_type == "oauth" {
invalid!(apply_imported_oauth_key_credentials( invalid!(apply_imported_oauth_key_credentials(
self, self,
&provider.provider_type,
previous_codex_credential_generation.as_deref(),
&raw_key, &raw_key,
normalized_auth_config.as_ref(), normalized_auth_config.as_ref(),
&mut updated, &mut updated,
@@ -1903,8 +1929,31 @@ impl<'a> AdminAppState<'a> {
imported_key.fingerprint.clone(), imported_key.fingerprint.clone(),
"fingerprint", "fingerprint",
)); ));
let Some(mut persisted) = let admin_update = build_provider_catalog_key_admin_cas_update(
self.update_provider_catalog_key(&updated).await? &existing_key,
updated.clone(),
&provider.provider_type,
);
if !self
.compare_and_update_provider_catalog_key_admin_state(&admin_update)
.await?
{
return Ok(Err((
http::StatusCode::CONFLICT,
json!({
"detail": format!(
"Provider '{provider_name}' 的 Key 已被其他请求更新,请重试"
)
}),
)));
}
let Some(mut persisted) = self
.read_provider_catalog_keys_by_ids(std::slice::from_ref(
&updated.id,
))
.await?
.into_iter()
.next()
else { else {
return Ok(Err(invalid_request(format!( return Ok(Err(invalid_request(format!(
"更新 Provider '{provider_name}' 的 Key 失败" "更新 Provider '{provider_name}' 的 Key 失败"
@@ -1926,16 +1975,15 @@ impl<'a> AdminAppState<'a> {
persisted = reloaded; persisted = reloaded;
} }
if oauth_credentials_supplied { if oauth_credentials_supplied {
if !self
.clear_provider_catalog_key_oauth_invalid_marker(&updated.id)
.await?
{
return Ok(Err(invalid_request(format!(
"更新 Provider '{provider_name}' 的 Key 失败"
))));
}
let Some(reloaded) = self let Some(reloaded) = self
.reset_provider_catalog_key_recovery_state(&updated.id) .reset_provider_catalog_key_recovery_state_fenced(
&updated.id,
updated.encrypted_auth_config.as_deref().ok_or_else(|| {
GatewayError::Internal(format!(
"OAuth Provider '{provider_name}' imported without auth_config"
))
})?,
)
.await? .await?
else { else {
return Ok(Err(invalid_request(format!( return Ok(Err(invalid_request(format!(
@@ -1975,6 +2023,8 @@ impl<'a> AdminAppState<'a> {
let oauth_credentials_supplied = if auth_type == "oauth" { let oauth_credentials_supplied = if auth_type == "oauth" {
invalid!(apply_imported_oauth_key_credentials( invalid!(apply_imported_oauth_key_credentials(
self, self,
&provider.provider_type,
None,
&raw_key, &raw_key,
normalized_auth_config.as_ref(), normalized_auth_config.as_ref(),
&mut record, &mut record,
@@ -867,6 +867,7 @@ mod tests {
candidate_id: None, candidate_id: None,
status_code: 200, status_code: 200,
headers: Default::default(), headers: Default::default(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(body), json_body: Some(body),
body_bytes_b64: None, body_bytes_b64: None,
@@ -237,6 +237,13 @@ pub(crate) struct LocalOAuthInvalidationEffect<'a> {
pub(crate) response_text: Option<&'a str>, 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)] #[derive(Debug, Clone, Copy)]
pub(crate) enum LocalExecutionEffect<'a> { pub(crate) enum LocalExecutionEffect<'a> {
AttemptFailure(LocalAttemptFailureEffect), AttemptFailure(LocalAttemptFailureEffect),
@@ -245,6 +252,7 @@ pub(crate) enum LocalExecutionEffect<'a> {
HealthSuccess(LocalHealthSuccessEffect), HealthSuccess(LocalHealthSuccessEffect),
AdaptiveSuccess(LocalAdaptiveSuccessEffect), AdaptiveSuccess(LocalAdaptiveSuccessEffect),
OauthInvalidation(LocalOAuthInvalidationEffect<'a>), OauthInvalidation(LocalOAuthInvalidationEffect<'a>),
OauthSuccess(LocalOAuthSuccessEffect<'a>),
PoolSuccessSync { PoolSuccessSync {
payload: &'a GatewaySyncReportRequest, payload: &'a GatewaySyncReportRequest,
}, },
@@ -255,6 +263,68 @@ pub(crate) enum LocalExecutionEffect<'a> {
PoolStreamTimeout, 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 { struct PoolFeedbackContext {
pool_config: AdminProviderPoolConfig, pool_config: AdminProviderPoolConfig,
sticky_session_token: Option<String>, sticky_session_token: Option<String>,
@@ -302,6 +372,9 @@ pub(crate) async fn apply_local_execution_effect(
LocalExecutionEffect::OauthInvalidation(effect) => { LocalExecutionEffect::OauthInvalidation(effect) => {
record_oauth_invalidation_effect(state, context, effect).await; record_oauth_invalidation_effect(state, context, effect).await;
} }
LocalExecutionEffect::OauthSuccess(effect) => {
record_oauth_success_effect(state, context, effect).await;
}
LocalExecutionEffect::PoolSuccessSync { payload } => { LocalExecutionEffect::PoolSuccessSync { payload } => {
record_sync_pool_success_effect(state, context, payload).await; record_sync_pool_success_effect(state, context, payload).await;
release_pool_key_lease_effect(state, context).await; release_pool_key_lease_effect(state, context).await;
@@ -349,6 +422,16 @@ fn report_context_string_field<'a>(
.filter(|value| !value.is_empty()) .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> { fn local_scheduler_affinity_cache_key(report_context: Option<&Value>) -> Option<String> {
let client_session_affinity = local_client_session_affinity(report_context); let client_session_affinity = local_client_session_affinity(report_context);
let policy_context = scheduler_affinity_policy_context_from_report_context(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 { ) else {
return; 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 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 .await
{ {
warn!( 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 { fn execution_plan_uses_codex_agent_identity(plan: &ExecutionPlan) -> bool {
execution_plan_authorization(plan) execution_plan_authorization(plan)
.is_some_and(crate::provider_transport::is_codex_agent_identity_authorization) .is_some_and(crate::provider_transport::is_codex_agent_identity_authorization)
@@ -1802,6 +1963,7 @@ mod tests {
StoredProviderCatalogProvider, StoredProviderCatalogProvider,
}; };
use aether_test_support::ManagedRedisServer; use aether_test_support::ManagedRedisServer;
use aether_usage_runtime::GatewaySyncReportRequest;
use serde_json::{json, Value}; use serde_json::{json, Value};
use super::{ use super::{
@@ -1811,11 +1973,13 @@ mod tests {
pool_score_hard_state_for_status, resolve_pool_feedback_context, pool_score_hard_state_for_status, resolve_pool_feedback_context,
LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect,
LocalExecutionEffect, LocalExecutionEffectContext, LocalHealthFailureEffect, LocalExecutionEffect, LocalExecutionEffectContext, LocalHealthFailureEffect,
LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalPoolErrorEffect, LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalOAuthSuccessEffect,
ProviderKeyEffectLockPool, LocalPoolErrorEffect, ProviderKeyEffectLockPool,
}; };
use crate::data::{GatewayDataConfig, GatewayDataState}; 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::scheduler::affinity::SCHEDULER_AFFINITY_TTL;
use crate::AppState; use crate::AppState;
use aether_scheduler_core::{ 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] #[tokio::test]
async fn oauth_invalidation_marks_generic_codex_403_as_token_invalid() { async fn oauth_invalidation_marks_generic_codex_403_as_token_invalid() {
let state = codex_state(); let state = codex_state();
@@ -3581,7 +4016,7 @@ mod tests {
} }
#[tokio::test] #[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 state = codex_state_with_auto_remove();
let plan = sample_codex_plan(); let plan = sample_codex_plan();
@@ -3600,6 +4035,39 @@ mod tests {
) )
.await; .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 let stored_keys = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id)) .read_provider_catalog_keys_by_ids(std::slice::from_ref(&plan.key_id))
.await .await
@@ -3709,6 +4177,155 @@ mod tests {
assert_eq!(stored_key.oauth_invalid_reason, None); 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] #[tokio::test]
async fn health_failure_updates_codex_key_for_current_bearer_request() { async fn health_failure_updates_codex_key_for_current_bearer_request() {
let state = codex_state(); let state = codex_state();
+4 -4
View File
@@ -33,10 +33,10 @@ pub(crate) use self::classifier::{
LocalTransportFailoverClassification, LocalTransportFailoverClassification,
}; };
pub(crate) use self::effects::{ pub(crate) use self::effects::{
apply_local_execution_effect, LocalAdaptiveRateLimitEffect, LocalAdaptiveSuccessEffect, apply_local_execution_effect, spawn_local_oauth_success_effect, LocalAdaptiveRateLimitEffect,
LocalAttemptFailureEffect, LocalExecutionEffect, LocalExecutionEffectContext, LocalAdaptiveSuccessEffect, LocalAttemptFailureEffect, LocalExecutionEffect,
LocalHealthFailureEffect, LocalHealthSuccessEffect, LocalOAuthInvalidationEffect, LocalExecutionEffectContext, LocalHealthFailureEffect, LocalHealthSuccessEffect,
LocalPoolErrorEffect, LocalOAuthInvalidationEffect, LocalOAuthSuccessEffect, LocalPoolErrorEffect,
}; };
pub(crate) use self::health::{ pub(crate) use self::health::{
project_local_failure_health, project_local_key_circuit_closed, 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_string(),
local_failover_policy_to_value(&local_failover_policy_from_transport(transport)), 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) 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] #[test]
fn transport_error_failover_defaults_to_continue_and_accepts_explicit_stop() { fn transport_error_failover_defaults_to_continue_and_accepts_explicit_stop() {
let default_policy = let default_policy =
@@ -1,6 +1,6 @@
use std::collections::{BTreeMap, HashMap}; use std::collections::BTreeMap;
use std::sync::{Mutex, OnceLock}; use std::sync::OnceLock;
use std::time::{Duration, Instant}; use std::time::Duration;
use aether_admin::provider::quota as admin_provider_quota_pure; use aether_admin::provider::quota as admin_provider_quota_pure;
use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyRuntimeMetadataUpdate; 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::log_ids::short_request_id;
use crate::{AppState, GatewayError}; 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; 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_CHINESE_WAIT_DURATION_RE: OnceLock<Regex> = OnceLock::new();
static GROK_ENGLISH_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> { fn report_context_key_id(report_context: Option<&Value>) -> Option<String> {
report_context report_context
.and_then(|context| context.get("key_id")) .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) .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( fn report_context_provider_response_headers(
report_context: Option<&Value>, report_context: Option<&Value>,
) -> Option<BTreeMap<String, String>> { ) -> Option<BTreeMap<String, String>> {
@@ -91,86 +95,6 @@ fn report_context_provider_response_headers(
(!out.is_empty()).then_some(out) (!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( fn merge_metadata_object(
current: Option<&Value>, current: Option<&Value>,
section_key: &str, 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) { async fn apply_local_sync_report_effect(state: &AppState, payload: &GatewaySyncReportRequest) {
apply_local_gemini_file_mapping_report_effect(state, payload).await; apply_local_gemini_file_mapping_report_effect(state, payload).await;
if let Err(err) = sync_codex_quota_from_response_headers( if (200..300).contains(&payload.status_code) {
state, if let Err(err) = sync_codex_quota_from_response_headers(
payload.report_context.as_ref(), state,
&payload.headers, payload.report_context.as_ref(),
) &payload.headers,
.await )
{ .await
warn!( {
event_name = "codex_realtime_quota_sync_failed", warn!(
log_type = "ops", event_name = "codex_realtime_quota_sync_failed",
report_kind = %payload.report_kind, log_type = "ops",
report_request_id = %short_request_id(report_request_id(payload.report_context.as_ref())), report_kind = %payload.report_kind,
error = ?err, report_request_id = %short_request_id(report_request_id(payload.report_context.as_ref())),
"gateway failed to persist codex realtime quota from sync response headers" error = ?err,
); "gateway failed to persist codex realtime quota from sync response headers"
);
}
} }
if let Err(err) = sync_grok_quota_from_report_context( if let Err(err) = sync_grok_quota_from_report_context(
state, 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) { async fn apply_local_stream_report_effect(state: &AppState, payload: &GatewayStreamReportRequest) {
if let Err(err) = sync_codex_quota_from_response_headers( if (200..300).contains(&payload.status_code) {
state, if let Err(err) = sync_codex_quota_from_response_headers(
payload.report_context.as_ref(), state,
&payload.headers, payload.report_context.as_ref(),
) &payload.headers,
.await )
{ .await
warn!( {
event_name = "codex_realtime_quota_sync_failed", warn!(
log_type = "ops", event_name = "codex_realtime_quota_sync_failed",
report_kind = %payload.report_kind, log_type = "ops",
report_request_id = %short_request_id(report_request_id(payload.report_context.as_ref())), report_kind = %payload.report_kind,
error = ?err, report_request_id = %short_request_id(report_request_id(payload.report_context.as_ref())),
"gateway failed to persist codex realtime quota from stream response headers" error = ?err,
); "gateway failed to persist codex realtime quota from stream response headers"
);
}
} }
if let Err(err) = sync_grok_quota_from_report_context( if let Err(err) = sync_grok_quota_from_report_context(
state, state,
@@ -840,103 +768,112 @@ async fn sync_codex_quota_from_response_headers(
}; };
let now_unix_secs = current_unix_secs(); 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 provider_headers = report_context_provider_response_headers(report_context);
let parsed_from_provider_headers = provider_headers.as_ref().and_then(|headers| { 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 let Some(parsed) = parsed_from_provider_headers.or_else(|| {
.or_else(|| 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)
else { }) else {
return Ok(false);
};
let Some(incoming_fingerprint) = fingerprint_codex_payload(&parsed) else {
return Ok(false); 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(); for attempt in 0..RUNTIME_METADATA_CAS_MAX_ATTEMPTS {
if get_cached_codex_quota_fingerprint(&key_id, now).as_deref() let Some(key) = state
== Some(incoming_fingerprint.as_str()) .read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
{ .await?
return Ok(false); .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(&current_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) Ok(false)
} }
#[cfg(test)] #[cfg(test)]
pub(crate) fn clear_local_report_effect_caches_for_tests() { 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();
}
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
@@ -950,6 +887,151 @@ mod tests {
use crate::data::GatewayDataState; 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] #[tokio::test]
async fn gemini_report_metadata_write_preserves_adaptive_and_other_provider_state() { async fn gemini_report_metadata_write_preserves_adaptive_and_other_provider_state() {
let provider = StoredProviderCatalogProvider::new( let provider = StoredProviderCatalogProvider::new(
+15
View File
@@ -662,6 +662,21 @@ impl AppState {
Ok(updated) Ok(updated)
} }
pub(crate) async fn compare_and_update_provider_catalog_key_admin_state(
&self,
update: &provider_catalog::ProviderCatalogKeyAdminCasUpdate,
) -> Result<bool, GatewayError> {
let updated = self
.data
.compare_and_update_provider_catalog_key_admin_state(update)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
// A conflict means another instance changed credentials. Invalidate on
// both outcomes before the caller reloads or reports the conflict.
self.invalidate_provider_routing_caches();
Ok(updated)
}
pub(crate) async fn update_provider_catalog_keys( pub(crate) async fn update_provider_catalog_keys(
&self, &self,
keys: &[provider_catalog::StoredProviderCatalogKey], keys: &[provider_catalog::StoredProviderCatalogKey],
+4 -1
View File
@@ -38,7 +38,10 @@ pub(crate) use self::cache::{
PROVIDER_TRANSPORT_SNAPSHOT_CACHE_STALE_TTL, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_STALE_TTL, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL,
}; };
pub use self::cors::FrontdoorCorsConfig; pub use self::cors::FrontdoorCorsConfig;
pub(crate) use self::oauth::AgentIdentityAuthConfigFence; pub(crate) use self::oauth::{
provider_transport_context_allows_credential_rotation, AgentIdentityAuthConfigFence,
CodexRuntimeOAuthObservation, ProviderTransportCredentialFence,
};
pub(crate) use self::types::{ pub(crate) use self::types::{
AdminWalletMutationOutcome, GatewayAdminPaymentCallbackView, GatewayUserPreferenceView, AdminWalletMutationOutcome, GatewayAdminPaymentCallbackView, GatewayUserPreferenceView,
GatewayUserSessionView, LocalExecutionRuntimeMissDiagnostic, LocalMutationOutcome, GatewayUserSessionView, LocalExecutionRuntimeMissDiagnostic, LocalMutationOutcome,
File diff suppressed because it is too large Load Diff
@@ -774,6 +774,105 @@ async fn generic_key_routes_reject_agent_identity_credential_writes() {
gateway_handle.abort(); gateway_handle.abort();
} }
#[tokio::test]
async fn generic_codex_key_credential_switch_rotates_generation_and_clears_quota() {
let mut existing_key = sample_key(
"key-codex-existing",
"provider-codex",
"openai:responses",
"old-oauth-access-token",
);
existing_key.auth_type = "oauth".to_string();
existing_key.encrypted_auth_config = Some(
encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
r#"{"provider_type":"codex","refresh_token":"old-refresh-token"}"#,
)
.expect("old auth config should encrypt"),
);
existing_key.upstream_metadata = Some(json!({
"codex": {
"credential_generation": "generation-before-switch",
"primary_used_percent": 80.0,
},
"unrelated": {"preserved": true},
}));
existing_key.status_snapshot = Some(json!({
"oauth": {"status": "valid"},
"quota": {"used_ratio": 0.8},
}));
let mut provider = sample_provider("provider-codex", "codex", 10);
provider.provider_type = "codex".to_string();
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![],
vec![existing_key],
));
let gateway = build_router_with_state(
AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog_repository.clone(),
)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.put(format!(
"{gateway_url}/api/admin/endpoints/keys/key-codex-existing"
))
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"auth_type": "api_key",
"api_key": "new-codex-api-key"
}))
.send()
.await
.expect("credential switch should complete");
assert_eq!(response.status(), StatusCode::OK);
let reloaded = provider_catalog_repository
.list_keys_by_ids(&["key-codex-existing".to_string()])
.await
.expect("key should reload");
assert_eq!(reloaded.len(), 1);
let key = &reloaded[0];
assert_eq!(key.auth_type, "api_key");
let codex = key
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.get("codex"))
.and_then(serde_json::Value::as_object)
.expect("codex metadata should exist");
assert_eq!(codex.len(), 1, "unexpected Codex metadata: {codex:?}");
assert_ne!(
codex
.get(aether_admin::provider::quota::CODEX_CREDENTIAL_GENERATION_KEY)
.and_then(serde_json::Value::as_str),
Some("generation-before-switch")
);
assert_eq!(
key.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/unrelated/preserved")),
Some(&json!(true))
);
assert_eq!(
key.status_snapshot
.as_ref()
.and_then(|snapshot| snapshot.get("quota")),
Some(&serde_json::Value::Null)
);
gateway_handle.abort();
}
#[tokio::test] #[tokio::test]
async fn provider_key_concurrent_limit_create_and_list_responses() { async fn provider_key_concurrent_limit_create_and_list_responses() {
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed( let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
@@ -1,7 +1,9 @@
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::sync::{Arc, Mutex}; use std::sync::{Arc, Mutex};
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY}; use aether_crypto::{
decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY,
};
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository; use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository;
use aether_data_contracts::repository::provider_catalog::{ use aether_data_contracts::repository::provider_catalog::{
@@ -121,6 +123,7 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_codex_with_trusted_a
"1900500000".to_string(), "1900500000".to_string(),
), ),
]), ]),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"plan_type": "plus", "plan_type": "plus",
@@ -295,6 +298,467 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_codex_with_trusted_a
upstream_handle.abort(); upstream_handle.abort();
} }
#[test]
fn gateway_codex_quota_refresh_persists_after_automatic_oauth_token_refresh() {
run_provider_quota_test(
"gateway_codex_quota_refresh_persists_after_automatic_oauth_token_refresh",
gateway_codex_quota_refresh_persists_after_automatic_oauth_token_refresh_impl,
);
}
async fn gateway_codex_quota_refresh_persists_after_automatic_oauth_token_refresh_impl() {
let token_hits = Arc::new(Mutex::new(0usize));
let token_server = Router::new().route(
"/oauth/token",
post({
let token_hits = Arc::clone(&token_hits);
move || {
let token_hits = Arc::clone(&token_hits);
async move {
*token_hits.lock().expect("mutex should lock") += 1;
Json(json!({
"access_token": "refreshed-codex-access-token",
"refresh_token": "rotated-codex-refresh-token",
"token_type": "Bearer",
"expires_in": 3_600,
"account_id": "acct-quota-refresh",
"plan_type": "plus"
}))
}
}
}),
);
let seen_requests = Arc::new(Mutex::new(Vec::<(String, String)>::new()));
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any({
let seen_requests = Arc::clone(&seen_requests);
move |request: Request| {
let seen_requests = Arc::clone(&seen_requests);
async move {
let plan: aether_contracts::ExecutionPlan = serde_json::from_slice(
&to_bytes(request.into_body(), usize::MAX)
.await
.expect("body should read"),
)
.expect("plan should parse");
seen_requests.lock().expect("mutex should lock").push((
plan.url.clone(),
plan.headers
.get("authorization")
.cloned()
.unwrap_or_default(),
));
let body_json = match plan.url.as_str() {
"https://chatgpt.com/backend-api/wham/usage" => json!({
"plan_type": "plus",
"rate_limit": {
"primary_window": {
"used_percent": 23.0,
"reset_at": 1_900_000_000u64,
"window_minutes": 300
}
}
}),
"https://chatgpt.com/backend-api/wham/rate-limit-reset-credits" => {
json!({"available_count": 0, "credits": []})
}
url => panic!("unexpected execution runtime URL: {url}"),
};
let result = aether_contracts::ExecutionResult {
request_id: plan.request_id,
candidate_id: None,
status_code: 200,
headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody {
json_body: Some(body_json),
body_bytes_b64: None,
}),
telemetry: None,
error: None,
};
(StatusCode::OK, Json(result))
}
}
}),
);
let mut key = sample_key(
"key-codex-expired-quota",
"provider-codex-expired-quota",
"openai:responses",
"expired-codex-access-token",
);
key.auth_type = "oauth".to_string();
key.expires_at_unix_secs = Some(1);
key.encrypted_auth_config = Some(
encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
&json!({
"provider_type": "codex",
"refresh_token": "expired-codex-refresh-token",
"expires_at": 1,
"account_id": "acct-quota-refresh",
"plan_type": "plus"
})
.to_string(),
)
.expect("auth config should encrypt"),
);
key.upstream_metadata = Some(json!({
"codex": {"credential_generation": "credential-quota-refresh"}
}));
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![StoredProviderCatalogProvider::new(
"provider-codex-expired-quota".to_string(),
"codex".to_string(),
Some("https://example.com".to_string()),
"codex".to_string(),
)
.expect("provider should build")],
vec![sample_endpoint(
"endpoint-codex-expired-quota",
"provider-codex-expired-quota",
"openai:responses",
"https://chatgpt.com/backend-api",
)],
vec![key],
));
let (token_url, token_handle) = start_server(token_server).await;
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let oauth_refresh =
crate::provider_transport::LocalOAuthRefreshCoordinator::with_adapters_for_tests(vec![
Arc::new(
crate::provider_transport::oauth_refresh::GenericOAuthRefreshAdapter::default()
.with_token_url_for_tests("codex", format!("{token_url}/oauth/token")),
),
]);
let gateway = build_router_with_state(
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog_repository.clone(),
)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
)
.with_oauth_refresh_coordinator_for_tests(oauth_refresh),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!(
"{gateway_url}/api/admin/endpoints/providers/provider-codex-expired-quota/refresh-quota"
))
.header(GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["success"], 1, "payload={payload}");
assert_eq!(payload["failed"], 0, "payload={payload}");
assert_eq!(payload["results"][0]["status"], "success");
assert_eq!(*token_hits.lock().expect("mutex should lock"), 1);
assert_eq!(
seen_requests.lock().expect("mutex should lock").as_slice(),
[
(
"https://chatgpt.com/backend-api/wham/usage".to_string(),
"Bearer refreshed-codex-access-token".to_string(),
),
(
"https://chatgpt.com/backend-api/wham/rate-limit-reset-credits".to_string(),
"Bearer refreshed-codex-access-token".to_string(),
),
]
);
let reloaded = provider_catalog_repository
.list_keys_by_ids(&["key-codex-expired-quota".to_string()])
.await
.expect("key should reload");
let persisted = reloaded.first().expect("key should remain installed");
let decrypted_api_key = decrypt_python_fernet_ciphertext(
DEVELOPMENT_ENCRYPTION_KEY,
persisted
.encrypted_api_key
.as_deref()
.expect("api key should persist"),
)
.expect("api key should decrypt");
assert_eq!(decrypted_api_key, "refreshed-codex-access-token");
let decrypted_auth_config = decrypt_python_fernet_ciphertext(
DEVELOPMENT_ENCRYPTION_KEY,
persisted
.encrypted_auth_config
.as_deref()
.expect("auth config should persist"),
)
.expect("auth config should decrypt");
let auth_config: serde_json::Value =
serde_json::from_str(&decrypted_auth_config).expect("auth config should parse");
assert_eq!(auth_config["refresh_token"], "rotated-codex-refresh-token");
assert_eq!(
persisted
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/codex/credential_generation")),
Some(&json!("credential-quota-refresh"))
);
assert_eq!(
persisted
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/codex/primary_used_percent")),
Some(&json!(23.0))
);
gateway_handle.abort();
execution_runtime_handle.abort();
token_handle.abort();
}
#[test]
fn gateway_codex_reset_credit_retries_until_same_window_usage_drop_is_authoritative() {
run_provider_quota_test(
"gateway_codex_reset_credit_retries_until_same_window_usage_drop_is_authoritative",
gateway_codex_reset_credit_retries_until_same_window_usage_drop_is_authoritative_impl,
);
}
async fn gateway_codex_reset_credit_retries_until_same_window_usage_drop_is_authoritative_impl() {
const RESET_FENCE_UNIX_MS: u64 = 1_800_000_000_000;
const RESET_AT_UNIX_SECS: u64 = 2_000_000_000;
let usage_hits = Arc::new(Mutex::new(0usize));
let detail_hits = Arc::new(Mutex::new(0usize));
let seen_urls = Arc::new(Mutex::new(Vec::<String>::new()));
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any({
let usage_hits = Arc::clone(&usage_hits);
let detail_hits = Arc::clone(&detail_hits);
let seen_urls = Arc::clone(&seen_urls);
move |request: Request| {
let usage_hits = Arc::clone(&usage_hits);
let detail_hits = Arc::clone(&detail_hits);
let seen_urls = Arc::clone(&seen_urls);
async move {
let plan: aether_contracts::ExecutionPlan = serde_json::from_slice(
&to_bytes(request.into_body(), usize::MAX)
.await
.expect("body should read"),
)
.expect("plan should parse");
seen_urls
.lock()
.expect("mutex should lock")
.push(plan.url.clone());
assert_eq!(
plan.headers.get("authorization").map(String::as_str),
Some("Bearer codex-reset-access-token")
);
let (body_json, response_observation) = match plan.url.as_str() {
"https://chatgpt.com/backend-api/wham/rate-limit-reset-credits/consume" => {
assert_eq!(plan.method, "POST");
assert_eq!(
plan.body.json_body,
Some(json!({"redeem_request_id": "reset-e2e"}))
);
(
json!({"outcome": "reset"}),
Some(aether_contracts::ExecutionResponseObservation {
request_started_at_unix_ms: RESET_FENCE_UNIX_MS - 100,
response_headers_observed_at_unix_ms: RESET_FENCE_UNIX_MS,
request_order_id: "consume-reset-e2e".to_string(),
}),
)
}
"https://chatgpt.com/backend-api/wham/usage" => {
let hit = {
let mut hits = usage_hits.lock().expect("mutex should lock");
*hits += 1;
*hits
};
let used_percent = match hit {
1 => 100.0,
2 => 0.0,
_ => panic!("unexpected wham/usage request #{hit}"),
};
(
json!({
"plan_type": "plus",
"rate_limit": {
"primary_window": {
"used_percent": used_percent,
"reset_at": RESET_AT_UNIX_SECS,
"window_minutes": 300
}
}
}),
Some(aether_contracts::ExecutionResponseObservation {
request_started_at_unix_ms: RESET_FENCE_UNIX_MS
+ (hit as u64 * 1_000),
response_headers_observed_at_unix_ms: RESET_FENCE_UNIX_MS
+ (hit as u64 * 1_000)
+ 100,
request_order_id: format!("usage-reset-e2e-{hit}"),
}),
)
}
"https://chatgpt.com/backend-api/wham/rate-limit-reset-credits" => {
*detail_hits.lock().expect("mutex should lock") += 1;
(json!({"available_count": 0, "credits": []}), None)
}
url => panic!("unexpected execution runtime URL: {url}"),
};
let result = aether_contracts::ExecutionResult {
request_id: plan.request_id,
candidate_id: None,
status_code: 200,
headers: BTreeMap::new(),
response_observation,
body: Some(aether_contracts::ResponseBody {
json_body: Some(body_json),
body_bytes_b64: None,
}),
telemetry: None,
error: None,
};
(StatusCode::OK, Json(result))
}
}
}),
);
let mut key = sample_key(
"key-codex-reset",
"provider-codex-reset",
"openai:responses",
"codex-reset-access-token",
);
key.auth_type = "oauth".to_string();
key.expires_at_unix_secs = Some(4_102_444_800);
key.encrypted_auth_config = Some(
encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
&json!({
"provider_type": "codex",
"refresh_token": "codex-reset-refresh-token",
"expires_at": 4_102_444_800u64,
"account_id": "acct-reset-e2e",
"plan_type": "plus"
})
.to_string(),
)
.expect("auth config should encrypt"),
);
key.upstream_metadata = Some(json!({
"codex": {
"credential_generation": "credential-reset-e2e",
"plan_type": "plus",
"primary_used_percent": 100.0,
"primary_reset_at": RESET_AT_UNIX_SECS,
"primary_window_minutes": 300,
"updated_at": (RESET_FENCE_UNIX_MS / 1_000) - 1,
"account_quota_request_started_at_unix_ms": RESET_FENCE_UNIX_MS - 1_000,
"account_quota_request_id": "usage-before-reset"
}
}));
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![StoredProviderCatalogProvider::new(
"provider-codex-reset".to_string(),
"codex".to_string(),
Some("https://example.com".to_string()),
"codex".to_string(),
)
.expect("provider should build")],
vec![sample_endpoint(
"endpoint-codex-reset",
"provider-codex-reset",
"openai:responses",
"https://chatgpt.com/backend-api",
)],
vec![key],
));
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let gateway = build_router_with_state(
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog_repository.clone(),
)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!(
"{gateway_url}/api/admin/endpoints/keys/key-codex-reset/codex-reset-credit/consume"
))
.header(GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"idempotency_key": "reset-e2e",
"expected_credential_generation": "credential-reset-e2e"
}))
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["status"], "success", "payload={payload}");
assert_eq!(payload["outcome"], "reset");
assert_eq!(payload["refresh_status"], "success");
assert_eq!(*usage_hits.lock().expect("mutex should lock"), 2);
assert_eq!(*detail_hits.lock().expect("mutex should lock"), 2);
assert_eq!(
seen_urls.lock().expect("mutex should lock").as_slice(),
[
"https://chatgpt.com/backend-api/wham/rate-limit-reset-credits/consume",
"https://chatgpt.com/backend-api/wham/usage",
"https://chatgpt.com/backend-api/wham/rate-limit-reset-credits",
"https://chatgpt.com/backend-api/wham/usage",
"https://chatgpt.com/backend-api/wham/rate-limit-reset-credits",
]
);
let reloaded = provider_catalog_repository
.list_keys_by_ids(&["key-codex-reset".to_string()])
.await
.expect("key should reload");
let codex = reloaded[0]
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.get("codex"))
.expect("codex metadata should persist");
assert_eq!(codex["primary_used_percent"], json!(0.0));
assert_eq!(codex["primary_reset_at"], json!(RESET_AT_UNIX_SECS));
assert_eq!(codex["account_quota_reset_pending"], json!(false));
assert_eq!(
codex["account_quota_reset_processed_ids"],
json!(["reset-e2e"])
);
gateway_handle.abort();
execution_runtime_handle.abort();
}
#[tokio::test] #[tokio::test]
async fn gateway_marks_codex_quota_exhausted_when_wham_usage_returns_payment_required() { async fn gateway_marks_codex_quota_exhausted_when_wham_usage_returns_payment_required() {
let upstream = Router::new().route( let upstream = Router::new().route(
@@ -318,6 +782,7 @@ async fn gateway_marks_codex_quota_exhausted_when_wham_usage_returns_payment_req
candidate_id: None, candidate_id: None,
status_code: 402, status_code: 402,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"error": { "error": {
@@ -439,6 +904,7 @@ async fn gateway_auto_removes_codex_key_when_quota_proves_oauth_invalid() {
candidate_id: None, candidate_id: None,
status_code: 401, status_code: 401,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"error": { "error": {
@@ -594,6 +1060,7 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_requested_codex_keys
"1900000000".to_string(), "1900000000".to_string(),
), ),
]), ]),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"plan_type": "plus", "plan_type": "plus",
@@ -740,6 +1207,7 @@ async fn gateway_refreshes_admin_provider_quota_for_codex_proxy_with_extended_ti
candidate_id: None, candidate_id: None,
status_code: 200, status_code: 200,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"plan_type": "plus", "plan_type": "plus",
@@ -897,6 +1365,7 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_kiro_with_trusted_ad
candidate_id: None, candidate_id: None,
status_code: 200, status_code: 200,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"subscriptionInfo": { "subscriptionInfo": {
@@ -1280,6 +1749,7 @@ async fn gateway_refresh_kiro_quota_reconciles_missing_fixed_endpoint_before_ref
candidate_id: None, candidate_id: None,
status_code: 200, status_code: 200,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"subscriptionInfo": { "subscriptionInfo": {
@@ -1443,6 +1913,7 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_gemini_cli_with_trus
candidate_id: None, candidate_id: None,
status_code: 200, status_code: 200,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"buckets": [ "buckets": [
@@ -1840,6 +2311,7 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_antigravity_with_tru
candidate_id: None, candidate_id: None,
status_code: 200, status_code: 200,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"models": { "models": {
@@ -3440,11 +3440,16 @@ async fn gateway_completes_admin_provider_oauth_key_locally_with_trusted_admin_p
let token_hits_clone = Arc::clone(&token_hits); let token_hits_clone = Arc::clone(&token_hits);
let seen_token = Arc::new(Mutex::new(None::<SeenTokenRequest>)); let seen_token = Arc::new(Mutex::new(None::<SeenTokenRequest>));
let seen_token_clone = Arc::clone(&seen_token); let seen_token_clone = Arc::clone(&seen_token);
let namespace_race_repository = Arc::new(Mutex::new(
None::<Arc<InMemoryProviderCatalogReadRepository>>,
));
let namespace_race_repository_clone = Arc::clone(&namespace_race_repository);
let token_server = Router::new().route( let token_server = Router::new().route(
"/oauth/token", "/oauth/token",
any(move |request: Request| { any(move |request: Request| {
let token_hits_inner = Arc::clone(&token_hits_clone); let token_hits_inner = Arc::clone(&token_hits_clone);
let seen_token_inner = Arc::clone(&seen_token_clone); let seen_token_inner = Arc::clone(&seen_token_clone);
let namespace_race_repository_inner = Arc::clone(&namespace_race_repository_clone);
async move { async move {
*token_hits_inner.lock().expect("mutex should lock") += 1; *token_hits_inner.lock().expect("mutex should lock") += 1;
let (parts, body) = request.into_parts(); let (parts, body) = request.into_parts();
@@ -3459,6 +3464,23 @@ async fn gateway_completes_admin_provider_oauth_key_locally_with_trusted_admin_p
body: String::from_utf8(raw_body.to_vec()) body: String::from_utf8(raw_body.to_vec())
.expect("token request body should be utf8"), .expect("token request body should be utf8"),
}); });
let repository = namespace_race_repository_inner
.lock()
.expect("mutex should lock")
.clone()
.expect("provider catalog repository should be installed");
repository
.upsert_key_upstream_metadata_namespace(
"key-codex-oauth",
"codex",
&json!({
"primary_used_percent": 80.0,
"account_quota_request_id": "concurrent-refresh"
}),
Some(1_800_000_001),
)
.await
.expect("concurrent Codex namespace update should succeed");
Json(json!({ Json(json!({
"access_token": "new-codex-access-token", "access_token": "new-codex-access-token",
"refresh_token": "new-codex-refresh-token", "refresh_token": "new-codex-refresh-token",
@@ -3492,6 +3514,33 @@ async fn gateway_completes_admin_provider_oauth_key_locally_with_trusted_admin_p
key.circuit_breaker_by_format = Some(json!({ key.circuit_breaker_by_format = Some(json!({
"openai:chat": {"state": "open"} "openai:chat": {"state": "open"}
})); }));
key.upstream_metadata = Some(json!({
"codex": {
"primary_used_percent": 100.0,
"primary_reset_at": 4_200_000_000u64,
"account_quota_reset_fence_unix_ms": 1_800_000_000_000u64,
"account_quota_reset_fence_id": "old-account-fence",
"account_quota_reset_processed_ids": ["old-account-reset"],
"account_quota_reset_pending": true
},
"unrelated_runtime": {
"preserved": true
}
}));
key.status_snapshot = Some(json!({
"oauth": {
"code": "invalid"
},
"quota": {
"provider_type": "codex",
"usage_ratio": 1.0,
"windows": [{
"kind": "primary",
"usage": 100.0,
"reset_at": 4_200_000_000u64
}]
}
}));
let score_identity = PoolMemberIdentity::provider_api_key("provider-codex", "key-codex-oauth"); let score_identity = PoolMemberIdentity::provider_api_key("provider-codex", "key-codex-oauth");
let score_scope = provider_key_pool_score_scope(); let score_scope = provider_key_pool_score_scope();
@@ -3510,6 +3559,8 @@ async fn gateway_completes_admin_provider_oauth_key_locally_with_trusted_admin_p
vec![], vec![],
vec![key], vec![key],
)); ));
*namespace_race_repository.lock().expect("mutex should lock") =
Some(Arc::clone(&provider_catalog_repository));
let pool_score_repository = let pool_score_repository =
Arc::new(InMemoryPoolMemberScoreRepository::seed(vec![invalid_score])); Arc::new(InMemoryPoolMemberScoreRepository::seed(vec![invalid_score]));
@@ -3593,6 +3644,37 @@ async fn gateway_completes_admin_provider_oauth_key_locally_with_trusted_admin_p
assert_eq!(persisted.error_count, Some(0)); assert_eq!(persisted.error_count, Some(0));
assert_eq!(persisted.health_by_format, Some(json!({}))); assert_eq!(persisted.health_by_format, Some(json!({})));
assert_eq!(persisted.circuit_breaker_by_format, Some(json!({}))); assert_eq!(persisted.circuit_breaker_by_format, Some(json!({})));
assert_eq!(
persisted
.upstream_metadata
.as_ref()
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("codex"))
.and_then(Value::as_object)
.map(|codex| codex.keys().cloned().collect::<Vec<_>>()),
Some(vec!["credential_generation".to_string()]),
"explicit Codex reauthorization must not carry quota/reset state across accounts"
);
assert!(persisted
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/codex/credential_generation"))
.and_then(Value::as_str)
.is_some_and(|generation| !generation.is_empty()));
assert_eq!(
persisted
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.get("unrelated_runtime")),
Some(&json!({"preserved": true}))
);
assert_eq!(
persisted
.status_snapshot
.as_ref()
.and_then(|snapshot| snapshot.get("quota")),
Some(&Value::Null)
);
let scores = pool_score_repository let scores = pool_score_repository
.get_pool_member_scores_by_ids(&GetPoolMemberScoresByIdsQuery { .get_pool_member_scores_by_ids(&GetPoolMemberScoresByIdsQuery {
ids: vec![provider_key_pool_score_id(&score_identity, &score_scope)], ids: vec![provider_key_pool_score_id(&score_identity, &score_scope)],
@@ -6164,6 +6246,7 @@ async fn gateway_refreshes_admin_provider_oauth_key_locally_with_trusted_admin_p
candidate_id: None, candidate_id: None,
status_code: 401, status_code: 401,
headers: std::collections::BTreeMap::new(), headers: std::collections::BTreeMap::new(),
response_observation: None,
body: None, body: None,
telemetry: None, telemetry: None,
error: None, error: None,
@@ -6523,6 +6606,7 @@ async fn gateway_manual_codex_oauth_refresh_reconciles_missing_fixed_endpoint_im
candidate_id: None, candidate_id: None,
status_code: 200, status_code: 200,
headers: std::collections::BTreeMap::new(), headers: std::collections::BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"plan_type": "plus", "plan_type": "plus",
@@ -6757,6 +6841,7 @@ async fn run_gateway_manual_kiro_oauth_refresh_maintenance_endpoint_test(
candidate_id: None, candidate_id: None,
status_code: 200, status_code: 200,
headers: std::collections::BTreeMap::new(), headers: std::collections::BTreeMap::new(),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"subscriptionInfo": { "subscriptionInfo": {
@@ -2317,6 +2317,17 @@ async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_imp
) )
.expect("auth config should encrypt"), .expect("auth config should encrypt"),
); );
existing_key.upstream_metadata = Some(json!({
"codex": {
"credential_generation": "generation-before-import",
"primary_used_percent": 90.0,
},
"unrelated": {"preserved": true},
}));
existing_key.status_snapshot = Some(json!({
"oauth": {"status": "invalid"},
"quota": {"used_ratio": 0.9},
}));
let score_identity = let score_identity =
PoolMemberIdentity::provider_api_key("provider-codex-existing", "key-codex-existing"); PoolMemberIdentity::provider_api_key("provider-codex-existing", "key-codex-existing");
@@ -2417,6 +2428,31 @@ async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_imp
assert_eq!(key.error_count, Some(0)); assert_eq!(key.error_count, Some(0));
assert_eq!(key.health_by_format, Some(json!({}))); assert_eq!(key.health_by_format, Some(json!({})));
assert_eq!(key.circuit_breaker_by_format, Some(json!({}))); assert_eq!(key.circuit_breaker_by_format, Some(json!({})));
let codex = key
.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.get("codex"))
.and_then(Value::as_object)
.expect("Codex metadata should exist");
assert_eq!(codex.len(), 1);
assert_ne!(
codex
.get(aether_admin::provider::quota::CODEX_CREDENTIAL_GENERATION_KEY)
.and_then(Value::as_str),
Some("generation-before-import")
);
assert_eq!(
key.upstream_metadata
.as_ref()
.and_then(|metadata| metadata.pointer("/unrelated/preserved")),
Some(&json!(true))
);
assert_eq!(
key.status_snapshot
.as_ref()
.and_then(|snapshot| snapshot.get("quota")),
Some(&Value::Null)
);
assert_eq!( assert_eq!(
decrypt_python_fernet_ciphertext( decrypt_python_fernet_ciphertext(
DEVELOPMENT_ENCRYPTION_KEY, DEVELOPMENT_ENCRYPTION_KEY,
@@ -657,6 +657,7 @@ fn embedding_execution_result(plan: &ExecutionPlan) -> ExecutionResult {
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code: 200, status_code: 200,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]), headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"object": "list", "object": "list",
@@ -679,6 +680,7 @@ fn gemini_embedding_execution_result(plan: &ExecutionPlan) -> ExecutionResult {
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code: 200, status_code: 200,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]), headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"model": "gemini-embedding-2-preview", "model": "gemini-embedding-2-preview",
@@ -703,6 +705,7 @@ fn vertex_gemini_embedding_execution_result(plan: &ExecutionPlan) -> ExecutionRe
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code: 200, status_code: 200,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]), headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"predictions": [ "predictions": [
@@ -727,6 +730,7 @@ fn gemini_batch_embedding_execution_result(plan: &ExecutionPlan) -> ExecutionRes
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code: 200, status_code: 200,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]), headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"model": "gemini-embedding-2-preview", "model": "gemini-embedding-2-preview",
@@ -752,6 +756,7 @@ fn aliyun_embedding_execution_result(plan: &ExecutionPlan) -> ExecutionResult {
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code: 200, status_code: 200,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]), headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"output": { "output": {
@@ -151,6 +151,7 @@ fn rerank_execution_result(plan: &ExecutionPlan) -> ExecutionResult {
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code: 200, status_code: 200,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]), headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"model": "upstream-rerank", "model": "upstream-rerank",
@@ -148,6 +148,7 @@ async fn gateway_executes_gemini_files_download_via_local_decision_gate_with_loc
"content-type".to_string(), "content-type".to_string(),
"application/octet-stream".to_string(), "application/octet-stream".to_string(),
)]), )]),
response_observation: None,
}, },
}, },
StreamFrame { StreamFrame {
@@ -774,6 +774,7 @@ async fn gateway_records_failed_usage_when_all_local_openai_chat_candidates_exha
"content-type".to_string(), "content-type".to_string(),
"application/json".to_string(), "application/json".to_string(),
)]), )]),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"error": { "error": {
@@ -1056,6 +1057,7 @@ async fn sync_transport_error_policy_stops_or_retries_candidates_end_to_end_impl
"content-type".to_string(), "content-type".to_string(),
"application/json".to_string(), "application/json".to_string(),
)]), )]),
response_observation: None,
body: Some(aether_contracts::ResponseBody { body: Some(aether_contracts::ResponseBody {
json_body: Some(json!({ json_body: Some(json!({
"id": "chatcmpl-transport-policy", "id": "chatcmpl-transport-policy",
@@ -212,6 +212,7 @@ async fn gateway_executes_openai_video_content_from_reconstructed_data_task_with
"content-type".to_string(), "content-type".to_string(),
"video/mp4".to_string(), "video/mp4".to_string(),
)]), )]),
response_observation: None,
}, },
}, },
StreamFrame { StreamFrame {
File diff suppressed because it is too large Load Diff
+6 -1
View File
@@ -2,7 +2,10 @@ use std::collections::BTreeMap;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use crate::{ExecutionError, ExecutionStreamTerminalSummary, ExecutionTelemetry}; use crate::{
ExecutionError, ExecutionResponseObservation, ExecutionStreamTerminalSummary,
ExecutionTelemetry,
};
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")] #[serde(rename_all = "snake_case")]
@@ -21,6 +24,8 @@ pub enum StreamFramePayload {
status_code: u16, status_code: u16,
#[serde(default)] #[serde(default)]
headers: BTreeMap<String, String>, headers: BTreeMap<String, String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
response_observation: Option<ExecutionResponseObservation>,
}, },
Data { Data {
#[serde(default, skip_serializing_if = "Option::is_none")] #[serde(default, skip_serializing_if = "Option::is_none")]
+1 -1
View File
@@ -19,7 +19,7 @@ pub use plan::{
TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY, TRANSPORT_HTTP_MODE_H2C_PRIOR_KNOWLEDGE, TRANSPORT_HTTP_MODE_HTTP1_ONLY,
TRANSPORT_POOL_SCOPE_KEY, TRANSPORT_POOL_SCOPE_KEY,
}; };
pub use result::{ExecutionResult, ExecutionTelemetry, ResponseBody}; pub use result::{ExecutionResponseObservation, ExecutionResult, ExecutionTelemetry, ResponseBody};
pub use usage::{ pub use usage::{
ExecutionStreamTerminalSummary, StandardizedUsage, USAGE_SERVER_NOW_UNIX_MS_HEADER, ExecutionStreamTerminalSummary, StandardizedUsage, USAGE_SERVER_NOW_UNIX_MS_HEADER,
}; };
+9
View File
@@ -5,6 +5,13 @@ use serde_json::Value;
use crate::ExecutionError; use crate::ExecutionError;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct ExecutionResponseObservation {
pub request_started_at_unix_ms: u64,
pub response_headers_observed_at_unix_ms: u64,
pub request_order_id: String,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct ExecutionTelemetry { pub struct ExecutionTelemetry {
#[serde(default, skip_serializing_if = "Option::is_none")] #[serde(default, skip_serializing_if = "Option::is_none")]
@@ -32,6 +39,8 @@ pub struct ExecutionResult {
#[serde(default)] #[serde(default)]
pub headers: BTreeMap<String, String>, pub headers: BTreeMap<String, String>,
#[serde(default, skip_serializing_if = "Option::is_none")] #[serde(default, skip_serializing_if = "Option::is_none")]
pub response_observation: Option<ExecutionResponseObservation>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub body: Option<ResponseBody>, pub body: Option<ResponseBody>,
#[serde(default, skip_serializing_if = "Option::is_none")] #[serde(default, skip_serializing_if = "Option::is_none")]
pub telemetry: Option<ExecutionTelemetry>, pub telemetry: Option<ExecutionTelemetry>,
@@ -8,8 +8,8 @@ use sqlx::{
}; };
use aether_data_contracts::repository::provider_catalog::{ use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdminCasUpdate,
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
ProviderCatalogReadRepository, ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogReadRepository, ProviderCatalogUpstreamMetadataNamespaceUpdate,
@@ -936,6 +936,84 @@ WHERE id = ?
self.reload_key(&key.id, "updated").await self.reload_key(&key.id, "updated").await
} }
pub async fn compare_and_update_key_admin_state(
&self,
update: &ProviderCatalogKeyAdminCasUpdate,
) -> Result<bool, DataLayerError> {
validate_admin_key_cas_update(update)?;
let key = &update.key;
let updated_at = key.updated_at_unix_secs.unwrap_or_else(current_unix_secs) as i64;
let rotation_json = update
.codex_rotation
.as_ref()
.map(serde_json::to_string)
.transpose()
.map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys Codex rotation is not serializable: {err}"
))
})?;
let mut builder = QueryBuilder::<MySql>::new("UPDATE provider_api_keys SET provider_id = ");
push_admin_key_assignments(&mut builder, key, updated_at)?;
if update.reset_oauth_runtime {
builder.push(", oauth_invalid_at = NULL, oauth_invalid_reason = NULL, error_count = 0");
}
if let Some(rotation_json) = rotation_json.as_deref() {
builder
.push(", upstream_metadata = JSON_SET(COALESCE(upstream_metadata, JSON_OBJECT()), '$.codex', CAST(")
.push_bind(rotation_json)
.push(" AS JSON))");
}
if update.codex_rotation.is_some() || update.reset_oauth_runtime {
builder.push(", status_snapshot = ");
match (update.codex_rotation.is_some(), update.reset_oauth_runtime) {
(true, true) => builder.push("JSON_SET(COALESCE(status_snapshot, JSON_OBJECT()), '$.quota', CAST('null' AS JSON), '$.oauth', CAST('null' AS JSON))"),
(true, false) => builder.push("JSON_SET(COALESCE(status_snapshot, JSON_OBJECT()), '$.quota', CAST('null' AS JSON))"),
(false, true) => builder.push("JSON_SET(COALESCE(status_snapshot, JSON_OBJECT()), '$.oauth', CAST('null' AS JSON))"),
(false, false) => unreachable!(),
};
}
builder
.push(" WHERE id = ")
.push_bind(&key.id)
.push(" AND BINARY api_key <=> BINARY ")
.push_bind(update.expected_credential.encrypted_api_key.as_deref())
.push(" AND BINARY auth_config <=> BINARY ")
.push_bind(update.expected_encrypted_auth_config.as_deref())
.push(" AND BINARY auth_type = BINARY ")
.push_bind(&update.expected_credential.auth_type)
.push(" AND BINARY provider_id = BINARY ")
.push_bind(&update.expected_credential.provider_id)
.push(
" AND EXISTS (SELECT 1 FROM providers WHERE BINARY providers.id = BINARY provider_api_keys.provider_id AND BINARY providers.provider_type = BINARY ",
)
.push_bind(&update.expected_credential.provider_type)
.push(")");
if update.codex_rotation.is_some() {
builder
.push(" AND JSON_TYPE(COALESCE(upstream_metadata, JSON_OBJECT())) = 'OBJECT'")
.push(" AND NOT (BINARY api_key <=> BINARY ")
.push_bind(key.encrypted_api_key.as_deref())
.push(" AND BINARY auth_config <=> BINARY ")
.push_bind(key.encrypted_auth_config.as_deref())
.push(" AND BINARY auth_type = BINARY ")
.push_bind(&key.auth_type)
.push(" AND BINARY provider_id = BINARY ")
.push_bind(&key.provider_id)
.push(")");
}
if update.codex_rotation.is_some() || update.reset_oauth_runtime {
builder.push(" AND JSON_TYPE(COALESCE(status_snapshot, JSON_OBJECT())) = 'OBJECT'");
}
let rows_affected = builder
.build()
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
pub async fn update_keys( pub async fn update_keys(
&self, &self,
keys: &[StoredProviderCatalogKey], keys: &[StoredProviderCatalogKey],
@@ -1007,6 +1085,10 @@ WHERE id = ?
|| expected.auth_type.trim().is_empty() || expected.auth_type.trim().is_empty()
|| expected.provider_id.trim().is_empty() || expected.provider_id.trim().is_empty()
|| expected.provider_type.trim().is_empty() || expected.provider_type.trim().is_empty()
|| delete
.expected_upstream_metadata_namespace
.as_ref()
.is_some_and(|expected| expected.namespace.trim().is_empty())
{ {
return Err(DataLayerError::InvalidInput( return Err(DataLayerError::InvalidInput(
"provider catalog OAuth credential CAS delete contains empty fields".to_string(), "provider catalog OAuth credential CAS delete contains empty fields".to_string(),
@@ -1030,6 +1112,33 @@ WHERE id = ?
) )
.push_bind(&expected.provider_type) .push_bind(&expected.provider_type)
.push(")"); .push(")");
if let Some(expected) = delete.expected_upstream_metadata_namespace.as_ref() {
let namespace_path = format!(
"$.{}",
serde_json::to_string(&expected.namespace).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.upstream_metadata namespace is not serializable: {err}"
))
})?
);
let expected_value = expected
.expected_value
.as_ref()
.map(serde_json::to_string)
.transpose()
.map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.upstream_metadata expected value is not serializable: {err}"
))
})?;
builder
.push(" AND JSON_TYPE(COALESCE(NULLIF(upstream_metadata, ''), '{}')) = 'OBJECT'")
.push(" AND JSON_EXTRACT(COALESCE(NULLIF(upstream_metadata, ''), '{}'), ")
.push_bind(namespace_path)
.push(") <=> CAST(")
.push_bind(expected_value)
.push(" AS JSON)");
}
let rows_affected = builder let rows_affected = builder
.build() .build()
.execute(&self.pool) .execute(&self.pool)
@@ -1339,6 +1448,25 @@ WHERE id = ?
"provider catalog OAuth credential fence must not contain empty fields".to_string(), "provider catalog OAuth credential fence must not contain empty fields".to_string(),
)); ));
} }
if update
.expected_upstream_metadata_namespace
.as_ref()
.is_some_and(|expected| expected.namespace.trim().is_empty())
{
return Err(DataLayerError::InvalidInput(
"provider catalog OAuth runtime metadata namespace must not be empty".to_string(),
));
}
if update
.upstream_metadata_namespace_to_remove
.as_deref()
.is_some_and(|namespace| namespace.trim().is_empty())
{
return Err(DataLayerError::InvalidInput(
"provider catalog OAuth runtime metadata namespace to remove must not be empty"
.to_string(),
));
}
if !update.status_snapshot_patch.is_object() { if !update.status_snapshot_patch.is_object() {
return Err(DataLayerError::InvalidInput( return Err(DataLayerError::InvalidInput(
"provider catalog status snapshot patch must be an object".to_string(), "provider catalog status snapshot patch must be an object".to_string(),
@@ -1353,6 +1481,22 @@ WHERE id = ?
"provider catalog upstream metadata patch must be an object".to_string(), "provider catalog upstream metadata patch must be an object".to_string(),
)); ));
} }
if update
.upstream_metadata_namespace_to_remove
.as_ref()
.is_some_and(|namespace| {
update
.upstream_metadata_patch
.as_ref()
.and_then(serde_json::Value::as_object)
.is_some_and(|patch| patch.contains_key(namespace))
})
{
return Err(DataLayerError::InvalidInput(
"provider catalog OAuth runtime metadata namespace cannot be patched and removed in the same update"
.to_string(),
));
}
let mut builder = let mut builder =
QueryBuilder::<MySql>::new("UPDATE provider_api_keys SET oauth_invalid_at = "); QueryBuilder::<MySql>::new("UPDATE provider_api_keys SET oauth_invalid_at = ");
builder builder
@@ -1375,9 +1519,29 @@ WHERE id = ?
"provider_api_keys.expires_at", "provider_api_keys.expires_at",
)?); )?);
} }
if let Some(metadata_patch) = update.upstream_metadata_patch.as_ref() { if update.upstream_metadata_patch.is_some()
|| update.upstream_metadata_namespace_to_remove.is_some()
{
builder.push(", upstream_metadata = "); builder.push(", upstream_metadata = ");
push_upstream_metadata_shallow_patch(&mut builder, metadata_patch)?; if update.upstream_metadata_namespace_to_remove.is_some() {
builder.push("JSON_REMOVE(");
}
if let Some(metadata_patch) = update.upstream_metadata_patch.as_ref() {
push_upstream_metadata_shallow_patch(&mut builder, metadata_patch)?;
} else {
builder.push("COALESCE(NULLIF(upstream_metadata, ''), '{}')");
}
if let Some(namespace) = update.upstream_metadata_namespace_to_remove.as_ref() {
let namespace_path = format!(
"$.{}",
serde_json::to_string(namespace).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.upstream_metadata namespace is not serializable: {err}"
))
})?
);
builder.push(", ").push_bind(namespace_path).push(")");
}
} }
builder.push(", status_snapshot = "); builder.push(", status_snapshot = ");
push_status_snapshot_shallow_patch(&mut builder, &update.status_snapshot_patch)?; push_status_snapshot_shallow_patch(&mut builder, &update.status_snapshot_patch)?;
@@ -1411,6 +1575,39 @@ WHERE id = ?
.push_bind(&expected.provider_type) .push_bind(&expected.provider_type)
.push(")"); .push(")");
} }
if update.expected_upstream_metadata_namespace.is_some()
|| update.upstream_metadata_patch.is_some()
|| update.upstream_metadata_namespace_to_remove.is_some()
{
builder
.push(" AND JSON_TYPE(COALESCE(NULLIF(upstream_metadata, ''), '{}')) = 'OBJECT'");
}
if let Some(expected) = update.expected_upstream_metadata_namespace.as_ref() {
let namespace_path = format!(
"$.{}",
serde_json::to_string(&expected.namespace).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.upstream_metadata namespace is not serializable: {err}"
))
})?
);
let expected_value = expected
.expected_value
.as_ref()
.map(serde_json::to_string)
.transpose()
.map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.upstream_metadata expected value is not serializable: {err}"
))
})?;
builder
.push(" AND JSON_EXTRACT(COALESCE(NULLIF(upstream_metadata, ''), '{}'), ")
.push_bind(namespace_path)
.push(") <=> CAST(")
.push_bind(expected_value)
.push(" AS JSON)");
}
let rows_affected = builder let rows_affected = builder
.build() .build()
.execute(&self.pool) .execute(&self.pool)
@@ -1618,6 +1815,7 @@ WHERE id = ?
) )
.push(" WHERE id = ") .push(" WHERE id = ")
.push_bind(&update.key_id) .push_bind(&update.key_id)
.push(" AND JSON_TYPE(COALESCE(NULLIF(upstream_metadata, ''), '{}')) = 'OBJECT'")
.push(" AND JSON_EXTRACT(COALESCE(NULLIF(upstream_metadata, ''), '{}'), ") .push(" AND JSON_EXTRACT(COALESCE(NULLIF(upstream_metadata, ''), '{}'), ")
.push_bind(namespace_path) .push_bind(namespace_path)
.push(") <=> CAST(") .push(") <=> CAST(")
@@ -1897,6 +2095,13 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
Self::update_key(self, key).await Self::update_key(self, key).await
} }
async fn compare_and_update_key_admin_state(
&self,
update: &ProviderCatalogKeyAdminCasUpdate,
) -> Result<bool, DataLayerError> {
Self::compare_and_update_key_admin_state(self, update).await
}
async fn update_keys( async fn update_keys(
&self, &self,
keys: &[StoredProviderCatalogKey], keys: &[StoredProviderCatalogKey],
@@ -2426,6 +2631,10 @@ SET
model_exclude_patterns = ?, model_exclude_patterns = ?,
updated_at = ? updated_at = ?
WHERE id = ? WHERE id = ?
AND BINARY provider_id = BINARY ?
AND BINARY auth_type = BINARY ?
AND BINARY api_key <=> BINARY ?
AND BINARY auth_config <=> BINARY ?
"# "#
} }
@@ -2500,7 +2709,152 @@ fn key_update_query(
"provider_api_keys.model_exclude_patterns", "provider_api_keys.model_exclude_patterns",
)?) )?)
.bind(updated_at) .bind(updated_at)
.bind(&key.id)) .bind(&key.id)
.bind(&key.provider_id)
.bind(&key.auth_type)
.bind(&key.encrypted_api_key)
.bind(&key.encrypted_auth_config))
}
fn push_admin_key_assignments<'args>(
builder: &mut QueryBuilder<'args, MySql>,
key: &'args StoredProviderCatalogKey,
updated_at: i64,
) -> Result<(), DataLayerError> {
builder
.push_bind(&key.provider_id)
.push(", name = ")
.push_bind(&key.name)
.push(", api_key = ")
.push_bind(&key.encrypted_api_key)
.push(", auth_type = ")
.push_bind(&key.auth_type)
.push(", capabilities = ")
.push_bind(optional_json_to_string(
&key.capabilities,
"provider_api_keys.capabilities",
)?)
.push(", is_active = ")
.push_bind(key.is_active)
.push(", api_formats = ")
.push_bind(optional_json_to_string(
&key.api_formats,
"provider_api_keys.api_formats",
)?)
.push(", auth_type_by_format = ")
.push_bind(optional_json_to_string(
&key.auth_type_by_format,
"provider_api_keys.auth_type_by_format",
)?)
.push(", allow_auth_channel_mismatch_formats = ")
.push_bind(optional_json_to_string(
&key.allow_auth_channel_mismatch_formats,
"provider_api_keys.allow_auth_channel_mismatch_formats",
)?)
.push(", auth_config = ")
.push_bind(&key.encrypted_auth_config)
.push(", note = ")
.push_bind(&key.note)
.push(", internal_priority = ")
.push_bind(key.internal_priority)
.push(", rate_multipliers = ")
.push_bind(optional_json_to_string(
&key.rate_multipliers,
"provider_api_keys.rate_multipliers",
)?)
.push(", global_priority_by_format = ")
.push_bind(optional_json_to_string(
&key.global_priority_by_format,
"provider_api_keys.global_priority_by_format",
)?)
.push(", allowed_models = ")
.push_bind(optional_json_to_string(
&key.allowed_models,
"provider_api_keys.allowed_models",
)?)
.push(", expires_at = ")
.push_bind(optional_i64_from_u64(
key.expires_at_unix_secs,
"provider_api_keys.expires_at",
)?)
.push(", cache_ttl_minutes = ")
.push_bind(key.cache_ttl_minutes)
.push(", max_probe_interval_minutes = ")
.push_bind(key.max_probe_interval_minutes)
.push(", proxy = ")
.push_bind(optional_json_to_string(
&key.proxy,
"provider_api_keys.proxy",
)?)
.push(", fingerprint = ")
.push_bind(optional_json_to_string(
&key.fingerprint,
"provider_api_keys.fingerprint",
)?)
.push(", rpm_limit = ")
.push_bind(optional_i64_from_u32(key.rpm_limit))
.push(", concurrent_limit = ")
.push_bind(key.concurrent_limit)
.push(", auto_fetch_models = ")
.push_bind(key.auto_fetch_models)
.push(", locked_models = ")
.push_bind(optional_json_to_string(
&key.locked_models,
"provider_api_keys.locked_models",
)?)
.push(", model_include_patterns = ")
.push_bind(optional_json_to_string(
&key.model_include_patterns,
"provider_api_keys.model_include_patterns",
)?)
.push(", model_exclude_patterns = ")
.push_bind(optional_json_to_string(
&key.model_exclude_patterns,
"provider_api_keys.model_exclude_patterns",
)?)
.push(", updated_at = ")
.push_bind(updated_at);
Ok(())
}
fn validate_admin_key_cas_update(
update: &ProviderCatalogKeyAdminCasUpdate,
) -> Result<(), DataLayerError> {
validate_key(&update.key)?;
let expected = &update.expected_credential;
if expected.auth_type.trim().is_empty()
|| expected.provider_id.trim().is_empty()
|| expected.provider_type.trim().is_empty()
|| expected
.encrypted_api_key
.as_deref()
.is_some_and(|value| value.trim().is_empty())
|| update
.expected_encrypted_auth_config
.as_deref()
.is_some_and(|value| value.trim().is_empty())
{
return Err(DataLayerError::InvalidInput(
"provider catalog admin credential fence contains empty fields".to_string(),
));
}
let Some(rotation) = update.codex_rotation.as_ref() else {
return Ok(());
};
let valid_rotation = expected.provider_type.eq_ignore_ascii_case("codex")
&& rotation.as_object().is_some_and(|object| {
object.len() == 1
&& object
.get("credential_generation")
.and_then(serde_json::Value::as_str)
.is_some_and(|generation| !generation.trim().is_empty())
});
if !valid_rotation {
return Err(DataLayerError::InvalidInput(
"provider catalog Codex rotation must contain only credential_generation".to_string(),
));
}
Ok(())
} }
fn optional_json_from_string( fn optional_json_from_string(
@@ -2880,6 +3234,34 @@ mod tests {
] { ] {
assert!(!sql.contains(runtime_assignment)); assert!(!sql.contains(runtime_assignment));
} }
assert!(sql.contains("binary api_key <=> binary ?"));
assert!(sql.contains("binary auth_config <=> binary ?"));
}
#[test]
fn admin_credential_cas_has_atomic_rotation_guards() {
let source = include_str!("provider_catalog.rs");
for predicate in [
"BINARY api_key <=> BINARY ",
"BINARY auth_config <=> BINARY ",
"JSON_TYPE(COALESCE(upstream_metadata, JSON_OBJECT())) = 'OBJECT'",
"JSON_TYPE(COALESCE(status_snapshot, JSON_OBJECT())) = 'OBJECT'",
"JSON_SET(COALESCE(upstream_metadata, JSON_OBJECT()), '$.codex'",
"oauth_invalid_at = NULL, oauth_invalid_reason = NULL, error_count = 0",
] {
assert!(
source.contains(predicate),
"missing admin CAS guard: {predicate}"
);
}
}
#[test]
fn runtime_metadata_cas_requires_an_object_metadata_root() {
let source = include_str!("provider_catalog.rs");
assert!(
source.contains("JSON_TYPE(COALESCE(NULLIF(upstream_metadata, ''), '{}')) = 'OBJECT'")
);
} }
#[tokio::test] #[tokio::test]
async fn empty_id_lists_do_not_connect_to_lazy_pool() { async fn empty_id_lists_do_not_connect_to_lazy_pool() {
@@ -9,8 +9,8 @@ use sqlx::{
}; };
use aether_data_contracts::repository::provider_catalog::{ use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdminCasUpdate,
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
ProviderCatalogReadRepository, ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogReadRepository, ProviderCatalogUpstreamMetadataNamespaceUpdate,
@@ -380,6 +380,10 @@ SET
auth_type_by_format = $27, auth_type_by_format = $27,
allow_auth_channel_mismatch_formats = $28 allow_auth_channel_mismatch_formats = $28
WHERE id = $1 WHERE id = $1
AND provider_id = $2
AND auth_type = $4
AND api_key IS NOT DISTINCT FROM $5
AND auth_config IS NOT DISTINCT FROM $6
"#; "#;
const KEY_RUNTIME_HEALTH_CAS_SQL: &str = r#" const KEY_RUNTIME_HEALTH_CAS_SQL: &str = r#"
@@ -405,6 +409,7 @@ SET
ELSE TO_TIMESTAMP($5::double precision) ELSE TO_TIMESTAMP($5::double precision)
END END
WHERE id = $1 WHERE id = $1
AND jsonb_typeof(COALESCE(upstream_metadata, '{}'::jsonb)) = 'object'
AND (COALESCE(upstream_metadata, '{}'::jsonb) -> $2) AND (COALESCE(upstream_metadata, '{}'::jsonb) -> $2)
IS NOT DISTINCT FROM $6::jsonb IS NOT DISTINCT FROM $6::jsonb
"#; "#;
@@ -500,6 +505,112 @@ fn key_update_query(key: &StoredProviderCatalogKey) -> Query<'_, Postgres, PgArg
.bind(&key.allow_auth_channel_mismatch_formats) .bind(&key.allow_auth_channel_mismatch_formats)
} }
fn push_admin_key_assignments<'args>(
builder: &mut QueryBuilder<'args, Postgres>,
key: &'args StoredProviderCatalogKey,
) {
builder
.push_bind(&key.provider_id)
.push(", api_formats = ")
.push_bind(&key.api_formats)
.push(", auth_type = ")
.push_bind(&key.auth_type)
.push(", api_key = ")
.push_bind(&key.encrypted_api_key)
.push(", auth_config = ")
.push_bind(&key.encrypted_auth_config)
.push(", name = ")
.push_bind(&key.name)
.push(", note = ")
.push_bind(&key.note)
.push(", rate_multipliers = ")
.push_bind(&key.rate_multipliers)
.push(", internal_priority = ")
.push_bind(key.internal_priority)
.push(", global_priority_by_format = ")
.push_bind(&key.global_priority_by_format)
.push(", rpm_limit = ")
.push_bind(key.rpm_limit.map(|value| value as i32))
.push(", concurrent_limit = ")
.push_bind(key.concurrent_limit)
.push(", allowed_models = ")
.push_bind(&key.allowed_models)
.push(", capabilities = ")
.push_bind(&key.capabilities)
.push(", cache_ttl_minutes = ")
.push_bind(key.cache_ttl_minutes)
.push(", max_probe_interval_minutes = ")
.push_bind(key.max_probe_interval_minutes)
.push(", auto_fetch_models = ")
.push_bind(key.auto_fetch_models)
.push(", locked_models = ")
.push_bind(&key.locked_models)
.push(", model_include_patterns = ")
.push_bind(&key.model_include_patterns)
.push(", model_exclude_patterns = ")
.push_bind(&key.model_exclude_patterns)
.push(", proxy = ")
.push_bind(&key.proxy)
.push(", fingerprint = ")
.push_bind(&key.fingerprint)
.push(", expires_at = CASE WHEN ")
.push_bind(key.expires_at_unix_secs.map(|value| value as f64))
.push("::double precision IS NULL THEN NULL ELSE TO_TIMESTAMP(")
.push_bind(key.expires_at_unix_secs.map(|value| value as f64))
.push("::double precision) END, is_active = ")
.push_bind(key.is_active)
.push(", updated_at = CASE WHEN ")
.push_bind(key.updated_at_unix_secs.map(|value| value as f64))
.push("::double precision IS NULL THEN NOW() ELSE TO_TIMESTAMP(")
.push_bind(key.updated_at_unix_secs.map(|value| value as f64))
.push("::double precision) END, auth_type_by_format = ")
.push_bind(&key.auth_type_by_format)
.push(", allow_auth_channel_mismatch_formats = ")
.push_bind(&key.allow_auth_channel_mismatch_formats);
}
fn validate_admin_key_cas_update(
update: &ProviderCatalogKeyAdminCasUpdate,
) -> Result<(), DataLayerError> {
validate_key_for_update(&update.key)?;
let expected = &update.expected_credential;
if update.key.name.trim().is_empty()
|| update.key.auth_type.trim().is_empty()
|| expected.auth_type.trim().is_empty()
|| expected.provider_id.trim().is_empty()
|| expected.provider_type.trim().is_empty()
|| expected
.encrypted_api_key
.as_deref()
.is_some_and(|value| value.trim().is_empty())
|| update
.expected_encrypted_auth_config
.as_deref()
.is_some_and(|value| value.trim().is_empty())
{
return Err(DataLayerError::InvalidInput(
"provider catalog admin credential fence contains empty fields".to_string(),
));
}
let Some(rotation) = update.codex_rotation.as_ref() else {
return Ok(());
};
let valid_rotation = expected.provider_type.eq_ignore_ascii_case("codex")
&& rotation.as_object().is_some_and(|object| {
object.len() == 1
&& object
.get("credential_generation")
.and_then(serde_json::Value::as_str)
.is_some_and(|generation| !generation.trim().is_empty())
});
if !valid_rotation {
return Err(DataLayerError::InvalidInput(
"provider catalog Codex rotation must contain only credential_generation".to_string(),
));
}
Ok(())
}
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct SqlxProviderCatalogReadRepository { pub struct SqlxProviderCatalogReadRepository {
pool: PgPool, pool: PgPool,
@@ -901,11 +1012,29 @@ WHERE id = $1
|| expected.provider_id.trim().is_empty() || expected.provider_id.trim().is_empty()
|| expected.provider_type.trim().is_empty() || expected.provider_type.trim().is_empty()
}) })
|| update
.expected_upstream_metadata_namespace
.as_ref()
.is_some_and(|expected| expected.namespace.trim().is_empty())
|| update
.upstream_metadata_namespace_to_remove
.as_deref()
.is_some_and(|namespace| namespace.trim().is_empty())
|| !update.status_snapshot_patch.is_object() || !update.status_snapshot_patch.is_object()
|| update || update
.upstream_metadata_patch .upstream_metadata_patch
.as_ref() .as_ref()
.is_some_and(|patch| !patch.is_object()) .is_some_and(|patch| !patch.is_object())
|| update
.upstream_metadata_namespace_to_remove
.as_ref()
.is_some_and(|namespace| {
update
.upstream_metadata_patch
.as_ref()
.and_then(serde_json::Value::as_object)
.is_some_and(|patch| patch.contains_key(namespace))
})
{ {
return Err(DataLayerError::InvalidInput( return Err(DataLayerError::InvalidInput(
"provider catalog OAuth runtime CAS requires key_id, auth_config, and object status patch" "provider catalog OAuth runtime CAS requires key_id, auth_config, and object status patch"
@@ -932,27 +1061,40 @@ SET
ELSE TO_TIMESTAMP($7::double precision) ELSE TO_TIMESTAMP($7::double precision)
END, END,
upstream_metadata = CASE upstream_metadata = CASE
WHEN $8::jsonb IS NULL THEN upstream_metadata WHEN $8::jsonb IS NULL AND $9::text IS NULL THEN upstream_metadata
ELSE COALESCE(upstream_metadata, '{}'::jsonb) || $8::jsonb WHEN $9::text IS NULL THEN COALESCE(upstream_metadata, '{}'::jsonb) || $8::jsonb
ELSE (COALESCE(upstream_metadata, '{}'::jsonb) || COALESCE($8::jsonb, '{}'::jsonb)) - $9
END, END,
status_snapshot = (COALESCE(status_snapshot::jsonb, '{}'::jsonb) || $9::jsonb)::json, status_snapshot = (COALESCE(status_snapshot::jsonb, '{}'::jsonb) || $10::jsonb)::json,
error_count = CASE WHEN $10::boolean THEN 0 ELSE error_count END, error_count = CASE WHEN $11::boolean THEN 0 ELSE error_count END,
updated_at = CASE updated_at = CASE
WHEN $11::double precision IS NULL THEN NOW() WHEN $12::double precision IS NULL THEN NOW()
ELSE TO_TIMESTAMP($11::double precision) ELSE TO_TIMESTAMP($12::double precision)
END END
WHERE id = $1 WHERE id = $1
AND auth_config IS NOT DISTINCT FROM $12 AND auth_config IS NOT DISTINCT FROM $13
AND ($13::boolean IS FALSE OR api_key IS NOT DISTINCT FROM $14)
AND ($15::text IS NULL OR auth_type = $15)
AND ($16::text IS NULL OR provider_id = $16)
AND ( AND (
$17::text IS NULL ($8::jsonb IS NULL AND $9::text IS NULL)
OR jsonb_typeof(COALESCE(upstream_metadata, '{}'::jsonb)) = 'object'
)
AND ($14::boolean IS FALSE OR api_key IS NOT DISTINCT FROM $15)
AND ($16::text IS NULL OR auth_type = $16)
AND ($17::text IS NULL OR provider_id = $17)
AND (
$18::text IS NULL
OR EXISTS ( OR EXISTS (
SELECT 1 SELECT 1
FROM providers FROM providers
WHERE providers.id = provider_api_keys.provider_id WHERE providers.id = provider_api_keys.provider_id
AND providers.provider_type = $17 AND providers.provider_type = $18
)
)
AND (
$19::boolean IS FALSE
OR (
jsonb_typeof(COALESCE(upstream_metadata, '{}'::jsonb)) = 'object'
AND (COALESCE(upstream_metadata, '{}'::jsonb) -> $20)
IS NOT DISTINCT FROM $21::jsonb
) )
) )
"#, "#,
@@ -970,6 +1112,7 @@ WHERE id = $1
.map(|value| value as f64), .map(|value| value as f64),
) )
.bind(update.upstream_metadata_patch.as_ref()) .bind(update.upstream_metadata_patch.as_ref())
.bind(update.upstream_metadata_namespace_to_remove.as_deref())
.bind(&update.status_snapshot_patch) .bind(&update.status_snapshot_patch)
.bind(update.reset_error_count) .bind(update.reset_error_count)
.bind(update.updated_at_unix_secs.map(|value| value as f64)) .bind(update.updated_at_unix_secs.map(|value| value as f64))
@@ -999,6 +1142,19 @@ WHERE id = $1
.as_ref() .as_ref()
.map(|expected| expected.provider_type.as_str()), .map(|expected| expected.provider_type.as_str()),
) )
.bind(update.expected_upstream_metadata_namespace.is_some())
.bind(
update
.expected_upstream_metadata_namespace
.as_ref()
.map(|expected| expected.namespace.as_str()),
)
.bind(
update
.expected_upstream_metadata_namespace
.as_ref()
.and_then(|expected| expected.expected_value.as_ref()),
)
.execute(&self.pool) .execute(&self.pool)
.await .await
.map_postgres_err()? .map_postgres_err()?
@@ -2008,6 +2164,76 @@ WHERE id = $1
}) })
} }
pub async fn compare_and_update_key_admin_state(
&self,
update: &ProviderCatalogKeyAdminCasUpdate,
) -> Result<bool, DataLayerError> {
validate_admin_key_cas_update(update)?;
let key = &update.key;
let mut builder =
QueryBuilder::<Postgres>::new("UPDATE provider_api_keys SET provider_id = ");
push_admin_key_assignments(&mut builder, key);
if update.reset_oauth_runtime {
builder.push(", oauth_invalid_at = NULL, oauth_invalid_reason = NULL, error_count = 0");
}
if let Some(rotation) = update.codex_rotation.as_ref() {
builder
.push(", upstream_metadata = jsonb_set(COALESCE(upstream_metadata, '{}'::jsonb), '{codex}', ")
.push_bind(rotation)
.push("::jsonb, true)");
}
if update.codex_rotation.is_some() || update.reset_oauth_runtime {
builder.push(", status_snapshot = ");
match (update.codex_rotation.is_some(), update.reset_oauth_runtime) {
(true, true) => builder.push("jsonb_set(jsonb_set(COALESCE(status_snapshot::jsonb, '{}'::jsonb), '{quota}', 'null'::jsonb, true), '{oauth}', 'null'::jsonb, true)::json"),
(true, false) => builder.push("jsonb_set(COALESCE(status_snapshot::jsonb, '{}'::jsonb), '{quota}', 'null'::jsonb, true)::json"),
(false, true) => builder.push("jsonb_set(COALESCE(status_snapshot::jsonb, '{}'::jsonb), '{oauth}', 'null'::jsonb, true)::json"),
(false, false) => unreachable!(),
};
}
builder
.push(" WHERE id = ")
.push_bind(&key.id)
.push(" AND api_key IS NOT DISTINCT FROM ")
.push_bind(update.expected_credential.encrypted_api_key.as_deref())
.push(" AND auth_config IS NOT DISTINCT FROM ")
.push_bind(update.expected_encrypted_auth_config.as_deref())
.push(" AND auth_type = ")
.push_bind(&update.expected_credential.auth_type)
.push(" AND provider_id = ")
.push_bind(&update.expected_credential.provider_id)
.push(
" AND EXISTS (SELECT 1 FROM providers WHERE providers.id = provider_api_keys.provider_id AND providers.provider_type = ",
)
.push_bind(&update.expected_credential.provider_type)
.push(")");
if update.codex_rotation.is_some() {
builder
.push(" AND jsonb_typeof(COALESCE(upstream_metadata, '{}'::jsonb)) = 'object'")
.push(" AND NOT (api_key IS NOT DISTINCT FROM ")
.push_bind(key.encrypted_api_key.as_deref())
.push(" AND auth_config IS NOT DISTINCT FROM ")
.push_bind(key.encrypted_auth_config.as_deref())
.push(" AND auth_type = ")
.push_bind(&key.auth_type)
.push(" AND provider_id = ")
.push_bind(&key.provider_id)
.push(")");
}
if update.codex_rotation.is_some() || update.reset_oauth_runtime {
builder.push(
" AND jsonb_typeof(COALESCE(status_snapshot::jsonb, '{}'::jsonb)) = 'object'",
);
}
let rows_affected = builder
.build()
.execute(&self.pool)
.await
.map_postgres_err()?
.rows_affected();
Ok(rows_affected > 0)
}
pub async fn update_keys( pub async fn update_keys(
&self, &self,
keys: &[StoredProviderCatalogKey], keys: &[StoredProviderCatalogKey],
@@ -2088,6 +2314,10 @@ WHERE id = $1
|| expected.auth_type.trim().is_empty() || expected.auth_type.trim().is_empty()
|| expected.provider_id.trim().is_empty() || expected.provider_id.trim().is_empty()
|| expected.provider_type.trim().is_empty() || expected.provider_type.trim().is_empty()
|| delete
.expected_upstream_metadata_namespace
.as_ref()
.is_some_and(|expected| expected.namespace.trim().is_empty())
{ {
return Err(DataLayerError::InvalidInput( return Err(DataLayerError::InvalidInput(
"provider catalog OAuth credential CAS delete contains empty fields".to_string(), "provider catalog OAuth credential CAS delete contains empty fields".to_string(),
@@ -2107,6 +2337,14 @@ WHERE id = $1
WHERE providers.id = provider_api_keys.provider_id WHERE providers.id = provider_api_keys.provider_id
AND providers.provider_type = $6 AND providers.provider_type = $6
) )
AND (
$7::boolean IS FALSE
OR (
jsonb_typeof(COALESCE(upstream_metadata, '{}'::jsonb)) = 'object'
AND (COALESCE(upstream_metadata, '{}'::jsonb) -> $8)
IS NOT DISTINCT FROM $9::jsonb
)
)
"#, "#,
) )
.bind(&delete.key_id) .bind(&delete.key_id)
@@ -2115,6 +2353,19 @@ WHERE id = $1
.bind(&expected.auth_type) .bind(&expected.auth_type)
.bind(&expected.provider_id) .bind(&expected.provider_id)
.bind(&expected.provider_type) .bind(&expected.provider_type)
.bind(delete.expected_upstream_metadata_namespace.is_some())
.bind(
delete
.expected_upstream_metadata_namespace
.as_ref()
.map(|expected| expected.namespace.as_str()),
)
.bind(
delete
.expected_upstream_metadata_namespace
.as_ref()
.and_then(|expected| expected.expected_value.as_ref()),
)
.execute(&self.pool) .execute(&self.pool)
.await .await
.map_postgres_err()? .map_postgres_err()?
@@ -2663,6 +2914,13 @@ impl ProviderCatalogWriteRepository for SqlxProviderCatalogReadRepository {
Self::update_key(self, key).await Self::update_key(self, key).await
} }
async fn compare_and_update_key_admin_state(
&self,
update: &ProviderCatalogKeyAdminCasUpdate,
) -> Result<bool, DataLayerError> {
Self::compare_and_update_key_admin_state(self, update).await
}
async fn update_keys( async fn update_keys(
&self, &self,
keys: &[StoredProviderCatalogKey], keys: &[StoredProviderCatalogKey],
@@ -3407,6 +3665,7 @@ mod tests {
fn runtime_metadata_cas_compares_only_the_requested_namespace() { fn runtime_metadata_cas_compares_only_the_requested_namespace() {
let sql = super::KEY_RUNTIME_METADATA_CAS_SQL.to_ascii_lowercase(); let sql = super::KEY_RUNTIME_METADATA_CAS_SQL.to_ascii_lowercase();
assert!(sql.contains("upstream_metadata, '{}'::jsonb) -> $2")); assert!(sql.contains("upstream_metadata, '{}'::jsonb) -> $2"));
assert!(sql.contains("jsonb_typeof(coalesce(upstream_metadata, '{}'::jsonb)) = 'object'"));
assert!(sql.contains("is not distinct from $6::jsonb")); assert!(sql.contains("is not distinct from $6::jsonb"));
assert!(sql.contains("status_snapshot::jsonb")); assert!(sql.contains("status_snapshot::jsonb"));
assert!(!sql.contains("is_active")); assert!(!sql.contains("is_active"));
@@ -3436,5 +3695,26 @@ mod tests {
} }
assert!(sql.contains("is_active = $25")); assert!(sql.contains("is_active = $25"));
assert!(sql.contains("rpm_limit = $12")); assert!(sql.contains("rpm_limit = $12"));
assert!(sql.contains("api_key is not distinct from $5"));
assert!(sql.contains("auth_config is not distinct from $6"));
}
#[test]
fn admin_credential_cas_has_atomic_rotation_guards() {
let source = include_str!("provider_catalog.rs");
for predicate in [
"api_key IS NOT DISTINCT FROM ",
"auth_config IS NOT DISTINCT FROM ",
"jsonb_typeof(COALESCE(upstream_metadata, '{}'::jsonb)) = 'object'",
"jsonb_typeof(COALESCE(status_snapshot::jsonb, '{}'::jsonb)) = 'object'",
"jsonb_set(COALESCE(upstream_metadata, '{}'::jsonb), '{codex}'",
"jsonb_set(COALESCE(status_snapshot::jsonb, '{}'::jsonb), '{quota}', 'null'::jsonb, true)::json",
"oauth_invalid_at = NULL, oauth_invalid_reason = NULL, error_count = 0",
] {
assert!(
source.contains(predicate),
"missing admin CAS guard: {predicate}"
);
}
} }
} }
File diff suppressed because it is too large Load Diff
@@ -4,10 +4,12 @@ mod types;
pub use snapshot::ProviderCatalogSnapshot; pub use snapshot::ProviderCatalogSnapshot;
pub use types::{ pub use types::{
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate,
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthCredentialFence, ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthCredentialFence,
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository,
ProviderCatalogUpstreamMetadataNamespaceExpectation,
ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository, ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage, StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
@@ -79,13 +79,50 @@ pub struct ProviderCatalogKeyOAuthCredentialFence {
pub provider_type: String, pub provider_type: String,
} }
/// Administrator-owned key replacement fenced by the exact credential state
/// observed while the edit was prepared. This prevents an older admin request
/// from restoring credentials that a concurrent request already replaced.
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct ProviderCatalogKeyAdminCasUpdate {
pub expected_encrypted_auth_config: Option<String>,
pub expected_credential: ProviderCatalogKeyOAuthCredentialFence,
/// Full requested key. Repositories merge its administrator-owned fields
/// while preserving the currently stored runtime-owned fields.
pub key: StoredProviderCatalogKey,
/// Optional replacement for the complete Codex metadata namespace. When
/// present it must be an object containing only a non-empty string
/// `credential_generation`; repositories also clear the quota snapshot in
/// the same atomic write.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub codex_rotation: Option<serde_json::Value>,
/// Clear OAuth invalid markers in the same credential replacement write.
/// This must be true whenever new credential material supersedes the old
/// credential, so a post-CAS unfenced cleanup is never required.
#[serde(default)]
pub reset_oauth_runtime: bool,
}
/// Atomic key deletion fenced by the exact OAuth credential generation that /// Atomic key deletion fenced by the exact OAuth credential generation that
/// produced the terminal failure. /// produced the terminal failure and, when supplied, one metadata namespace.
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct ProviderCatalogKeyOAuthCredentialCasDelete { pub struct ProviderCatalogKeyOAuthCredentialCasDelete {
pub key_id: String, pub key_id: String,
pub expected_encrypted_auth_config: Option<String>, pub expected_encrypted_auth_config: Option<String>,
pub expected_credential: ProviderCatalogKeyOAuthCredentialFence, pub expected_credential: ProviderCatalogKeyOAuthCredentialFence,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub expected_upstream_metadata_namespace:
Option<ProviderCatalogUpstreamMetadataNamespaceExpectation>,
}
/// Optional single-namespace metadata fence for an OAuth runtime CAS.
///
/// The outer option on the owning update controls whether the namespace is
/// compared. Within an expectation, `None` requires the namespace to be absent.
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct ProviderCatalogUpstreamMetadataNamespaceExpectation {
pub namespace: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub expected_value: Option<serde_json::Value>,
} }
/// Agent/runtime-owned OAuth state update fenced by the exact encrypted /// Agent/runtime-owned OAuth state update fenced by the exact encrypted
@@ -98,6 +135,11 @@ pub struct ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
pub expected_encrypted_auth_config: Option<String>, pub expected_encrypted_auth_config: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")] #[serde(default, skip_serializing_if = "Option::is_none")]
pub expected_credential: Option<ProviderCatalogKeyOAuthCredentialFence>, pub expected_credential: Option<ProviderCatalogKeyOAuthCredentialFence>,
/// Optional metadata namespace value compared atomically with the OAuth
/// credential fence. `expected_value: None` means the namespace must be absent.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub expected_upstream_metadata_namespace:
Option<ProviderCatalogUpstreamMetadataNamespaceExpectation>,
pub encrypted_auth_config: String, pub encrypted_auth_config: String,
/// Optional access-token ciphertext replacement owned by refresh success. /// Optional access-token ciphertext replacement owned by refresh success.
#[serde(default, skip_serializing_if = "Option::is_none")] #[serde(default, skip_serializing_if = "Option::is_none")]
@@ -110,6 +152,10 @@ pub struct ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
/// Top-level runtime metadata namespaces to merge in the same fenced write. /// Top-level runtime metadata namespaces to merge in the same fenced write.
#[serde(default, skip_serializing_if = "Option::is_none")] #[serde(default, skip_serializing_if = "Option::is_none")]
pub upstream_metadata_patch: Option<serde_json::Value>, pub upstream_metadata_patch: Option<serde_json::Value>,
/// Optional top-level runtime metadata namespace to remove in the same
/// fenced write. It must not also appear in `upstream_metadata_patch`.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub upstream_metadata_namespace_to_remove: Option<String>,
pub status_snapshot_patch: serde_json::Value, pub status_snapshot_patch: serde_json::Value,
#[serde(default)] #[serde(default)]
pub reset_error_count: bool, pub reset_error_count: bool,
@@ -822,6 +868,18 @@ pub trait ProviderCatalogWriteRepository: Send + Sync {
key: &StoredProviderCatalogKey, key: &StoredProviderCatalogKey,
) -> Result<StoredProviderCatalogKey, crate::DataLayerError>; ) -> Result<StoredProviderCatalogKey, crate::DataLayerError>;
/// Compare-and-swap administrator-owned key configuration. Credential
/// rotation, Codex namespace replacement, and quota invalidation must be
/// committed atomically with the configuration update.
async fn compare_and_update_key_admin_state(
&self,
_update: &ProviderCatalogKeyAdminCasUpdate,
) -> Result<bool, crate::DataLayerError> {
Err(crate::DataLayerError::InvalidConfiguration(
"provider catalog admin CAS updates are not supported by this repository".to_string(),
))
}
async fn update_keys( async fn update_keys(
&self, &self,
keys: &[StoredProviderCatalogKey], keys: &[StoredProviderCatalogKey],
File diff suppressed because it is too large Load Diff
@@ -3,10 +3,12 @@ mod memory;
#[allow(unused_imports)] #[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::provider_catalog::{ pub(crate) use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate,
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthCredentialFence, ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthCredentialFence,
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot, ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot,
ProviderCatalogUpstreamMetadataNamespaceExpectation,
ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository, ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage, StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
@@ -28,6 +28,11 @@ mod tests {
payload: StreamFramePayload::Headers { payload: StreamFramePayload::Headers {
status_code: 200, status_code: 200,
headers: BTreeMap::from([("content-type".into(), "text/event-stream".into())]), headers: BTreeMap::from([("content-type".into(), "text/event-stream".into())]),
response_observation: Some(aether_contracts::ExecutionResponseObservation {
request_started_at_unix_ms: 100,
response_headers_observed_at_unix_ms: 125,
request_order_id: "0198-order-1".to_string(),
}),
}, },
}; };
@@ -1475,6 +1475,7 @@ mod tests {
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code: self.status_code, status_code: self.status_code,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: Some(self.response_body.clone()), json_body: Some(self.response_body.clone()),
body_bytes_b64: None, body_bytes_b64: None,
@@ -1526,6 +1527,7 @@ mod tests {
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code, status_code,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: Some(response_body), json_body: Some(response_body),
body_bytes_b64: None, body_bytes_b64: None,
@@ -1582,6 +1584,7 @@ mod tests {
candidate_id: plan.candidate_id.clone(), candidate_id: plan.candidate_id.clone(),
status_code, status_code,
headers: BTreeMap::new(), headers: BTreeMap::new(),
response_observation: None,
body: Some(ResponseBody { body: Some(ResponseBody {
json_body: Some(response_body), json_body: Some(response_body),
body_bytes_b64: None, body_bytes_b64: None,
+3
View File
@@ -321,6 +321,7 @@ export async function refreshProviderQuota(
export interface ConsumeCodexResetCreditPayload { export interface ConsumeCodexResetCreditPayload {
idempotency_key: string idempotency_key: string
expected_credential_generation: string | null
} }
export interface ConsumeCodexResetCreditResult { export interface ConsumeCodexResetCreditResult {
@@ -331,6 +332,8 @@ export interface ConsumeCodexResetCreditResult {
| 'already_redeemed' | 'already_redeemed'
| 'nothing_to_reset' | 'nothing_to_reset'
| 'no_credit' | 'no_credit'
| 'historical_replay'
| 'credential_changed'
| 'unknown' | 'unknown'
| 'error' | 'error'
| string | string
@@ -323,6 +323,7 @@ export interface EndpointAPIKey {
// Codex 上游元数据类型 // Codex 上游元数据类型
export interface CodexUpstreamMetadata { export interface CodexUpstreamMetadata {
credential_generation?: string
updated_at?: number // 更新时间(Unix 时间戳) updated_at?: number // 更新时间(Unix 时间戳)
plan_type?: string // 套餐类型 plan_type?: string // 套餐类型
primary_used_percent?: number // 周限额窗口使用百分比 primary_used_percent?: number // 周限额窗口使用百分比
@@ -348,6 +349,10 @@ export interface CodexUpstreamMetadata {
has_credits?: boolean // 是否有积分 has_credits?: boolean // 是否有积分
credits_balance?: number // 积分余额 credits_balance?: number // 积分余额
reset_credits?: QuotaResetCreditsSnapshot | null // Codex earned rate-limit reset credits reset_credits?: QuotaResetCreditsSnapshot | null // Codex earned rate-limit reset credits
account_quota_reset_reservation?: {
idempotency_key?: string | null
generation?: number | null
} | null
} }
export interface AntigravityModelQuota { export interface AntigravityModelQuota {
@@ -292,7 +292,7 @@
</div> </div>
</div> </div>
<div <div
v-if="getCodexResetCreditAvailableCount(key) !== null" v-if="getCodexResetCreditAvailableCount(key) !== null || hasPendingCodexResetCredit(key)"
class="mt-3 border-t border-border/60 pt-2" class="mt-3 border-t border-border/60 pt-2"
> >
<div class="flex flex-wrap items-center gap-x-1 gap-y-1 text-[10px] leading-4 text-muted-foreground"> <div class="flex flex-wrap items-center gap-x-1 gap-y-1 text-[10px] leading-4 text-muted-foreground">
@@ -303,7 +303,9 @@
:disabled="consumingCodexResetCreditKeyId === key.id" :disabled="consumingCodexResetCreditKeyId === key.id"
@click="handleConsumeCodexResetCredit(key)" @click="handleConsumeCodexResetCredit(key)"
> >
{{ consumingCodexResetCreditKeyId === key.id ? legacyT('重置中...') : legacyT('点击以进行重置') }} {{ consumingCodexResetCreditKeyId === key.id
? legacyT('重置中...')
: legacyT(hasPendingCodexResetCredit(key) ? '继续确认重置' : '点击以进行重置') }}
</button> </button>
<span <span
v-else v-else
@@ -1044,12 +1046,16 @@ import {
} from '@/utils/providerKeyStatus' } from '@/utils/providerKeyStatus'
import { getGeminiCliAccountCreditsText } from '@/utils/providerKeyQuota' import { getGeminiCliAccountCreditsText } from '@/utils/providerKeyQuota'
import { import {
clearPendingCodexResetCreditIdempotencyKeyForOutcome,
createCodexResetCreditIdempotencyKey, createCodexResetCreditIdempotencyKey,
formatCodexResetCreditCount as formatCodexResetCreditCountLabel, formatCodexResetCreditCount as formatCodexResetCreditCountLabel,
formatCodexResetCreditExpiresAt, formatCodexResetCreditExpiresAt,
getCodexResetCreditAvailableCount as getCodexResetCreditAvailableCountFromSnapshot, getCodexResetCreditAvailableCount as getCodexResetCreditAvailableCountFromSnapshot,
getCodexResetCreditReservationIdempotencyKey,
getVisibleCodexResetCreditItems as getVisibleCodexResetCreditItemsFromSnapshot, getVisibleCodexResetCreditItems as getVisibleCodexResetCreditItemsFromSnapshot,
mergeCodexQuotaDisplays, mergeCodexQuotaDisplays,
readPendingCodexResetCreditIdempotencyKey,
rememberPendingCodexResetCreditIdempotencyKey,
} from './codex-reset-credit-display' } from './codex-reset-credit-display'
// , // ,
@@ -1741,16 +1747,36 @@ function codexResetCreditOutcomeFeedback(
} }
} }
function codexResetCreditActiveIdempotencyKeyFromError(error: unknown): string | null {
if (typeof error !== 'object' || error === null || !('response' in error)) return null
const response = (error as { response?: { data?: unknown } }).response
if (typeof response?.data !== 'object' || response.data === null) return null
const activeKey = (response.data as Record<string, unknown>).active_idempotency_key
return typeof activeKey === 'string' && activeKey.trim() ? activeKey.trim() : null
}
function codexResetCreditCredentialChangedFromError(error: unknown): boolean {
if (typeof error !== 'object' || error === null || !('response' in error)) return false
const response = (error as { response?: { data?: unknown } }).response
if (typeof response?.data !== 'object' || response.data === null) return false
return (response.data as Record<string, unknown>).outcome === 'credential_changed'
}
async function handleConsumeCodexResetCredit(key: EndpointAPIKey) { async function handleConsumeCodexResetCredit(key: EndpointAPIKey) {
if (!canConsumeCodexResetCredit(key)) return if (!canConsumeCodexResetCredit(key)) return
const credentialGeneration = getCodexCredentialGeneration(key)
if (credentialGeneration === undefined) return
const pendingIdempotencyKey = getPendingCodexResetCreditIdempotencyKey(key)
const earliest = getVisibleCodexResetCreditItems(key)[0] const earliest = getVisibleCodexResetCreditItems(key)[0]
const detailMessage = earliest const detailMessage = earliest
? `\n当前最早过期项:${earliest.displayKey}${formatCodexResetCreditExpiresAt(earliest.expiresAt)} 过期。` ? `\n当前最早过期项:${earliest.displayKey}${formatCodexResetCreditExpiresAt(earliest.expiresAt)} 过期。`
: '' : ''
const confirmed = await confirm({ const confirmed = await confirm({
title: legacyT('确认使用 Codex 重置机会'), title: legacyT('确认使用 Codex 重置机会'),
message: `${legacyT('将消耗 1 次 Codex 重置机会。操作完成后会重新刷新账号配额状态。')}${detailMessage}`, message: pendingIdempotencyKey
? legacyT('将继续确认上次尚未完成的 Codex 重置请求。')
: `${legacyT('将消耗 1 次 Codex 重置机会。操作完成后会重新刷新账号配额状态。')}${detailMessage}`,
confirmText: legacyT('确认重置'), confirmText: legacyT('确认重置'),
cancelText: legacyT('取消'), cancelText: legacyT('取消'),
variant: 'warning', variant: 'warning',
@@ -1759,10 +1785,15 @@ async function handleConsumeCodexResetCredit(key: EndpointAPIKey) {
consumingCodexResetCreditKeyId.value = key.id consumingCodexResetCreditKeyId.value = key.id
try { try {
const idempotencyKey = createCodexResetCreditIdempotencyKey() const idempotencyKey = pendingIdempotencyKey
|| readPendingCodexResetCreditIdempotencyKey(key.id, credentialGeneration)
|| createCodexResetCreditIdempotencyKey()
rememberPendingCodexResetCreditIdempotencyKey(key.id, idempotencyKey, credentialGeneration)
const result = await consumeCodexResetCredit(key.id, { const result = await consumeCodexResetCredit(key.id, {
idempotency_key: idempotencyKey, idempotency_key: idempotencyKey,
expected_credential_generation: credentialGeneration,
}) })
clearPendingCodexResetCreditIdempotencyKeyForOutcome(key.id, result.outcome)
applyQuotaResults([{ applyQuotaResults([{
key_id: result.key_id, key_id: result.key_id,
status: result.refresh_status === 'success' ? 'success' : result.status, status: result.refresh_status === 'success' ? 'success' : result.status,
@@ -1781,6 +1812,16 @@ async function handleConsumeCodexResetCredit(key: EndpointAPIKey) {
} }
emit('refresh') emit('refresh')
} catch (err: unknown) { } catch (err: unknown) {
const activeIdempotencyKey = codexResetCreditActiveIdempotencyKeyFromError(err)
if (codexResetCreditCredentialChangedFromError(err)) {
clearPendingCodexResetCreditIdempotencyKey(key.id)
} else if (activeIdempotencyKey) {
rememberPendingCodexResetCreditIdempotencyKey(
key.id,
activeIdempotencyKey,
credentialGeneration,
)
}
showError(localizedApiError(err, 'Codex 重置机会使用失败'), legacyT('错误')) showError(localizedApiError(err, 'Codex 重置机会使用失败'), legacyT('错误'))
await Promise.all([loadProvider(), loadEndpoints()]) await Promise.all([loadProvider(), loadEndpoints()])
emit('refresh') emit('refresh')
@@ -1909,6 +1950,9 @@ function getCodexQuotaDisplayFromMetadata(metadata: CodexUpstreamMetadata | null
if (!metadata) return null if (!metadata) return null
const display: CodexUpstreamMetadata = {} const display: CodexUpstreamMetadata = {}
if (metadata.credential_generation?.trim()) {
display.credential_generation = metadata.credential_generation.trim()
}
if (metadata.plan_type) display.plan_type = metadata.plan_type if (metadata.plan_type) display.plan_type = metadata.plan_type
const numberFields: (keyof CodexUpstreamMetadata)[] = [ const numberFields: (keyof CodexUpstreamMetadata)[] = [
@@ -1938,6 +1982,9 @@ function getCodexQuotaDisplayFromMetadata(metadata: CodexUpstreamMetadata | null
numberFields.forEach(field => copyCodexNumberField(display, metadata, field)) numberFields.forEach(field => copyCodexNumberField(display, metadata, field))
if (metadata.has_credits !== undefined) display.has_credits = metadata.has_credits if (metadata.has_credits !== undefined) display.has_credits = metadata.has_credits
if (metadata.reset_credits) display.reset_credits = metadata.reset_credits if (metadata.reset_credits) display.reset_credits = metadata.reset_credits
if (metadata.account_quota_reset_reservation) {
display.account_quota_reset_reservation = metadata.account_quota_reset_reservation
}
return Object.keys(display).length > 0 ? display : null return Object.keys(display).length > 0 ? display : null
} }
@@ -2025,6 +2072,7 @@ function hasCodexQuotaDisplayData(key: EndpointAPIKey): boolean {
|| codex.spark_primary_used_percent !== undefined || codex.spark_primary_used_percent !== undefined
|| codex.spark_secondary_used_percent !== undefined || codex.spark_secondary_used_percent !== undefined
|| codexDisplayHasResetCredits(codex) || codexDisplayHasResetCredits(codex)
|| hasPendingCodexResetCredit(key)
) )
} }
@@ -2052,10 +2100,29 @@ function getVisibleCodexResetCreditItems(key: EndpointAPIKey) {
return getVisibleCodexResetCreditItemsFromSnapshot(getCodexResetCreditsDisplay(key)) return getVisibleCodexResetCreditItemsFromSnapshot(getCodexResetCreditsDisplay(key))
} }
function getPendingCodexResetCreditIdempotencyKey(key: EndpointAPIKey): string | null {
const serverReservation = getCodexResetCreditReservationIdempotencyKey(getCodexQuotaDisplay(key))
if (serverReservation) return serverReservation
const credentialGeneration = getCodexCredentialGeneration(key)
return credentialGeneration === undefined
? null
: readPendingCodexResetCreditIdempotencyKey(key.id, credentialGeneration)
}
function getCodexCredentialGeneration(key: EndpointAPIKey): string | null | undefined {
const codex = key.upstream_metadata?.codex
if (!codex || typeof codex !== 'object') return undefined
return codex.credential_generation?.trim() || null
}
function hasPendingCodexResetCredit(key: EndpointAPIKey): boolean {
return getPendingCodexResetCreditIdempotencyKey(key) !== null
}
function canConsumeCodexResetCredit(key: EndpointAPIKey): boolean { function canConsumeCodexResetCredit(key: EndpointAPIKey): boolean {
return provider.value?.provider_type === 'codex' return provider.value?.provider_type === 'codex'
&& getCodexResetCreditAvailableCount(key) !== null && getCodexCredentialGeneration(key) !== undefined
&& (getCodexResetCreditAvailableCount(key) ?? 0) > 0 && (hasPendingCodexResetCredit(key) || (getCodexResetCreditAvailableCount(key) ?? 0) > 0)
&& !consumingCodexResetCreditKeyId.value && !consumingCodexResetCreditKeyId.value
} }
@@ -1,12 +1,18 @@
import { describe, expect, it } from 'vitest' import { describe, expect, it } from 'vitest'
import { import {
clearPendingCodexResetCreditIdempotencyKey,
clearPendingCodexResetCreditIdempotencyKeyForOutcome,
createCodexResetCreditIdempotencyKey, createCodexResetCreditIdempotencyKey,
formatCodexResetCreditCount, formatCodexResetCreditCount,
formatCodexResetCreditExpiresAt, formatCodexResetCreditExpiresAt,
getCodexResetCreditAvailableCount, getCodexResetCreditAvailableCount,
getCodexResetCreditReservationIdempotencyKey,
getVisibleCodexResetCreditItems, getVisibleCodexResetCreditItems,
isCodexResetCreditTerminalOutcome,
mergeCodexQuotaDisplays, mergeCodexQuotaDisplays,
readPendingCodexResetCreditIdempotencyKey,
rememberPendingCodexResetCreditIdempotencyKey,
} from '@/features/providers/components/codex-reset-credit-display' } from '@/features/providers/components/codex-reset-credit-display'
import type { QuotaResetCreditsSnapshot } from '@/api/endpoints/types' import type { QuotaResetCreditsSnapshot } from '@/api/endpoints/types'
@@ -59,6 +65,19 @@ describe('codex reset credit display helpers', () => {
expect(getVisibleCodexResetCreditItems(snapshot, 1_700_000_000)).toEqual([]) expect(getVisibleCodexResetCreditItems(snapshot, 1_700_000_000)).toEqual([])
}) })
it('recovers the active idempotency key from a persisted server reservation', () => {
expect(getCodexResetCreditReservationIdempotencyKey({
reset_credits: { available_count: 0 },
account_quota_reset_reservation: {
idempotency_key: ' server-active-attempt ',
generation: 3,
},
})).toBe('server-active-attempt')
expect(getCodexResetCreditReservationIdempotencyKey({
account_quota_reset_reservation: { idempotency_key: ' ' },
})).toBeNull()
})
it('sorts available detail items by remaining time and labels visible items with short ordinal keys', () => { it('sorts available detail items by remaining time and labels visible items with short ordinal keys', () => {
const snapshot: QuotaResetCreditsSnapshot = { const snapshot: QuotaResetCreditsSnapshot = {
available_count: 7, available_count: 7,
@@ -139,4 +158,164 @@ describe('codex reset credit display helpers', () => {
getRandomValues: array => array, getRandomValues: array => array,
})).toBe('existing-random-uuid') })).toBe('existing-random-uuid')
}) })
it('keeps an unresolved reset idempotency key until a terminal response clears it', () => {
const values = new Map<string, string>()
const storage = {
getItem: (key: string) => values.get(key) ?? null,
setItem: (key: string, value: string) => values.set(key, value),
removeItem: (key: string) => values.delete(key),
}
rememberPendingCodexResetCreditIdempotencyKey(
'key-1',
'reset-attempt-1',
'credential-v1',
storage,
)
expect(readPendingCodexResetCreditIdempotencyKey('key-1', 'credential-v1', storage))
.toBe('reset-attempt-1')
expect(readPendingCodexResetCreditIdempotencyKey('key-2', 'credential-v1', storage)).toBeNull()
clearPendingCodexResetCreditIdempotencyKey('key-1', storage)
expect(readPendingCodexResetCreditIdempotencyKey('key-1', 'credential-v1', storage)).toBeNull()
})
it('clears pending idempotency keys for compatibility history replays', () => {
const values = new Map<string, string>()
const storage = {
getItem: (key: string) => values.get(key) ?? null,
setItem: (key: string, value: string) => values.set(key, value),
removeItem: (key: string) => values.delete(key),
}
rememberPendingCodexResetCreditIdempotencyKey(
'key-1',
'reset-attempt-1',
'credential-v1',
storage,
)
expect(isCodexResetCreditTerminalOutcome('historical_replay')).toBe(true)
expect(clearPendingCodexResetCreditIdempotencyKeyForOutcome(
'key-1',
'historical_replay',
storage,
)).toBe(true)
expect(readPendingCodexResetCreditIdempotencyKey('key-1', 'credential-v1', storage)).toBeNull()
rememberPendingCodexResetCreditIdempotencyKey(
'key-1',
'reset-attempt-2',
'credential-v1',
storage,
)
expect(clearPendingCodexResetCreditIdempotencyKeyForOutcome('key-1', 'unknown', storage))
.toBe(false)
expect(readPendingCodexResetCreditIdempotencyKey('key-1', 'credential-v1', storage))
.toBe('reset-attempt-2')
expect(isCodexResetCreditTerminalOutcome('unknown')).toBe(false)
expect(isCodexResetCreditTerminalOutcome('error')).toBe(false)
})
it('keeps the active idempotency key in memory when session storage is unavailable', () => {
const unavailableStorage = {
getItem: () => { throw new Error('storage unavailable') },
setItem: () => { throw new Error('storage unavailable') },
removeItem: () => { throw new Error('storage unavailable') },
}
rememberPendingCodexResetCreditIdempotencyKey(
'key-without-storage',
'reset-attempt-original',
'credential-v1',
unavailableStorage,
)
expect(readPendingCodexResetCreditIdempotencyKey(
'key-without-storage',
'credential-v1',
unavailableStorage,
))
.toBe('reset-attempt-original')
rememberPendingCodexResetCreditIdempotencyKey(
'key-without-storage',
'reset-attempt-from-conflict',
'credential-v1',
unavailableStorage,
)
expect(readPendingCodexResetCreditIdempotencyKey(
'key-without-storage',
'credential-v1',
unavailableStorage,
))
.toBe('reset-attempt-from-conflict')
clearPendingCodexResetCreditIdempotencyKey('key-without-storage', unavailableStorage)
expect(readPendingCodexResetCreditIdempotencyKey(
'key-without-storage',
'credential-v1',
unavailableStorage,
))
.toBeNull()
})
it('does not resurrect a cleared key when session storage removal fails', () => {
const values = new Map<string, string>()
const partiallyUnavailableStorage = {
getItem: (key: string) => values.get(key) ?? null,
setItem: (key: string, value: string) => values.set(key, value),
removeItem: () => { throw new Error('storage removal unavailable') },
}
rememberPendingCodexResetCreditIdempotencyKey(
'key-with-stale-storage',
'terminal-attempt',
'credential-v1',
partiallyUnavailableStorage,
)
clearPendingCodexResetCreditIdempotencyKey(
'key-with-stale-storage',
partiallyUnavailableStorage,
)
expect(readPendingCodexResetCreditIdempotencyKey(
'key-with-stale-storage',
'credential-v1',
partiallyUnavailableStorage,
)).toBeNull()
})
it('drops a pending attempt when the credential generation changes', () => {
const values = new Map<string, string>()
const storage = {
getItem: (key: string) => values.get(key) ?? null,
setItem: (key: string, value: string) => values.set(key, value),
removeItem: (key: string) => values.delete(key),
}
rememberPendingCodexResetCreditIdempotencyKey(
'key-rotated',
'attempt-for-account-a',
'credential-a',
storage,
)
expect(readPendingCodexResetCreditIdempotencyKey('key-rotated', 'credential-b', storage))
.toBeNull()
expect([...values.values()]).toEqual([])
})
it('never replays a legacy v1 pending value without a credential generation', () => {
const values = new Map<string, string>([
['aether:codex-reset-credit-pending:v1:key-legacy', 'legacy-attempt'],
])
const storage = {
getItem: (key: string) => values.get(key) ?? null,
setItem: (key: string, value: string) => values.set(key, value),
removeItem: (key: string) => values.delete(key),
}
expect(readPendingCodexResetCreditIdempotencyKey('key-legacy', 'credential-current', storage))
.toBeNull()
expect([...values.values()]).toEqual([])
})
}) })
@@ -66,6 +66,84 @@ interface CodexResetCreditCrypto {
getRandomValues: (array: Uint8Array) => Uint8Array getRandomValues: (array: Uint8Array) => Uint8Array
} }
interface CodexResetCreditPendingStorage {
getItem(key: string): string | null
setItem(key: string, value: string): void
removeItem(key: string): void
}
const CODEX_RESET_CREDIT_PENDING_STORAGE_PREFIX = 'aether:codex-reset-credit-pending:v2:'
const CODEX_RESET_CREDIT_LEGACY_PENDING_STORAGE_PREFIX = 'aether:codex-reset-credit-pending:v1:'
interface CodexResetCreditPendingAttempt {
idempotencyKey: string
credentialGeneration: string | null
}
const pendingCodexResetCreditAttempts = new Map<string, CodexResetCreditPendingAttempt | null>()
const CODEX_RESET_CREDIT_TERMINAL_OUTCOMES = new Set([
'reset',
'already_redeemed',
'nothing_to_reset',
'no_credit',
'historical_replay',
])
function codexResetCreditPendingStorageKey(keyId: string): string {
return `${CODEX_RESET_CREDIT_PENDING_STORAGE_PREFIX}${keyId.trim()}`
}
function codexResetCreditLegacyPendingStorageKey(keyId: string): string {
return `${CODEX_RESET_CREDIT_LEGACY_PENDING_STORAGE_PREFIX}${keyId.trim()}`
}
function normalizeCodexCredentialGeneration(value: string | null | undefined): string | null {
const normalized = value?.trim()
return normalized ? normalized : null
}
function parseCodexResetCreditPendingAttempt(value: string): CodexResetCreditPendingAttempt | null {
try {
const parsed = JSON.parse(value) as unknown
if (typeof parsed !== 'object' || parsed === null) return null
const record = parsed as Record<string, unknown>
const idempotencyKey = typeof record.idempotencyKey === 'string'
? record.idempotencyKey.trim()
: ''
const generation = record.credentialGeneration
if (!idempotencyKey || idempotencyKey.length > 256) return null
if (generation !== null && typeof generation !== 'string') return null
const credentialGeneration = normalizeCodexCredentialGeneration(generation)
if (typeof generation === 'string' && credentialGeneration === null) return null
return { idempotencyKey, credentialGeneration }
} catch {
return null
}
}
function resolveCodexResetCreditPendingStorage(
storage: CodexResetCreditPendingStorage | undefined,
): CodexResetCreditPendingStorage | undefined {
if (storage) return storage
try {
return globalThis.sessionStorage
} catch {
return undefined
}
}
export function isCodexResetCreditTerminalOutcome(outcome: string): boolean {
return CODEX_RESET_CREDIT_TERMINAL_OUTCOMES.has(outcome)
}
export function getCodexResetCreditReservationIdempotencyKey(
display: CodexUpstreamMetadata | null | undefined,
): string | null {
const value = display?.account_quota_reset_reservation?.idempotency_key?.trim()
return value && value.length <= 256 ? value : null
}
export function createCodexResetCreditIdempotencyKey( export function createCodexResetCreditIdempotencyKey(
cryptoSource: CodexResetCreditCrypto | undefined = globalThis.crypto, cryptoSource: CodexResetCreditCrypto | undefined = globalThis.crypto,
): string { ): string {
@@ -82,6 +160,91 @@ export function createCodexResetCreditIdempotencyKey(
return `${hex.slice(0, 8)}-${hex.slice(8, 12)}-${hex.slice(12, 16)}-${hex.slice(16, 20)}-${hex.slice(20)}` return `${hex.slice(0, 8)}-${hex.slice(8, 12)}-${hex.slice(12, 16)}-${hex.slice(16, 20)}-${hex.slice(20)}`
} }
export function readPendingCodexResetCreditIdempotencyKey(
keyId: string,
expectedCredentialGeneration: string | null,
storage?: CodexResetCreditPendingStorage,
): string | null {
const storageKey = codexResetCreditPendingStorageKey(keyId)
const expectedGeneration = normalizeCodexCredentialGeneration(expectedCredentialGeneration)
const resolvedStorage = resolveCodexResetCreditPendingStorage(storage)
try {
resolvedStorage?.removeItem(codexResetCreditLegacyPendingStorageKey(keyId))
} catch {
// Legacy values are never read, so failed cleanup cannot replay them.
}
if (pendingCodexResetCreditAttempts.has(storageKey)) {
const attempt = pendingCodexResetCreditAttempts.get(storageKey)
if (attempt?.credentialGeneration === expectedGeneration) return attempt.idempotencyKey
pendingCodexResetCreditAttempts.set(storageKey, null)
try {
resolvedStorage?.removeItem(storageKey)
} catch {
// The in-memory tombstone still prevents a stale generation from replaying.
}
return null
}
try {
const value = resolvedStorage?.getItem(storageKey)?.trim()
const attempt = value ? parseCodexResetCreditPendingAttempt(value) : null
if (attempt?.credentialGeneration === expectedGeneration) {
pendingCodexResetCreditAttempts.set(storageKey, attempt)
return attempt.idempotencyKey
}
pendingCodexResetCreditAttempts.set(storageKey, null)
if (value) resolvedStorage?.removeItem(storageKey)
} catch {
// Fall through to the in-memory copy when browser storage is unavailable.
}
return null
}
export function rememberPendingCodexResetCreditIdempotencyKey(
keyId: string,
idempotencyKey: string,
credentialGeneration: string | null,
storage?: CodexResetCreditPendingStorage,
): void {
const normalized = idempotencyKey.trim()
if (!normalized || normalized.length > 256) return
const storageKey = codexResetCreditPendingStorageKey(keyId)
const attempt: CodexResetCreditPendingAttempt = {
idempotencyKey: normalized,
credentialGeneration: normalizeCodexCredentialGeneration(credentialGeneration),
}
pendingCodexResetCreditAttempts.set(storageKey, attempt)
try {
resolveCodexResetCreditPendingStorage(storage)?.setItem(storageKey, JSON.stringify(attempt))
} catch {
// The in-memory copy still permits retrying the same request in this page session.
}
}
export function clearPendingCodexResetCreditIdempotencyKey(
keyId: string,
storage?: CodexResetCreditPendingStorage,
): void {
const storageKey = codexResetCreditPendingStorageKey(keyId)
pendingCodexResetCreditAttempts.set(storageKey, null)
try {
const resolvedStorage = resolveCodexResetCreditPendingStorage(storage)
resolvedStorage?.removeItem(storageKey)
resolvedStorage?.removeItem(codexResetCreditLegacyPendingStorageKey(keyId))
} catch {
// Ignore unavailable browser storage after a terminal server response.
}
}
export function clearPendingCodexResetCreditIdempotencyKeyForOutcome(
keyId: string,
outcome: string,
storage?: CodexResetCreditPendingStorage,
): boolean {
if (!isCodexResetCreditTerminalOutcome(outcome)) return false
clearPendingCodexResetCreditIdempotencyKey(keyId, storage)
return true
}
function codexResetCreditRemainingSeconds( function codexResetCreditRemainingSeconds(
item: QuotaResetCreditSnapshot, item: QuotaResetCreditSnapshot,
snapshot: QuotaResetCreditsSnapshot, snapshot: QuotaResetCreditsSnapshot,