mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
fix(oauth): 修复 OAuth 刷新后 Token 有效期不更新问题 (#324)
- 持久化刷新后的 `expires_at` 到 Provider Key SQL 更新链路 - 手动刷新接口优先返回本次刷新得到的过期时间 - 按当前 Key 字段重建 OAuth 状态快照,避免旧快照覆盖 - 前端刷新后防止旧列表数据回退覆盖新有效期
This commit is contained in:
@@ -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,
|
||||||
}))
|
}))
|
||||||
|
|||||||
@@ -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>,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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),
|
||||||
|
|||||||
@@ -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()?
|
||||||
|
|||||||
@@ -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({
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user