fix(oauth): 系统导入时清理失效标记并强制刷新 refresh_token,前端徽章优先使用 account_id

- 后端 admin 系统导入 OAuth 凭据时重置 expires_at 与失效标记,并在存在 refresh_token 时触发一次本地刷新
- 前端 OAuth 徽章改为优先展示 account_id,采用去前缀后 5 字符的紧凑格式
This commit is contained in:
fawney19
2026-04-21 18:06:39 +08:00
parent 25a2b417be
commit 063ef02306
4 changed files with 319 additions and 17 deletions

View File

@@ -281,6 +281,83 @@ fn apply_imported_oauth_key_credentials(
};
}
// Importing OAuth credentials replaces the previous session state, so stale
// expiry/invalid markers must not survive across the overwrite.
record.expires_at_unix_secs = imported_oauth_expires_at_unix_secs(normalized_auth_config);
record.oauth_invalid_at_unix_secs = None;
record.oauth_invalid_reason = None;
Ok(())
}
fn imported_oauth_expires_at_unix_secs(normalized_auth_config: Option<&Value>) -> Option<u64> {
let object = normalized_auth_config?.as_object()?;
for field in ["expires_at", "expiresAt", "expiry", "exp"] {
let Some(value) = object.get(field) else {
continue;
};
match value {
Value::Number(number) => {
if let Some(expires_at) = number.as_u64() {
return Some(expires_at);
}
}
Value::String(raw) => {
if let Ok(expires_at) = raw.trim().parse::<u64>() {
return Some(expires_at);
}
}
_ => {}
}
}
None
}
fn imported_oauth_has_refresh_token(normalized_auth_config: Option<&Value>) -> bool {
normalized_auth_config
.and_then(Value::as_object)
.and_then(|object| object.get("refresh_token"))
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty())
}
async fn refresh_imported_oauth_key_after_persist(
state: &AdminAppState<'_>,
provider: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider,
key_id: &str,
) -> Result<(), GatewayError> {
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let Some(endpoint) =
crate::handlers::admin::provider::oauth::runtime::provider_oauth_runtime_endpoint_for_provider(
provider.provider_type.as_str(),
&endpoints,
)
else {
return Ok(());
};
let Some(transport) = state
.read_provider_transport_snapshot(&provider.id, &endpoint.id, key_id)
.await?
else {
return Ok(());
};
if !crate::provider_transport::supports_local_oauth_request_auth_resolution(&transport) {
return Ok(());
}
if let Err(error) = state.force_local_oauth_refresh_entry(&transport).await {
tracing::warn!(
provider_id = %provider.id,
provider_type = %provider.provider_type,
key_id = %key_id,
error = ?error,
"admin system import oauth refresh after credential import failed"
);
}
Ok(())
}
@@ -1130,6 +1207,16 @@ impl<'a> AdminAppState<'a> {
"更新 Provider '{provider_name}' 的 Key 失败"
))));
};
if auth_type == "oauth"
&& imported_oauth_has_refresh_token(normalized_auth_config.as_ref())
{
refresh_imported_oauth_key_after_persist(
self,
&provider,
&persisted.id,
)
.await?;
}
existing_keys[existing_index] = persisted;
stats.keys.updated += 1;
}
@@ -1170,6 +1257,11 @@ impl<'a> AdminAppState<'a> {
"创建 Provider '{provider_name}' 的 Key 失败"
))));
};
if auth_type == "oauth"
&& imported_oauth_has_refresh_token(normalized_auth_config.as_ref())
{
refresh_imported_oauth_key_after_persist(self, &provider, &created.id).await?;
}
existing_keys.push(created);
stats.keys.created += 1;
}

View File

