fix(oauth): 修复 OAuth 刷新后 Token 有效期不更新问题 (#324)

- 持久化刷新后的 `expires_at` 到 Provider Key SQL 更新链路
- 手动刷新接口优先返回本次刷新得到的过期时间
- 按当前 Key 字段重建 OAuth 状态快照,避免旧快照覆盖
- 前端刷新后防止旧列表数据回退覆盖新有效期
This commit is contained in:
AAEE86
2026-04-24 09:35:25 +08:00
committed by GitHub
parent 31e871fe1b
commit a9d10163af
8 changed files with 70 additions and 18 deletions

View File

@@ -24,8 +24,8 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
transport, transport,
} = request; } = request;
match state.force_local_oauth_refresh_entry(&transport).await { let refreshed_entry = match state.force_local_oauth_refresh_entry(&transport).await {
Ok(Some(_)) => {} Ok(Some(entry)) => Some(entry),
Ok(None) => { Ok(None) => {
return Ok(RefreshDispatch::Respond(response::control_error_response( return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::BAD_REQUEST, http::StatusCode::BAD_REQUEST,
@@ -75,7 +75,7 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
response::oauth_refresh_failed_bad_request_response(&message), response::oauth_refresh_failed_bad_request_response(&message),
)); ));
} }
} };
if !helpers::key_is_account_blocked(&key, OAUTH_ACCOUNT_BLOCK_PREFIX) { if !helpers::key_is_account_blocked(&key, OAUTH_ACCOUNT_BLOCK_PREFIX) {
let _ = state let _ = state
@@ -89,10 +89,25 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
.into_iter() .into_iter()
.next() .next()
.unwrap_or(key); .unwrap_or(key);
let refreshed_auth_config = helpers::refreshed_auth_config_object( let refreshed_auth_config = refreshed_entry
state, .as_ref()
refreshed_key.encrypted_auth_config.as_deref(), .and_then(|entry| entry.metadata.as_ref())
); .and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_else(|| {
helpers::refreshed_auth_config_object(
state,
refreshed_key.encrypted_auth_config.as_deref(),
)
});
let refreshed_expires_at_unix_secs = refreshed_entry
.as_ref()
.and_then(|entry| entry.expires_at_unix_secs)
.or_else(|| {
refreshed_auth_config
.get("expires_at")
.and_then(serde_json::Value::as_u64)
});
let (account_state_recheck_attempted, account_state_recheck_error) = state let (account_state_recheck_attempted, account_state_recheck_error) = state
.refresh_provider_oauth_account_state_after_update(&provider, &key_id, None) .refresh_provider_oauth_account_state_after_update(&provider, &key_id, None)
.await?; .await?;
@@ -100,6 +115,7 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
Ok(RefreshDispatch::Continue(RefreshSuccessContext { Ok(RefreshDispatch::Continue(RefreshSuccessContext {
provider_type, provider_type,
refreshed_auth_config, refreshed_auth_config,
refreshed_expires_at_unix_secs,
account_state_recheck_attempted, account_state_recheck_attempted,
account_state_recheck_error, account_state_recheck_error,
})) }))

View File

@@ -23,6 +23,7 @@ pub(super) struct RefreshRequestContext {
pub(super) struct RefreshSuccessContext { pub(super) struct RefreshSuccessContext {
pub(super) provider_type: String, pub(super) provider_type: String,
pub(super) refreshed_auth_config: Map<String, Value>, pub(super) refreshed_auth_config: Map<String, Value>,
pub(super) refreshed_expires_at_unix_secs: Option<u64>,
pub(super) account_state_recheck_attempted: bool, pub(super) account_state_recheck_attempted: bool,
pub(super) account_state_recheck_error: Option<String>, pub(super) account_state_recheck_error: Option<String>,
} }

View File

@@ -36,13 +36,14 @@ pub(super) fn oauth_refresh_failed_service_unavailable_response(
pub(super) fn admin_provider_oauth_refresh_success_response( pub(super) fn admin_provider_oauth_refresh_success_response(
success: RefreshSuccessContext, success: RefreshSuccessContext,
) -> Response<Body> { ) -> Response<Body> {
let expires_at = success
.refreshed_expires_at_unix_secs
.map(serde_json::Value::from)
.or_else(|| success.refreshed_auth_config.get("expires_at").cloned())
.unwrap_or(Value::Null);
Json(json!({ Json(json!({
"provider_type": success.provider_type, "provider_type": success.provider_type,
"expires_at": success "expires_at": expires_at,
.refreshed_auth_config
.get("expires_at")
.cloned()
.unwrap_or(Value::Null),
"has_refresh_token": success "has_refresh_token": success
.refreshed_auth_config .refreshed_auth_config
.get("refresh_token") .get("refresh_token")

View File

@@ -121,6 +121,10 @@ fn admin_pool_derive_oauth_expires_at(
return None; return None;
} }
if key.expires_at_unix_secs.is_some() {
return key.expires_at_unix_secs;
}
for field in ["expires_at", "expiresAt", "expiry", "exp"] { for field in ["expires_at", "expiresAt", "expiry", "exp"] {
let expires_at = admin_pool_json_to_u64(auth_config.and_then(|config| config.get(field))); let expires_at = admin_pool_json_to_u64(auth_config.and_then(|config| config.get(field)));
if expires_at.is_some() { if expires_at.is_some() {
@@ -128,7 +132,7 @@ fn admin_pool_derive_oauth_expires_at(
} }
} }
key.expires_at_unix_secs None
} }
fn admin_pool_derive_oauth_plan_type( fn admin_pool_derive_oauth_plan_type(

View File

@@ -1023,6 +1023,10 @@ pub(crate) fn provider_key_status_snapshot_payload(
let mut snapshot = provider_key_status_snapshot_object(Some(&payload)) let mut snapshot = provider_key_status_snapshot_object(Some(&payload))
.or_else(|| default_provider_key_status_snapshot().as_object().cloned()) .or_else(|| default_provider_key_status_snapshot().as_object().cloned())
.unwrap_or_default(); .unwrap_or_default();
snapshot.insert(
"oauth".to_string(),
build_provider_key_oauth_status_snapshot(key),
);
snapshot.insert( snapshot.insert(
"account".to_string(), "account".to_string(),
build_provider_key_account_status_snapshot(key, provider_type), build_provider_key_account_status_snapshot(key, provider_type),

View File

@@ -1366,6 +1366,7 @@ INSERT INTO provider_api_keys (
.bind(key.is_active) .bind(key.is_active)
.bind(key.created_at_unix_ms.map(|value| value as f64)) .bind(key.created_at_unix_ms.map(|value| value as f64))
.bind(key.updated_at_unix_secs.map(|value| value as f64)) .bind(key.updated_at_unix_secs.map(|value| value as f64))
.bind(key.expires_at_unix_secs.map(|value| value as f64))
.execute(&self.pool) .execute(&self.pool)
.await .await
.map_postgres_err()?; .map_postgres_err()?;
@@ -1735,6 +1736,10 @@ SET
proxy = $22, proxy = $22,
fingerprint = $23, fingerprint = $23,
upstream_metadata = $24, upstream_metadata = $24,
expires_at = CASE
WHEN $38::double precision IS NULL THEN NULL
ELSE TO_TIMESTAMP($38::double precision)
END,
oauth_invalid_at = CASE oauth_invalid_at = CASE
WHEN $25::double precision IS NULL THEN NULL WHEN $25::double precision IS NULL THEN NULL
ELSE TO_TIMESTAMP($25::double precision) ELSE TO_TIMESTAMP($25::double precision)
@@ -1803,6 +1808,7 @@ WHERE id = $1
.bind(key.last_rpm_peak.map(|value| value as i32)) .bind(key.last_rpm_peak.map(|value| value as i32))
.bind(key.is_active) .bind(key.is_active)
.bind(key.updated_at_unix_secs.map(|value| value as f64)) .bind(key.updated_at_unix_secs.map(|value| value as f64))
.bind(key.expires_at_unix_secs.map(|value| value as f64))
.execute(&self.pool) .execute(&self.pool)
.await .await
.map_postgres_err()? .map_postgres_err()?

View File

@@ -1703,19 +1703,28 @@ async function handleRefreshOAuth(key: EndpointAPIKey) {
refreshingOAuthKeyId.value = key.id refreshingOAuthKeyId.value = key.id
try { try {
const result = await refreshProviderOAuth(key.id) const result = await refreshProviderOAuth(key.id)
const refreshedExpiresAt = typeof result.expires_at === 'number' ? result.expires_at : null
let refreshedKey: EndpointAPIKey | null = null let refreshedKey: EndpointAPIKey | null = null
// 更新本地数据 // 更新本地数据
const keyInList = providerKeys.value.find(k => k.id === key.id) const keyInList = providerKeys.value.find(k => k.id === key.id)
if (keyInList) { if (keyInList) {
keyInList.oauth_expires_at = result.expires_at keyInList.oauth_expires_at = refreshedExpiresAt
} }
// 只重新加载 keys 数据,避免整个表格刷新 // 只重新加载 keys 数据,避免整个表格刷新
if (props.providerId) { if (props.providerId) {
const freshKeys = await getProviderKeys(props.providerId).catch(() => null) const freshKeys = await getProviderKeys(props.providerId).catch(() => null)
if (freshKeys) { if (freshKeys) {
providerKeys.value = freshKeys const mergedKeys = freshKeys.map((item) => {
syncCurrentSelections(endpoints.value, freshKeys) if (item.id !== key.id) return item
refreshedKey = freshKeys.find(item => item.id === key.id) ?? null if (refreshedExpiresAt == null) return item
if (typeof item.oauth_expires_at === 'number' && item.oauth_expires_at >= refreshedExpiresAt) {
return item
}
return { ...item, oauth_expires_at: refreshedExpiresAt }
})
providerKeys.value = mergedKeys
syncCurrentSelections(endpoints.value, mergedKeys)
refreshedKey = mergedKeys.find(item => item.id === key.id) ?? null
} }
} }
const feedback = getOAuthRefreshFeedback({ const feedback = getOAuthRefreshFeedback({

View File

@@ -2180,11 +2180,22 @@ async function handleRefreshOAuth(key: PoolKeyDetail) {
refreshingOAuthKeyId.value = key.key_id refreshingOAuthKeyId.value = key.key_id
try { try {
const result = await refreshProviderOAuth(key.key_id) const result = await refreshProviderOAuth(key.key_id)
const refreshedExpiresAt = typeof result.expires_at === 'number' ? result.expires_at : null
const target = keyPage.value.keys.find(k => k.key_id === key.key_id) const target = keyPage.value.keys.find(k => k.key_id === key.key_id)
if (target) { if (target) {
target.oauth_expires_at = result.expires_at ?? null target.oauth_expires_at = refreshedExpiresAt
} }
await loadKeys() await loadKeys()
if (refreshedExpiresAt != null) {
const reloadedTarget = keyPage.value.keys.find(k => k.key_id === key.key_id)
if (
reloadedTarget
&& (typeof reloadedTarget.oauth_expires_at !== 'number'
|| reloadedTarget.oauth_expires_at < refreshedExpiresAt)
) {
reloadedTarget.oauth_expires_at = refreshedExpiresAt
}
}
const refreshedKey = keyPage.value.keys.find(k => k.key_id === key.key_id) ?? null const refreshedKey = keyPage.value.keys.find(k => k.key_id === key.key_id) ?? null
const feedback = getOAuthRefreshFeedback({ const feedback = getOAuthRefreshFeedback({
accountStateRecheckAttempted: result.account_state_recheck_attempted, accountStateRecheckAttempted: result.account_state_recheck_attempted,