mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
fix: tighten OAuth auto cleanup signals
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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> {
|
||||
|
||||
@@ -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,
|
||||
})))
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
})))
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
})))
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -879,6 +879,7 @@ impl AppState {
|
||||
¤t_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(
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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>));
|
||||
|
||||
@@ -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(
|
||||
) {
|
||||
|
||||
Reference in New Issue
Block a user