@@ -1,6 +1,8 @@
use std::sync::{Arc, Mutex};
use aether_crypto::{decrypt_python_fernet_ciphertext, DEVELOPMENT_ENCRYPTION_KEY};
use aether_crypto::{
decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY,
};
use aether_data::repository::auth::InMemoryAuthApiKeySnapshotRepository;
use aether_data::repository::auth_modules::{
AuthModuleReadRepository, InMemoryAuthModuleReadRepository, StoredOAuthProviderModuleConfig,
@@ -17,12 +19,14 @@ use aether_data_contracts::repository::global_models::{
StoredPublicGlobalModel,
};
use aether_data_contracts::repository::provider_catalog::ProviderCatalogReadRepository;
use axum::body::Body;
use axum::routing::any;
use axum::body::{Body, Bytes};
use axum::http::HeaderMap;
use axum::routing::{any, post};
use axum::{extract::Request, Router};
use http::StatusCode;
use serde_json::{json, Value};
use super::super::helpers::{sample_endpoint, sample_key, sample_provider};
use super::super::{build_router_with_state, start_server, AppState};
use crate::constants::{
GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER,
@@ -1119,6 +1123,190 @@ async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_imp
gateway_handle.abort();
}
#[tokio::test]
async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_import_and_forces_refresh(
) {
#[derive(Debug, Clone)]
struct SeenRefreshRequest {
content_type: String,
body: String,
}
let seen_refresh = Arc::new(Mutex::new(None::<SeenRefreshRequest>));
let seen_refresh_clone = Arc::clone(&seen_refresh);
let refresh_hits = Arc::new(Mutex::new(0usize));
let refresh_hits_clone = Arc::clone(&refresh_hits);
let refresh_server = Router::new().route(
"/oauth/token",
post(move |headers: HeaderMap, body: Bytes| {
let seen_refresh_inner = Arc::clone(&seen_refresh_clone);
let refresh_hits_inner = Arc::clone(&refresh_hits_clone);
async move {
*refresh_hits_inner.lock().expect("mutex should lock") += 1;
*seen_refresh_inner.lock().expect("mutex should lock") = Some(SeenRefreshRequest {
content_type: headers
.get(http::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string(),
body: String::from_utf8(body.to_vec()).unwrap_or_default(),
});
axum::Json(json!({
"access_token": "oauth-access-token-refreshed",
"refresh_token": "oauth-refresh-token-refreshed",
"token_type": "Bearer",
"expires_in": 3600
}))
}
}),
);
let mut provider = sample_provider("provider-codex-existing", "oauth-import-provider", 10);
provider.provider_type = "codex".to_string();
let endpoint = sample_endpoint(
"endpoint-codex-existing",
"provider-codex-existing",
"openai:cli",
"https://chatgpt.com/backend-api/codex",
);
let mut existing_key = sample_key(
"key-codex-existing",
"provider-codex-existing",
"openai:cli",
"oauth-access-token-old",
);
existing_key.name = "oauth-primary".to_string();
existing_key.auth_type = "oauth".to_string();
existing_key.expires_at_unix_secs = Some(1);
existing_key.oauth_invalid_at_unix_secs = Some(1_700_000_000);
existing_key.oauth_invalid_reason =
Some("[REFRESH_FAILED] refresh_token 无效、已过期或已撤销,请重新登录授权".to_string());
existing_key.encrypted_auth_config = Some(
encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
r#"{"provider_type":"codex","refresh_token":"oauth-refresh-token-old","email":"alice@example.com","account_id":"acct-codex-123","plan_type":"plus","expires_at":1}"#,
)
.expect("auth config should encrypt"),
);
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![endpoint],
vec![existing_key],
));
let global_model_repository = Arc::new(InMemoryGlobalModelReadRepository::seed(Vec::<
StoredPublicGlobalModel,
>::new()));
let auth_module_repository = Arc::new(InMemoryAuthModuleReadRepository::seed(
Vec::<StoredOAuthProviderModuleConfig>::new(),
None,
));
let oauth_provider_repository = Arc::new(InMemoryOAuthProviderRepository::seed(Vec::<
StoredOAuthProviderConfig,
>::new()));
let (refresh_url, refresh_handle) = start_server(refresh_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!("{refresh_url}/oauth/token")),
),
]);
let data_state = GatewayDataState::with_provider_catalog_repository_for_tests(Arc::clone(
&provider_catalog_repository,
))
.with_global_model_repository_for_tests(Arc::clone(&global_model_repository))
.attach_auth_module_repository_for_tests(Arc::clone(&auth_module_repository))
.attach_oauth_provider_repository_for_tests(Arc::clone(&oauth_provider_repository))
.with_system_config_values_for_tests(Vec::<(String, Value)>::new())
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
let gateway = build_router_with_state(
AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(data_state)
.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/system/config/import"))
.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(&sample_oauth_system_import_payload(
"oauth-access-token-new",
"oauth-refresh-token-new",
))
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(*refresh_hits.lock().expect("mutex should lock"), 1);
let seen_refresh = seen_refresh
.lock()
.expect("mutex should lock")
.clone()
.expect("refresh request should be captured");
assert_eq!(
seen_refresh.content_type,
"application/x-www-form-urlencoded"
);
assert!(seen_refresh.body.contains("grant_type=refresh_token"));
assert!(seen_refresh
.body
.contains("refresh_token=oauth-refresh-token-new"));
let providers = provider_catalog_repository
.list_providers(false)
.await
.expect("providers should load");
assert_eq!(providers.len(), 1);
let keys = provider_catalog_repository
.list_keys_by_provider_ids(std::slice::from_ref(&providers[0].id))
.await
.expect("keys should load");
assert_eq!(keys.len(), 1);
let key = &keys[0];
assert_eq!(key.name, "oauth-primary");
assert_eq!(key.oauth_invalid_at_unix_secs, None);
assert_eq!(key.oauth_invalid_reason, None);
assert!(key.expires_at_unix_secs.is_some());
assert_eq!(
decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, &key.encrypted_api_key)
.expect("oauth access token should decrypt"),
"oauth-access-token-refreshed"
);
let auth_config = decrypt_python_fernet_ciphertext(
DEVELOPMENT_ENCRYPTION_KEY,
key.encrypted_auth_config
.as_deref()
.expect("oauth auth config should exist"),
)
.expect("oauth auth config should decrypt");
let auth_config: Value =
serde_json::from_str(&auth_config).expect("oauth auth config json should parse");
assert_eq!(auth_config["provider_type"], "codex");
assert_eq!(
auth_config["refresh_token"],
"oauth-refresh-token-refreshed"
);
assert_eq!(auth_config["email"], "alice@example.com");
assert_eq!(auth_config["account_id"], "acct-codex-123");
assert_eq!(auth_config["plan_type"], "plus");
assert_eq!(auth_config["token_type"], "Bearer");
assert_eq!(auth_config["expires_at"].as_u64(), key.expires_at_unix_secs);
gateway_handle.abort();
refresh_handle.abort();
}
#[tokio::test]
async fn gateway_skips_proxy_nodes_during_admin_system_config_import() {
let gateway = build_router_with_state(