fix: tighten OAuth auto cleanup signals

This commit is contained in:
fawney19
2026-05-20 10:34:22 +08:00
parent fbda210b84
commit 61bcbe826a
13 changed files with 922 additions and 81 deletions

View File

@@ -1,7 +1,7 @@
use aether_contracts::ExecutionPlan;
use tracing::warn;
use crate::AppState;
use crate::{provider_transport::LocalOAuthRefreshError, AppState};
pub(crate) async fn refresh_oauth_plan_auth_for_retry(
state: &AppState,
@@ -13,6 +13,8 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
if !status_may_be_oauth_invalid(status_code, response_text) {
return false;
}
let access_token_invalid_proven =
status_proves_access_token_invalid(status_code, response_text);
let transport = match state
.read_provider_transport_snapshot(&plan.provider_id, &plan.endpoint_id, &plan.key_id)
@@ -52,6 +54,46 @@ pub(crate) async fn refresh_oauth_plan_auth_for_retry(
true
}
Ok(None) => false,
Err(LocalOAuthRefreshError::HttpStatus {
status_code: refresh_status_code,
body_excerpt,
..
}) if matches!(refresh_status_code, 400 | 401 | 403) => {
if let Err(err) = state
.persist_local_oauth_refresh_failure_state(
&transport,
refresh_status_code,
body_excerpt.as_str(),
access_token_invalid_proven,
)
.await
{
warn!(
event_name = "local_oauth_retry_refresh_failure_persist_failed",
log_type = "ops",
trace_id = %trace_id,
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
status_code,
refresh_status_code,
error = ?err,
"gateway failed to persist oauth retry refresh failure"
);
}
warn!(
event_name = "local_oauth_retry_refresh_failed",
log_type = "ops",
trace_id = %trace_id,
provider_id = %plan.provider_id,
endpoint_id = %plan.endpoint_id,
key_id = %plan.key_id,
status_code,
refresh_status_code,
"gateway oauth retry refresh failed"
);
false
}
Err(err) => {
warn!(
event_name = "local_oauth_retry_refresh_failed",
@@ -86,9 +128,54 @@ fn status_may_be_oauth_invalid(status_code: u16, response_text: Option<&str>) ->
.any(|needle| response_text.contains(needle))
}
fn status_proves_access_token_invalid(status_code: u16, response_text: Option<&str>) -> bool {
if status_code == 401 {
return true;
}
if status_code != 403 {
return false;
}
let Some(response_text) = response_text else {
return false;
};
let response_text = response_text.to_ascii_lowercase();
[
"oauth_token_invalid",
"invalid_token",
"invalid access token",
"access token invalid",
"access token expired",
"expired access token",
"authentication token has been invalidated",
"token has been invalidated",
"security token included in the request is expired",
]
.iter()
.any(|needle| response_text.contains(needle))
}
#[cfg(test)]
mod tests {
use super::status_may_be_oauth_invalid;
use super::{
refresh_oauth_plan_auth_for_retry, status_may_be_oauth_invalid,
status_proves_access_token_invalid,
};
use std::collections::BTreeMap;
use std::sync::{Arc, Mutex};
use aether_contracts::{ExecutionPlan, RequestBody};
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogReadRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogProvider,
};
use axum::routing::post;
use axum::{extract::Request, Json, Router};
use http::StatusCode;
use serde_json::json;
use tokio::task::JoinHandle;
#[test]
fn recognizes_oauth_invalid_statuses() {
@@ -101,4 +188,198 @@ mod tests {
assert!(!status_may_be_oauth_invalid(403, Some("quota exceeded")));
assert!(!status_may_be_oauth_invalid(429, Some("token bucket")));
}
#[test]
fn separates_retry_candidate_from_access_token_invalid_proof() {
assert!(status_proves_access_token_invalid(401, None));
assert!(status_proves_access_token_invalid(
403,
Some("The security token included in the request is expired")
));
assert!(!status_proves_access_token_invalid(403, None));
assert!(!status_proves_access_token_invalid(
403,
Some("quota exceeded")
));
assert!(!status_proves_access_token_invalid(
429,
Some("token bucket")
));
}
#[tokio::test]
async fn auto_removes_request_proven_oauth_failure_after_terminal_refresh_failure() {
let token_hits = Arc::new(Mutex::new(0usize));
let token_hits_clone = Arc::clone(&token_hits);
let token_server = Router::new().route(
"/oauth/token",
post(move |_request: Request| {
let token_hits_inner = Arc::clone(&token_hits_clone);
async move {
*token_hits_inner.lock().expect("mutex should lock") += 1;
(
StatusCode::UNAUTHORIZED,
Json(json!({
"error": {
"message": "Your refresh token has already been used to generate a new access token. Please try signing in again.",
"type": "invalid_request_error",
"code": "refresh_token_reused"
}
})),
)
}
}),
);
let mut provider = StoredProviderCatalogProvider::new(
"provider-codex".to_string(),
"codex".to_string(),
Some("https://example.com".to_string()),
"codex".to_string(),
)
.expect("provider should build")
.with_routing_fields(10);
provider.config = Some(json!({
"pool_advanced": {
"auto_remove_banned_keys": true
}
}));
let endpoint = StoredProviderCatalogEndpoint::new(
"endpoint-codex-cli".to_string(),
"provider-codex".to_string(),
"openai:responses".to_string(),
None,
None,
true,
)
.expect("endpoint should build")
.with_transport_fields(
"https://chatgpt.com/backend-api/codex".to_string(),
None,
None,
None,
None,
None,
None,
None,
)
.expect("endpoint transport should build");
let encrypted_api_key =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "stale-codex-token")
.expect("api key ciphertext should build");
let mut key = StoredProviderCatalogKey::new(
"key-codex-oauth-retry".to_string(),
"provider-codex".to_string(),
"default".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build")
.with_transport_fields(
Some(json!(["openai:responses"])),
encrypted_api_key,
None,
None,
None,
None,
None,
None,
None,
)
.expect("key transport should build");
key.expires_at_unix_secs = Some(4_102_444_800);
key.encrypted_auth_config = Some(
encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
r#"{"provider_type":"codex","refresh_token":"used-refresh-token","email":"alice@example.com","account_id":"acct-codex-123","plan_type":"plus","expires_at":4102444800}"#,
)
.expect("auth config ciphertext should build"),
);
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![key],
));
let (token_url, token_handle) = start_test_server(token_server).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 state = crate::AppState::new()
.expect("state should build")
.with_data_state_for_tests(
crate::data::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 mut plan = ExecutionPlan {
request_id: "req-oauth-retry".to_string(),
candidate_id: None,
provider_name: Some("codex".to_string()),
provider_id: "provider-codex".to_string(),
endpoint_id: "endpoint-codex-cli".to_string(),
key_id: "key-codex-oauth-retry".to_string(),
method: "POST".to_string(),
url: "https://chatgpt.com/backend-api/codex/responses".to_string(),
headers: BTreeMap::from([(
"authorization".to_string(),
"Bearer stale-codex-token".to_string(),
)]),
content_type: Some("application/json".to_string()),
content_encoding: None,
body: RequestBody::from_json(json!({"model": "gpt-5"})),
stream: false,
client_api_format: "openai:responses".to_string(),
provider_api_format: "openai:responses".to_string(),
model_name: Some("gpt-5".to_string()),
proxy: None,
transport_profile: None,
timeouts: None,
};
let retried = refresh_oauth_plan_auth_for_retry(
&state,
&mut plan,
401,
Some(r#"{"error":"oauth_token_invalid"}"#),
"trace-oauth-retry",
)
.await;
assert!(!retried);
assert_eq!(*token_hits.lock().expect("mutex should lock"), 1);
let keys = provider_catalog_repository
.list_keys_by_ids(&["key-codex-oauth-retry".to_string()])
.await
.expect("keys should read");
assert!(keys.is_empty());
token_handle.abort();
}
async fn start_test_server(router: Router) -> (String, JoinHandle<()>) {
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("test server should bind");
let addr = listener
.local_addr()
.expect("test server address should resolve");
let handle = tokio::spawn(async move {
axum::serve(listener, router)
.await
.expect("test server should serve");
});
(format!("http://{addr}"), handle)
}
}

View File

@@ -1,7 +1,10 @@
use super::super::super::errors::{
merge_provider_oauth_refresh_failure_reason, normalize_provider_oauth_refresh_error_message,
};
use super::super::super::quota::shared::persist_provider_quota_refresh_state;
use super::super::super::quota::shared::{
persist_provider_quota_refresh_state, provider_auto_remove_banned_keys,
should_auto_remove_oauth_invalid_key,
};
use super::super::super::runtime::refresh_provider_oauth_account_state_after_update;
use super::helpers::{self, RefreshDispatch, RefreshRequestContext, RefreshSuccessContext};
use super::response;
@@ -76,6 +79,50 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
None,
)
.await?;
if provider_auto_remove_banned_keys(provider.config.as_ref()) {
let now_unix_secs = helpers::unix_now_secs();
let latest_key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next();
if latest_key.as_ref().is_some_and(|latest_key| {
should_auto_remove_oauth_invalid_key(
latest_key,
None,
false,
now_unix_secs,
)
}) {
state
.clear_admin_provider_pool_cooldown(&provider.id, &key_id)
.await;
state
.reset_admin_provider_pool_cost(&provider.id, &key_id)
.await;
if state.delete_provider_catalog_key(&key_id).await? {
let deleted_key_ids = [key_id.clone()];
state
.cleanup_deleted_provider_catalog_refs(
&provider.id,
&[],
&deleted_key_ids,
)
.await?;
tracing::info!(
trace_id = %trace_id,
key_id = %key_id,
provider_id = %provider.id,
provider_type = %provider_type,
event_name = "auto_removed_oauth_refresh_failed",
"gateway manual provider oauth refresh auto-removed unusable key"
);
return Ok(RefreshDispatch::Respond(
response::oauth_refresh_auto_removed_response(&error_reason),
));
}
}
}
}
}
return Ok(RefreshDispatch::Respond(

View File

@@ -24,6 +24,15 @@ pub(super) fn oauth_refresh_failed_bad_request_response(
)
}
pub(super) fn oauth_refresh_auto_removed_response(error_reason: impl AsRef<str>) -> Response<Body> {
Json(json!({
"status": "auto_removed",
"message": "已自动删除",
"detail": format!("Token 刷新失败且 OAuth 凭证已不可用,已自动删除:{}", error_reason.as_ref()),
}))
.into_response()
}
pub(super) fn oauth_refresh_failed_service_unavailable_response(
error_reason: impl Into<String>,
) -> Response<Body> {

View File

@@ -1,7 +1,8 @@
use super::shared::{
build_provider_quota_execution_plan, build_quota_snapshot_payload, coerce_json_f64,
coerce_json_string, default_provider_quota_execution_timeouts, execute_provider_quota_plan,
extract_execution_error_message, persist_provider_quota_refresh_state,
extract_execution_error_message, oauth_refresh_auto_removed_result,
persist_provider_quota_refresh_state, quota_key_auto_removed,
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
@@ -65,6 +66,7 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
let mut results = Vec::new();
let mut success_count = 0usize;
let mut failed_count = 0usize;
let mut auto_removed_count = 0usize;
for key in keys {
let transport = match state
@@ -87,6 +89,11 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
let authorization = match state.resolve_local_oauth_header_auth(&transport).await? {
Some(auth) => auth,
_ => {
if quota_key_auto_removed(state, &key.id).await? {
auto_removed_count += 1;
results.push(oauth_refresh_auto_removed_result(&key));
continue;
}
failed_count += 1;
results.push(json!({
"key_id": key.id,
@@ -247,9 +254,9 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
Ok(Some(json!({
"success": success_count,
"failed": failed_count,
"total": success_count + failed_count,
"total": results.len(),
"results": results,
"message": format!("已处理 {} 个 Key", success_count + failed_count),
"auto_removed": 0,
"message": format!("已处理 {} 个 Key", results.len()),
"auto_removed": auto_removed_count,
})))
}

View File

@@ -1,8 +1,8 @@
use super::shared::{
build_quota_snapshot_payload, default_provider_quota_execution_timeouts,
execute_provider_quota_plan, extract_execution_error_message,
persist_provider_quota_refresh_state, quota_refresh_success_invalid_state,
ProviderQuotaExecutionOutcome,
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
quota_key_auto_removed, quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::provider::shared::payloads::{
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX,
@@ -113,6 +113,7 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
let mut results = Vec::new();
let mut success_count = 0usize;
let mut failed_count = 0usize;
let mut auto_removed_count = 0usize;
for key in keys {
let transport = match state
@@ -135,6 +136,11 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
let authorization = match resolve_chatgpt_web_quota_auth(state, &transport).await? {
Some(auth) => auth,
None => {
if quota_key_auto_removed(state, &key.id).await? {
auto_removed_count += 1;
results.push(oauth_refresh_auto_removed_result(&key));
continue;
}
failed_count += 1;
results.push(json!({
"key_id": key.id,
@@ -291,9 +297,9 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
Ok(Some(json!({
"success": success_count,
"failed": failed_count,
"total": success_count + failed_count,
"total": results.len(),
"results": results,
"message": format!("已处理 {} 个 Key", success_count + failed_count),
"auto_removed": 0,
"message": format!("已处理 {} 个 Key", results.len()),
"auto_removed": auto_removed_count,
})))
}

View File

@@ -14,9 +14,9 @@ use self::parse::{
use self::plan::{build_codex_quota_request_spec, execute_codex_quota_plan};
use super::shared::{
build_quota_snapshot_payload, extract_execution_error_message,
persist_provider_quota_refresh_state, provider_auto_remove_banned_keys,
quota_refresh_success_invalid_state, should_auto_remove_oauth_invalid_key,
ProviderQuotaExecutionOutcome,
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
provider_auto_remove_banned_keys, quota_key_auto_removed, quota_refresh_success_invalid_state,
should_auto_remove_structured_reason, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::request::AdminAppState;
use crate::provider_key_auth::provider_key_is_oauth_managed;
@@ -75,12 +75,27 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
}
};
let resolved_oauth_auth =
if provider_key_is_oauth_managed(&key, provider.provider_type.as_str()) {
state.resolve_local_oauth_header_auth(&transport).await?
} else {
None
};
let is_oauth_managed = provider_key_is_oauth_managed(&key, provider.provider_type.as_str());
let resolved_oauth_auth = if is_oauth_managed {
state.resolve_local_oauth_header_auth(&transport).await?
} else {
None
};
if is_oauth_managed && quota_key_auto_removed(state, &key.id).await? {
auto_removed_count += 1;
results.push(oauth_refresh_auto_removed_result(&key));
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 request_spec = match build_codex_quota_request_spec(&transport, resolved_oauth_auth) {
Ok(request_spec) => request_spec,
@@ -261,22 +276,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
}
}
let auto_remove_key = if auto_remove_abnormal_keys {
state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key.id))
.await?
.into_iter()
.next()
.unwrap_or_else(|| key.clone())
} else {
key.clone()
};
let auto_removed = auto_remove_abnormal_keys
&& should_auto_remove_oauth_invalid_key(
&auto_remove_key,
oauth_invalid_reason.as_deref(),
now_unix_secs,
);
&& should_auto_remove_structured_reason(oauth_invalid_reason.as_deref());
if auto_removed {
if state.delete_provider_catalog_key(&key.id).await? {
auto_removed_count += 1;

View File

@@ -5,7 +5,8 @@ use self::parse::parse_kiro_usage_response;
use self::plan::execute_kiro_quota_plan;
use super::shared::{
build_quota_snapshot_payload, extract_execution_error_message,
persist_provider_quota_refresh_state, quota_refresh_success_invalid_state,
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
persist_quota_oauth_refresh_failure_state, quota_refresh_success_invalid_state,
ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::request::{AdminAppState, AdminLocalOAuthRefreshError};
@@ -95,6 +96,7 @@ pub(crate) async fn refresh_kiro_provider_quota_locally(
let mut results = Vec::new();
let mut success_count = 0usize;
let mut failed_count = 0usize;
let mut auto_removed_count = 0usize;
for key in keys {
let transport = match state
@@ -145,6 +147,13 @@ pub(crate) async fn refresh_kiro_provider_quota_locally(
}
},
Err(err) => {
if persist_quota_oauth_refresh_failure_state(state, &transport, &err).await?
|| super::shared::quota_key_auto_removed(state, &key.id).await?
{
auto_removed_count += 1;
results.push(oauth_refresh_auto_removed_result(&key));
continue;
}
failed_count += 1;
let mut payload = serde_json::Map::new();
payload.insert("key_id".to_string(), json!(key.id));
@@ -328,10 +337,10 @@ pub(crate) async fn refresh_kiro_provider_quota_locally(
Ok(Some(json!({
"success": success_count,
"failed": failed_count,
"total": success_count + failed_count,
"total": results.len(),
"results": results,
"message": format!("已处理 {} 个 Key", success_count + failed_count),
"auto_removed": 0,
"message": format!("已处理 {} 个 Key", results.len()),
"auto_removed": auto_removed_count,
})))
}

View File

@@ -1,7 +1,9 @@
use crate::handlers::admin::provider::shared::payloads::{
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
};
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::handlers::admin::request::{
AdminAppState, AdminGatewayProviderTransportSnapshot, AdminLocalOAuthRefreshError,
};
use crate::handlers::shared::{
sync_provider_key_oauth_status_snapshot, sync_provider_key_quota_status_snapshot,
};
@@ -44,22 +46,75 @@ pub(super) fn default_provider_quota_execution_timeouts(
}
}
pub(super) fn provider_auto_remove_banned_keys(config: Option<&serde_json::Value>) -> bool {
pub(crate) fn provider_auto_remove_banned_keys(config: Option<&serde_json::Value>) -> bool {
admin_provider_quota_pure::provider_auto_remove_banned_keys(config)
}
pub(super) fn should_auto_remove_oauth_invalid_key(
pub(super) fn should_auto_remove_structured_reason(reason: Option<&str>) -> bool {
admin_provider_quota_pure::should_auto_remove_structured_reason(reason)
}
pub(crate) fn should_auto_remove_oauth_invalid_key(
key: &StoredProviderCatalogKey,
candidate_reason: Option<&str>,
access_token_invalid_proven: bool,
now_unix_secs: u64,
) -> bool {
admin_provider_quota_pure::should_auto_remove_oauth_invalid_key(
key,
candidate_reason,
access_token_invalid_proven,
now_unix_secs,
)
}
pub(crate) async fn persist_quota_oauth_refresh_failure_state(
state: &AdminAppState<'_>,
transport: &AdminGatewayProviderTransportSnapshot,
err: &AdminLocalOAuthRefreshError,
) -> Result<bool, GatewayError> {
let AdminLocalOAuthRefreshError::HttpStatus {
status_code,
body_excerpt,
..
} = err
else {
return Ok(false);
};
if !matches!(*status_code, 400 | 401 | 403) {
return Ok(false);
}
state
.app()
.persist_local_oauth_refresh_failure_state(transport, *status_code, body_excerpt, false)
.await
}
pub(crate) async fn quota_key_auto_removed(
state: &AdminAppState<'_>,
key_id: &str,
) -> Result<bool, GatewayError> {
if key_id.trim().is_empty() {
return Ok(false);
}
Ok(state
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
.await?
.is_empty())
}
pub(crate) fn oauth_refresh_auto_removed_result(
key: &StoredProviderCatalogKey,
) -> serde_json::Value {
serde_json::json!({
"key_id": key.id,
"key_name": key.name,
"status": "auto_removed",
"message": "OAuth refresh 失败且凭证已不可用,已自动删除",
"auto_removed": true,
})
}
pub(crate) fn normalize_string_id_list(values: Option<Vec<String>>) -> Option<Vec<String>> {
admin_provider_quota_pure::normalize_string_id_list(values)
}

View File

@@ -879,6 +879,7 @@ impl AppState {
&current_transport,
status_code,
body_excerpt.as_str(),
false,
)
.await
{
@@ -1212,11 +1213,12 @@ impl AppState {
Ok(())
}
async fn persist_local_oauth_refresh_failure_state(
pub(crate) async fn persist_local_oauth_refresh_failure_state(
&self,
transport: &provider_transport::GatewayProviderTransportSnapshot,
status_code: u16,
body_excerpt: &str,
access_token_invalid_proven: bool,
) -> Result<bool, GatewayError> {
let key_id = transport.key.id.trim();
if key_id.is_empty() {
@@ -1248,44 +1250,71 @@ impl AppState {
) else {
return Ok(false);
};
if latest_key.oauth_invalid_reason.as_deref() == Some(merged_reason.as_str())
&& latest_key.oauth_invalid_at_unix_secs.is_some()
{
return Ok(false);
}
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
latest_key.oauth_invalid_at_unix_secs = latest_key
.oauth_invalid_at_unix_secs
.or(Some(now_unix_secs));
latest_key.oauth_invalid_reason = Some(merged_reason);
latest_key.updated_at_unix_secs = Some(now_unix_secs);
let current_status_snapshot = latest_key.status_snapshot.take();
latest_key.status_snapshot =
sync_provider_key_oauth_status_snapshot(current_status_snapshot, &latest_key);
let mut updated = false;
if latest_key.oauth_invalid_reason.as_deref() != Some(merged_reason.as_str())
|| latest_key.oauth_invalid_at_unix_secs.is_none()
{
latest_key.oauth_invalid_at_unix_secs = latest_key
.oauth_invalid_at_unix_secs
.or(Some(now_unix_secs));
latest_key.oauth_invalid_reason = Some(merged_reason);
latest_key.updated_at_unix_secs = Some(now_unix_secs);
let current_status_snapshot = latest_key.status_snapshot.take();
latest_key.status_snapshot =
sync_provider_key_oauth_status_snapshot(current_status_snapshot, &latest_key);
let updated = self
.update_provider_catalog_key(&latest_key)
.await?
.is_some();
if updated {
self.clear_provider_transport_snapshot_cache();
let _ = self.invalidate_local_oauth_refresh_entry(key_id).await;
updated = self
.update_provider_catalog_key(&latest_key)
.await?
.is_some();
if updated {
self.clear_provider_transport_snapshot_cache();
let _ = self.invalidate_local_oauth_refresh_entry(key_id).await;
}
}
let auto_removed = if admin_provider_quota_pure::provider_auto_remove_banned_keys(
transport.provider.config.as_ref(),
) && admin_provider_quota_pure::should_auto_remove_oauth_invalid_key(
&latest_key,
None,
access_token_invalid_proven,
now_unix_secs,
) {
self.clear_provider_transport_snapshot_cache();
if self.delete_provider_catalog_key(key_id).await? {
let deleted_key_ids = [key_id.to_string()];
self.cleanup_deleted_provider_catalog_refs(
&transport.provider.id,
&[],
&deleted_key_ids,
)
.await?;
let _ = self.invalidate_local_oauth_refresh_entry(key_id).await;
true
} else {
false
}
} else {
false
};
tracing::info!(
key_id = %key_id,
provider_id = %transport.provider.id,
provider_type = %transport.provider.provider_type,
status_code,
updated,
cleared_provider_transport_snapshot_cache = updated,
auto_removed,
cleared_provider_transport_snapshot_cache = updated || auto_removed,
"gateway local oauth refresh failure state persisted"
);
Ok(updated)
Ok(auto_removed)
}
async fn execute_local_oauth_http_request(

View File

@@ -1127,8 +1127,7 @@ async fn gateway_routes_openai_chat_stream_image_intent_to_openai_image_plan_wit
seen_plan.body_json["input"][0]["content"],
"Draw a city made of glass"
);
assert_eq!(seen_plan.body_json["tools"][0]["type"], "image_generation");
assert_eq!(seen_plan.body_json["tools"][0]["size"], "1024x1024");
assert!(seen_plan.body_json.get("tools").is_none());
gateway_handle.abort();
execution_runtime_handle.abort();
@@ -1184,8 +1183,7 @@ async fn gateway_routes_openai_responses_stream_image_intent_to_openai_image_pla
assert_eq!(seen_plan.auth_header, "Bearer sk-upstream-image-bridge");
assert_eq!(seen_plan.body_json["stream"], true);
assert_eq!(seen_plan.body_json["input"], "Draw a mountain observatory");
assert_eq!(seen_plan.body_json["tools"][0]["type"], "image_generation");
assert_eq!(seen_plan.body_json["tools"][0]["size"], "1024x1024");
assert!(seen_plan.body_json.get("tools").is_none());
gateway_handle.abort();
execution_runtime_handle.abort();

View File

@@ -8,7 +8,7 @@ use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogReadRepository, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use axum::body::{to_bytes, Body};
use axum::routing::any;
use axum::routing::{any, post};
use axum::{extract::Request, Json, Router};
use http::StatusCode;
use serde_json::json;
@@ -392,7 +392,7 @@ async fn gateway_marks_codex_quota_exhausted_when_wham_usage_returns_payment_req
}
#[tokio::test]
async fn gateway_auto_removes_codex_key_when_refresh_and_access_tokens_are_invalid() {
async fn gateway_retains_codex_key_when_quota_only_reports_oauth_invalid() {
let upstream = Router::new().route(
"/api/admin/endpoints/providers/provider-codex/refresh-quota",
any(move |_request: Request| async move {
@@ -496,15 +496,19 @@ async fn gateway_auto_removes_codex_key_when_refresh_and_access_tokens_are_inval
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["success"], 0);
assert_eq!(payload["failed"], 1);
assert_eq!(payload["auto_removed"], 1);
assert_eq!(payload["auto_removed"], 0);
assert_eq!(payload["results"][0]["status"], "auth_invalid");
assert_eq!(payload["results"][0]["auto_removed"], true);
assert!(payload["results"][0].get("auto_removed").is_none());
let reloaded = provider_catalog_repository
.list_keys_by_ids(&["key-codex-expired".to_string()])
.await
.expect("keys should read");
assert!(reloaded.is_empty());
assert_eq!(reloaded.len(), 1);
assert!(reloaded[0]
.oauth_invalid_reason
.as_deref()
.is_some_and(|reason| reason.starts_with("[OAUTH_EXPIRED]")));
gateway_handle.abort();
execution_runtime_handle.abort();
@@ -1070,6 +1074,159 @@ async fn gateway_refreshes_admin_provider_quota_locally_for_kiro_with_trusted_ad
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_auto_removes_kiro_quota_refresh_after_terminal_oauth_refresh_failure() {
let token_hits = Arc::new(Mutex::new(0usize));
let token_hits_clone = Arc::clone(&token_hits);
let token_server = Router::new().route(
"/refreshToken",
post(move |_request: Request| {
let token_hits_inner = Arc::clone(&token_hits_clone);
async move {
*token_hits_inner.lock().expect("mutex should lock") += 1;
(
StatusCode::UNAUTHORIZED,
Json(json!({
"message": "refresh token invalid"
})),
)
}
}),
);
let execution_hits = Arc::new(Mutex::new(0usize));
let execution_hits_clone = Arc::clone(&execution_hits);
let execution_runtime = Router::new().route(
"/v1/execute/sync",
any(move |_request: Request| {
let execution_hits_inner = Arc::clone(&execution_hits_clone);
async move {
*execution_hits_inner.lock().expect("mutex should lock") += 1;
(
StatusCode::OK,
Body::from("unexpected execution runtime hit"),
)
}
}),
);
let encrypted_auth_config = encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
&json!({
"provider_type": "kiro",
"auth_method": "social",
"refresh_token": "r".repeat(120),
"machine_id": "123e4567-e89b-12d3-a456-426614174000",
"kiro_version": "1.2.3",
"expires_at": 1u64
})
.to_string(),
)
.expect("auth config ciphertext should build");
let encrypted_api_key =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "__placeholder__")
.expect("api key ciphertext should build");
let mut key = StoredProviderCatalogKey::new(
"key-kiro-expired-refresh-failed".to_string(),
"provider-kiro-auto-remove".to_string(),
"default".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build")
.with_transport_fields(
Some(json!(["claude:messages"])),
encrypted_api_key,
Some(encrypted_auth_config),
None,
None,
None,
None,
None,
None,
)
.expect("key transport should build");
key.expires_at_unix_secs = Some(1);
let mut provider = StoredProviderCatalogProvider::new(
"provider-kiro-auto-remove".to_string(),
"kiro".to_string(),
Some("https://example.com".to_string()),
"kiro".to_string(),
)
.expect("provider should build");
provider.config = Some(json!({
"pool_advanced": {
"auto_remove_banned_keys": true
}
}));
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![sample_endpoint(
"endpoint-kiro-cli",
"provider-kiro-auto-remove",
"claude:messages",
"https://q.us-west-2.amazonaws.com",
)],
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::kiro::KiroOAuthRefreshAdapter::default()
.with_refresh_base_urls(Some(token_url), None),
)
as Arc<dyn crate::provider_transport::oauth_refresh::LocalOAuthRefreshAdapter>,
]);
let gateway = build_router_with_state(
build_state_with_execution_runtime_override(execution_runtime_url.clone())
.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-kiro-auto-remove/refresh-quota"
))
.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")
.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"], 0);
assert_eq!(payload["failed"], 0);
assert_eq!(payload["auto_removed"], 1);
assert_eq!(payload["total"], 1);
assert_eq!(payload["results"][0]["status"], "auto_removed");
assert_eq!(payload["results"][0]["auto_removed"], true);
assert_eq!(*token_hits.lock().expect("mutex should lock"), 1);
assert_eq!(*execution_hits.lock().expect("mutex should lock"), 0);
let reloaded = provider_catalog_repository
.list_keys_by_ids(&["key-kiro-expired-refresh-failed".to_string()])
.await
.expect("keys should read");
assert!(reloaded.is_empty());
gateway_handle.abort();
execution_runtime_handle.abort();
token_handle.abort();
}
#[tokio::test]
async fn gateway_refresh_kiro_quota_reconciles_missing_fixed_endpoint_before_refresh() {
let seen_endpoint_id = Arc::new(Mutex::new(None::<String>));

View File

@@ -5695,6 +5695,128 @@ async fn gateway_marks_manual_oauth_refresh_failures_as_invalid_in_pool_payload(
token_handle.abort();
}
#[tokio::test]
async fn gateway_auto_removes_manual_oauth_refresh_failure_after_access_token_expiry() {
let token_hits = Arc::new(Mutex::new(0usize));
let token_hits_clone = Arc::clone(&token_hits);
let token_server = Router::new().route(
"/oauth/token",
post(move |_request: Request| {
let token_hits_inner = Arc::clone(&token_hits_clone);
async move {
*token_hits_inner.lock().expect("mutex should lock") += 1;
(
StatusCode::UNAUTHORIZED,
Json(json!({
"error": {
"message": "Your refresh token has already been used to generate a new access token. Please try signing in again.",
"type": "invalid_request_error",
"param": serde_json::Value::Null,
"code": "refresh_token_reused"
}
})),
)
}
}),
);
let mut provider = sample_provider("provider-codex", "codex", 10).with_transport_fields(
true,
false,
true,
None,
None,
None,
None,
None,
Some(json!({
"pool_advanced": {
"enabled": true,
"auto_remove_banned_keys": true
}
})),
);
provider.provider_type = "codex".to_string();
let endpoint = sample_endpoint(
"endpoint-codex-cli",
"provider-codex",
"openai:responses",
"https://chatgpt.com/backend-api/codex",
);
let mut key = sample_key(
"key-codex-oauth-refresh-expired",
"provider-codex",
"openai:responses",
"stale-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,
r#"{"provider_type":"codex","refresh_token":"used-refresh-token","email":"alice@example.com","account_id":"acct-codex-123","plan_type":"plus","expires_at":1}"#,
)
.expect("auth config ciphertext should build"),
);
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![key],
));
let (token_url, token_handle) = start_server(token_server).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(
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),
)
.with_oauth_refresh_coordinator_for_tests(oauth_refresh),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let refresh_response = reqwest::Client::new()
.post(format!(
"{gateway_url}/api/admin/provider-oauth/keys/key-codex-oauth-refresh-expired/refresh"
))
.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")
.send()
.await
.expect("refresh request should succeed");
assert_eq!(refresh_response.status(), StatusCode::OK);
let refresh_payload: serde_json::Value = refresh_response
.json()
.await
.expect("refresh payload should parse");
assert_eq!(refresh_payload["status"], json!("auto_removed"));
assert_eq!(*token_hits.lock().expect("mutex should lock"), 1);
let keys = provider_catalog_repository
.list_keys_by_ids(&["key-codex-oauth-refresh-expired".to_string()])
.await
.expect("keys should read");
assert!(keys.is_empty());
gateway_handle.abort();
token_handle.abort();
}
#[tokio::test]
async fn gateway_refreshes_admin_provider_oauth_key_locally_via_execution_runtime_provider_proxy_before_system_proxy(
) {

View File

@@ -38,6 +38,32 @@ fn oauth_reason_has_tag(reason: Option<&str>, tag: &str) -> bool {
})
}
fn oauth_refresh_failure_is_terminal(reason: Option<&str>) -> bool {
reason
.map(str::trim)
.filter(|value| !value.is_empty())
.is_some_and(|reason| {
reason
.lines()
.map(str::trim)
.filter(|line| line.starts_with(OAUTH_REFRESH_FAILED_PREFIX))
.any(|line| {
let lowered = line.to_ascii_lowercase();
lowered.contains("invalid_grant")
|| lowered.contains("invalid_refresh_token")
|| lowered.contains("refresh_token_expired")
|| lowered.contains("could not validate your refresh token")
|| lowered.contains("refresh_token 无效")
|| lowered.contains("已过期或已撤销")
|| lowered.contains("已被使用并轮换")
|| (lowered.contains("refresh token")
&& ["expired", "revoked", "invalid", "reused"]
.iter()
.any(|keyword| lowered.contains(keyword)))
})
})
}
fn oauth_access_token_expired(key: &StoredProviderCatalogKey, now_unix_secs: u64) -> bool {
let now_unix_secs = if now_unix_secs == 0 {
std::time::SystemTime::now()
@@ -55,6 +81,7 @@ fn oauth_access_token_expired(key: &StoredProviderCatalogKey, now_unix_secs: u64
pub fn should_auto_remove_oauth_invalid_key(
key: &StoredProviderCatalogKey,
candidate_reason: Option<&str>,
access_token_invalid_proven: bool,
now_unix_secs: u64,
) -> bool {
if should_auto_remove_structured_reason(candidate_reason)
@@ -71,8 +98,13 @@ pub fn should_auto_remove_oauth_invalid_key(
if !refresh_token_failed {
return false;
}
if !oauth_refresh_failure_is_terminal(candidate_reason)
&& !oauth_refresh_failure_is_terminal(key.oauth_invalid_reason.as_deref())
{
return false;
}
oauth_reason_has_tag(candidate_reason, OAUTH_EXPIRED_PREFIX)
access_token_invalid_proven
|| oauth_reason_has_tag(key.oauth_invalid_reason.as_deref(), OAUTH_EXPIRED_PREFIX)
|| oauth_access_token_expired(key, now_unix_secs)
}
@@ -1243,13 +1275,15 @@ mod tests {
)
.expect("key should build");
key.expires_at_unix_secs = Some(1_000);
key.oauth_invalid_reason = Some(format!("{OAUTH_REFRESH_FAILED_PREFIX}Token 续期失败"));
key.oauth_invalid_reason = Some(format!(
"{OAUTH_REFRESH_FAILED_PREFIX}Token 续期失败 (401): refresh_token 无效、已过期或已撤销,请重新登录授权"
));
assert!(!super::should_auto_remove_oauth_invalid_key(
&key, None, 999
&key, None, false, 999
));
assert!(super::should_auto_remove_oauth_invalid_key(
&key, None, 1_000
&key, None, false, 1_000
));
}
@@ -1265,15 +1299,82 @@ mod tests {
)
.expect("key should build");
key.expires_at_unix_secs = Some(2_000);
key.oauth_invalid_reason = Some(format!("{OAUTH_REFRESH_FAILED_PREFIX}Token 续期失败"));
key.oauth_invalid_reason = Some(format!(
"{OAUTH_REFRESH_FAILED_PREFIX}Token 续期失败 (401): refresh_token 无效、已过期或已撤销,请重新登录授权"
));
assert!(super::should_auto_remove_oauth_invalid_key(
&key, None, true, 1_000,
));
}
#[test]
fn auto_remove_existing_oauth_expired_after_terminal_refresh_failure() {
let mut key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"key-1".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build");
key.expires_at_unix_secs = Some(2_000);
key.oauth_invalid_reason = Some(format!(
"{OAUTH_EXPIRED_PREFIX}access token invalid\n{OAUTH_REFRESH_FAILED_PREFIX}Token 续期失败 (401): refresh_token 无效、已过期或已撤销,请重新登录授权"
));
assert!(super::should_auto_remove_oauth_invalid_key(
&key, None, false, 1_000,
));
}
#[test]
fn candidate_oauth_expired_is_not_auto_remove_proof_by_itself() {
let mut key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"key-1".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build");
key.expires_at_unix_secs = Some(2_000);
key.oauth_invalid_reason = Some(format!(
"{OAUTH_REFRESH_FAILED_PREFIX}Token 续期失败 (401): refresh_token 无效、已过期或已撤销,请重新登录授权"
));
assert!(!super::should_auto_remove_oauth_invalid_key(
&key,
Some("[OAUTH_EXPIRED] access token invalid"),
false,
1_000,
));
}
#[test]
fn oauth_token_invalid_is_not_auto_remove_proof_by_itself() {
let mut key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"key-1".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build");
key.expires_at_unix_secs = Some(1_000);
key.oauth_invalid_reason = Some("oauth_token_invalid".to_string());
assert!(!super::should_auto_remove_oauth_invalid_key(
&key,
Some("oauth_token_invalid"),
false,
1_001,
));
}
#[test]
fn does_not_auto_remove_access_token_failure_without_refresh_failure() {
let mut key = StoredProviderCatalogKey::new(
@@ -1289,7 +1390,26 @@ mod tests {
key.oauth_invalid_reason = Some(format!("{OAUTH_EXPIRED_PREFIX}session expired"));
assert!(!super::should_auto_remove_oauth_invalid_key(
&key, None, 1_001
&key, None, false, 1_001
));
}
#[test]
fn does_not_auto_remove_non_terminal_refresh_failure() {
let mut key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"key-1".to_string(),
"oauth".to_string(),
None,
true,
)
.expect("key should build");
key.expires_at_unix_secs = Some(1_000);
key.oauth_invalid_reason = Some(format!("{OAUTH_REFRESH_FAILED_PREFIX}Token 续期失败"));
assert!(!super::should_auto_remove_oauth_invalid_key(
&key, None, true, 1_001
));
}