mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
fix(admin): 同步 OAuth 刷新后的状态快照有效期 (#300)
手动刷新 OAuth Token 后同步回写 status_snapshot.oauth 修复号池管理页面仍显示旧有效期的问题 补充刷新后快照更新的回归测试
This commit is contained in:
@@ -3,24 +3,179 @@ use super::{
|
|||||||
GatewayError, ProviderTransportSnapshotCacheKey, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_MAX_ENTRIES,
|
GatewayError, ProviderTransportSnapshotCacheKey, PROVIDER_TRANSPORT_SNAPSHOT_CACHE_MAX_ENTRIES,
|
||||||
PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL,
|
PROVIDER_TRANSPORT_SNAPSHOT_CACHE_TTL,
|
||||||
};
|
};
|
||||||
|
use crate::handlers::shared::default_provider_key_status_snapshot;
|
||||||
use crate::provider_transport::LocalOAuthHttpExecutor;
|
use crate::provider_transport::LocalOAuthHttpExecutor;
|
||||||
|
|
||||||
use super::super::provider_transport;
|
use super::super::provider_transport;
|
||||||
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
|
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
|
||||||
|
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||||
use base64::{engine::general_purpose::STANDARD, Engine as _};
|
use base64::{engine::general_purpose::STANDARD, Engine as _};
|
||||||
use flate2::read::{DeflateDecoder, GzDecoder};
|
use flate2::read::{DeflateDecoder, GzDecoder};
|
||||||
|
use serde_json::{json, Map, Value};
|
||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
use std::io::Read;
|
use std::io::Read;
|
||||||
use std::time::Duration;
|
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||||
|
|
||||||
use aether_crypto::encrypt_python_fernet_plaintext;
|
use aether_crypto::encrypt_python_fernet_plaintext;
|
||||||
|
|
||||||
const LOCAL_OAUTH_HTTP_TIMEOUT_MS: u64 = 30_000;
|
const LOCAL_OAUTH_HTTP_TIMEOUT_MS: u64 = 30_000;
|
||||||
|
const OAUTH_ACCOUNT_BLOCK_PREFIX: &str = "[ACCOUNT_BLOCK] ";
|
||||||
|
const OAUTH_EXPIRED_PREFIX: &str = "[OAUTH_EXPIRED] ";
|
||||||
|
const OAUTH_REFRESH_FAILED_PREFIX: &str = "[REFRESH_FAILED] ";
|
||||||
|
const OAUTH_REQUEST_FAILED_PREFIX: &str = "[REQUEST_FAILED] ";
|
||||||
|
|
||||||
struct GatewayLocalOAuthHttpExecutor<'a> {
|
struct GatewayLocalOAuthHttpExecutor<'a> {
|
||||||
state: &'a AppState,
|
state: &'a AppState,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn trimmed_reason(reason: Option<&str>) -> Option<String> {
|
||||||
|
reason
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn tagged_reason(reason: Option<&str>, prefix: &str) -> Option<String> {
|
||||||
|
reason.and_then(|value| {
|
||||||
|
value
|
||||||
|
.lines()
|
||||||
|
.map(str::trim)
|
||||||
|
.find_map(|line| line.strip_prefix(prefix))
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())
|
||||||
|
.map(ToOwned::to_owned)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn default_oauth_status_snapshot_value() -> Value {
|
||||||
|
default_provider_key_status_snapshot()
|
||||||
|
.get("oauth")
|
||||||
|
.cloned()
|
||||||
|
.unwrap_or_else(|| {
|
||||||
|
json!({
|
||||||
|
"code": "none",
|
||||||
|
"label": Value::Null,
|
||||||
|
"reason": Value::Null,
|
||||||
|
"expires_at": Value::Null,
|
||||||
|
"invalid_at": Value::Null,
|
||||||
|
"source": Value::Null,
|
||||||
|
"requires_reauth": false,
|
||||||
|
"expiring_soon": false,
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn build_oauth_status_snapshot_value(key: &StoredProviderCatalogKey) -> Value {
|
||||||
|
if !key.auth_type.trim().eq_ignore_ascii_case("oauth") {
|
||||||
|
return default_oauth_status_snapshot_value();
|
||||||
|
}
|
||||||
|
|
||||||
|
let now_unix_secs = SystemTime::now()
|
||||||
|
.duration_since(UNIX_EPOCH)
|
||||||
|
.ok()
|
||||||
|
.map(|duration| duration.as_secs())
|
||||||
|
.unwrap_or(0);
|
||||||
|
let expires_at_unix_secs = key.expires_at_unix_secs;
|
||||||
|
let invalid_at_unix_secs = key.oauth_invalid_at_unix_secs;
|
||||||
|
let invalid_reason = trimmed_reason(key.oauth_invalid_reason.as_deref());
|
||||||
|
|
||||||
|
if let Some(reason) = tagged_reason(invalid_reason.as_deref(), OAUTH_EXPIRED_PREFIX) {
|
||||||
|
return json!({
|
||||||
|
"code": "invalid",
|
||||||
|
"label": "已失效",
|
||||||
|
"reason": reason,
|
||||||
|
"expires_at": expires_at_unix_secs,
|
||||||
|
"invalid_at": invalid_at_unix_secs,
|
||||||
|
"source": "oauth_invalid",
|
||||||
|
"requires_reauth": true,
|
||||||
|
"expiring_soon": false,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
if let Some(reason) = tagged_reason(invalid_reason.as_deref(), OAUTH_REFRESH_FAILED_PREFIX) {
|
||||||
|
return json!({
|
||||||
|
"code": "invalid",
|
||||||
|
"label": "已失效",
|
||||||
|
"reason": reason,
|
||||||
|
"expires_at": expires_at_unix_secs,
|
||||||
|
"invalid_at": invalid_at_unix_secs,
|
||||||
|
"source": "oauth_refresh",
|
||||||
|
"requires_reauth": true,
|
||||||
|
"expiring_soon": false,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
if let Some(reason) = tagged_reason(invalid_reason.as_deref(), OAUTH_REQUEST_FAILED_PREFIX) {
|
||||||
|
return json!({
|
||||||
|
"code": "check_failed",
|
||||||
|
"label": "检查失败",
|
||||||
|
"reason": reason,
|
||||||
|
"expires_at": expires_at_unix_secs,
|
||||||
|
"invalid_at": Value::Null,
|
||||||
|
"source": "oauth_request",
|
||||||
|
"requires_reauth": false,
|
||||||
|
"expiring_soon": false,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
if invalid_reason
|
||||||
|
.as_deref()
|
||||||
|
.is_some_and(|reason| !reason.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX))
|
||||||
|
|| invalid_at_unix_secs.is_some()
|
||||||
|
{
|
||||||
|
return json!({
|
||||||
|
"code": "invalid",
|
||||||
|
"label": "已失效",
|
||||||
|
"reason": invalid_reason,
|
||||||
|
"expires_at": expires_at_unix_secs,
|
||||||
|
"invalid_at": invalid_at_unix_secs,
|
||||||
|
"source": "oauth_invalid",
|
||||||
|
"requires_reauth": true,
|
||||||
|
"expiring_soon": false,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
let Some(expires_at_unix_secs) = expires_at_unix_secs else {
|
||||||
|
return default_oauth_status_snapshot_value();
|
||||||
|
};
|
||||||
|
if expires_at_unix_secs <= now_unix_secs {
|
||||||
|
return json!({
|
||||||
|
"code": "expired",
|
||||||
|
"label": "已过期",
|
||||||
|
"reason": "Token 已过期,请重新授权",
|
||||||
|
"expires_at": expires_at_unix_secs,
|
||||||
|
"invalid_at": Value::Null,
|
||||||
|
"source": "expires_at",
|
||||||
|
"requires_reauth": true,
|
||||||
|
"expiring_soon": false,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
let expiring_soon = expires_at_unix_secs.saturating_sub(now_unix_secs) < 24 * 60 * 60;
|
||||||
|
json!({
|
||||||
|
"code": if expiring_soon { "expiring" } else { "valid" },
|
||||||
|
"label": if expiring_soon { "即将过期" } else { "有效" },
|
||||||
|
"reason": Value::Null,
|
||||||
|
"expires_at": expires_at_unix_secs,
|
||||||
|
"invalid_at": Value::Null,
|
||||||
|
"source": "expires_at",
|
||||||
|
"requires_reauth": false,
|
||||||
|
"expiring_soon": expiring_soon,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn sync_provider_key_oauth_status_snapshot(
|
||||||
|
status_snapshot: Option<Value>,
|
||||||
|
key: &StoredProviderCatalogKey,
|
||||||
|
) -> Option<Value> {
|
||||||
|
let mut snapshot = status_snapshot
|
||||||
|
.and_then(|value| match value {
|
||||||
|
Value::Object(object) => Some(object),
|
||||||
|
_ => None,
|
||||||
|
})
|
||||||
|
.or_else(|| default_provider_key_status_snapshot().as_object().cloned())
|
||||||
|
.unwrap_or_default();
|
||||||
|
snapshot.insert("oauth".to_string(), build_oauth_status_snapshot_value(key));
|
||||||
|
Some(Value::Object(snapshot))
|
||||||
|
}
|
||||||
|
|
||||||
#[async_trait::async_trait]
|
#[async_trait::async_trait]
|
||||||
impl<'a> provider_transport::LocalOAuthHttpExecutor for GatewayLocalOAuthHttpExecutor<'a> {
|
impl<'a> provider_transport::LocalOAuthHttpExecutor for GatewayLocalOAuthHttpExecutor<'a> {
|
||||||
async fn execute(
|
async fn execute(
|
||||||
@@ -555,13 +710,34 @@ impl AppState {
|
|||||||
.transpose()
|
.transpose()
|
||||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||||
|
|
||||||
self.update_provider_catalog_key_oauth_credentials(
|
let Some(mut latest_key) = self
|
||||||
key_id,
|
.data
|
||||||
encrypted_api_key.as_str(),
|
.list_provider_catalog_keys_by_ids(&[key_id.to_string()])
|
||||||
encrypted_auth_config.as_deref(),
|
.await
|
||||||
entry.expires_at_unix_secs,
|
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||||
)
|
.into_iter()
|
||||||
.await?;
|
.next()
|
||||||
|
else {
|
||||||
|
return Ok(());
|
||||||
|
};
|
||||||
|
|
||||||
|
latest_key.encrypted_api_key = encrypted_api_key;
|
||||||
|
latest_key.encrypted_auth_config = encrypted_auth_config;
|
||||||
|
latest_key.is_active = true;
|
||||||
|
latest_key.expires_at_unix_secs = entry.expires_at_unix_secs;
|
||||||
|
latest_key.oauth_invalid_at_unix_secs = None;
|
||||||
|
latest_key.oauth_invalid_reason = None;
|
||||||
|
latest_key.updated_at_unix_secs = Some(
|
||||||
|
SystemTime::now()
|
||||||
|
.duration_since(UNIX_EPOCH)
|
||||||
|
.ok()
|
||||||
|
.map(|duration| duration.as_secs())
|
||||||
|
.unwrap_or(0),
|
||||||
|
);
|
||||||
|
let current_status_snapshot = latest_key.status_snapshot.take();
|
||||||
|
latest_key.status_snapshot =
|
||||||
|
sync_provider_key_oauth_status_snapshot(current_status_snapshot, &latest_key);
|
||||||
|
self.update_provider_catalog_key(&latest_key).await?;
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -2709,6 +2709,36 @@ async fn gateway_refreshes_admin_provider_oauth_key_locally_with_trusted_admin_p
|
|||||||
key.auth_type = "oauth".to_string();
|
key.auth_type = "oauth".to_string();
|
||||||
key.oauth_invalid_at_unix_secs = Some(1_700_000_000);
|
key.oauth_invalid_at_unix_secs = Some(1_700_000_000);
|
||||||
key.oauth_invalid_reason = Some("[REFRESH_FAILED] stale token".to_string());
|
key.oauth_invalid_reason = Some("[REFRESH_FAILED] stale token".to_string());
|
||||||
|
key.status_snapshot = Some(json!({
|
||||||
|
"oauth": {
|
||||||
|
"code": "expired",
|
||||||
|
"label": "已过期",
|
||||||
|
"reason": "Token 已过期,请重新授权",
|
||||||
|
"expires_at": 1u64,
|
||||||
|
"invalid_at": 1_700_000_000u64,
|
||||||
|
"source": "expires_at",
|
||||||
|
"requires_reauth": true,
|
||||||
|
"expiring_soon": false
|
||||||
|
},
|
||||||
|
"account": {
|
||||||
|
"code": "ok",
|
||||||
|
"label": null,
|
||||||
|
"reason": null,
|
||||||
|
"blocked": false,
|
||||||
|
"source": null,
|
||||||
|
"recoverable": false
|
||||||
|
},
|
||||||
|
"quota": {
|
||||||
|
"code": "unknown",
|
||||||
|
"label": null,
|
||||||
|
"reason": null,
|
||||||
|
"exhausted": false,
|
||||||
|
"usage_ratio": null,
|
||||||
|
"updated_at": null,
|
||||||
|
"reset_seconds": null,
|
||||||
|
"plan_type": null
|
||||||
|
}
|
||||||
|
}));
|
||||||
key.encrypted_auth_config = Some(
|
key.encrypted_auth_config = Some(
|
||||||
encrypt_python_fernet_plaintext(
|
encrypt_python_fernet_plaintext(
|
||||||
DEVELOPMENT_ENCRYPTION_KEY,
|
DEVELOPMENT_ENCRYPTION_KEY,
|
||||||
@@ -2827,6 +2857,44 @@ async fn gateway_refreshes_admin_provider_oauth_key_locally_with_trusted_admin_p
|
|||||||
auth_config["email"],
|
auth_config["email"],
|
||||||
serde_json::Value::String("alice@example.com".to_string())
|
serde_json::Value::String("alice@example.com".to_string())
|
||||||
);
|
);
|
||||||
|
let status_snapshot = stored_key
|
||||||
|
.status_snapshot
|
||||||
|
.as_ref()
|
||||||
|
.and_then(serde_json::Value::as_object)
|
||||||
|
.expect("status snapshot should exist");
|
||||||
|
let oauth_snapshot = status_snapshot
|
||||||
|
.get("oauth")
|
||||||
|
.and_then(serde_json::Value::as_object)
|
||||||
|
.expect("oauth snapshot should exist");
|
||||||
|
assert_eq!(
|
||||||
|
oauth_snapshot
|
||||||
|
.get("code")
|
||||||
|
.and_then(serde_json::Value::as_str),
|
||||||
|
Some("expiring")
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
oauth_snapshot
|
||||||
|
.get("label")
|
||||||
|
.and_then(serde_json::Value::as_str),
|
||||||
|
Some("即将过期")
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
oauth_snapshot.get("expires_at"),
|
||||||
|
auth_config.get("expires_at")
|
||||||
|
);
|
||||||
|
assert_eq!(oauth_snapshot.get("reason"), Some(&serde_json::Value::Null));
|
||||||
|
assert_eq!(
|
||||||
|
oauth_snapshot
|
||||||
|
.get("requires_reauth")
|
||||||
|
.and_then(serde_json::Value::as_bool),
|
||||||
|
Some(false)
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
oauth_snapshot
|
||||||
|
.get("expiring_soon")
|
||||||
|
.and_then(serde_json::Value::as_bool),
|
||||||
|
Some(true)
|
||||||
|
);
|
||||||
|
|
||||||
gateway_handle.abort();
|
gateway_handle.abort();
|
||||||
token_handle.abort();
|
token_handle.abort();
|
||||||
|
|||||||
Reference in New Issue
Block a user