Merge remote-tracking branch 'origin/main' into codex/provider-policy-hardening

This commit is contained in:
elky
2026-09-05 16:21:10 +08:00
1228 changed files with 241230 additions and 31814 deletions
@@ -8,16 +8,13 @@ use crate::handlers::public::{
build_api_key_install_session_response, CreateApiKeyInstallSessionRequest,
};
use crate::GatewayError;
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
};
use axum::{body::Body, http, response::Response};
pub(super) async fn build_admin_create_api_key_install_session_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_headers: &http::HeaderMap,
remote_addr: &std::net::SocketAddr,
request_body: Option<&axum::body::Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_auth_api_key_data_reader() {
@@ -48,40 +45,26 @@ pub(super) async fn build_admin_create_api_key_install_session_response(
else {
return Ok(build_admin_api_keys_not_found_response());
};
let Some(ciphertext) = record
.key_encrypted
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(build_admin_api_keys_bad_request_response(
"该密钥没有存储完整密钥信息",
));
};
let Some(api_key) = state.decrypt_catalog_secret_with_fallbacks(ciphertext) else {
return Ok((
http::StatusCode::INTERNAL_SERVER_ERROR,
axum::Json(serde_json::json!({ "detail": "解密密钥失败" })),
)
.into_response());
};
let response = build_api_key_install_session_response(
state.app(),
request_context.public(),
request_headers,
record.api_key_id.clone(),
record.name.unwrap_or_else(|| "API Key".to_string()),
api_key,
remote_addr,
&record,
payload,
)
.await;
Ok(attach_admin_audit_response(
response,
"admin_standalone_api_key_install_session_created",
"create_standalone_api_key_install_session",
"api_key",
&api_key_id,
))
if response.status().is_success() {
Ok(attach_admin_audit_response(
response,
"admin_standalone_api_key_install_session_created",
"create_standalone_api_key_install_session",
"api_key",
&api_key_id,
))
} else {
Ok(response)
}
}
@@ -6,8 +6,7 @@ use super::super::users::{
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{
decrypt_catalog_secret_with_fallbacks, encrypt_catalog_secret_with_fallbacks, query_param_bool,
query_param_optional_bool, query_param_value,
query_param_bool, query_param_optional_bool, query_param_value,
};
use crate::GatewayError;
use axum::{
@@ -42,12 +41,14 @@ pub(crate) async fn maybe_build_local_admin_api_keys_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_headers: &http::HeaderMap,
remote_addr: &std::net::SocketAddr,
request_body: Option<&axum::body::Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
routes::maybe_build_local_admin_api_keys_routes_response(
state,
request_context,
request_headers,
remote_addr,
request_body,
)
.await
@@ -5,7 +5,9 @@ use super::shared::{
AdminStandaloneApiKeyToggleRequest, AdminStandaloneApiKeyUpdatePatch,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::handlers::admin::shared::{
attach_admin_audit_response, mark_sensitive_admin_response_no_store,
};
use crate::handlers::admin::users::{
default_admin_user_api_key_name, format_optional_unix_secs_iso8601,
generate_admin_user_api_key_plaintext, hash_admin_user_api_key, masked_user_api_key_display,
@@ -13,7 +15,9 @@ use crate::handlers::admin::users::{
normalize_admin_user_api_formats, normalize_admin_user_ip_rules,
normalize_admin_user_string_list,
};
use crate::handlers::shared::normalize_optional_api_key_concurrent_limit;
use crate::handlers::shared::{
normalize_optional_api_key_concurrent_limit, seal_auth_api_key_secret,
};
use crate::GatewayError;
use aether_admin::system::serialize_admin_system_users_export_wallet;
use axum::{
@@ -60,6 +64,48 @@ fn normalize_standalone_initial_balance(
Ok((initial_balance_usd, false))
}
/// Compensate the rows created by a standalone API-key request when wallet
/// provisioning cannot complete. The wallet is deleted first, and only when
/// its exact API-key owner and untouched state still match. If a wallet has
/// become funded or otherwise referenced, preserve both rows rather than
/// deleting the key and leaving an orphaned financial account.
async fn compensate_standalone_api_key_creation(
state: &AdminAppState<'_>,
api_key_id: &str,
) -> Result<(), GatewayError> {
let wallet = state
.find_wallet(aether_data::repository::wallet::WalletLookupKey::ApiKeyId(
api_key_id,
))
.await?;
if let Some(wallet) = wallet {
let removed = state
.delete_wallet_if_unreferenced(
&wallet.id,
aether_data::repository::wallet::WalletLookupKey::ApiKeyId(api_key_id),
)
.await?;
if !removed
&& state
.find_wallet(aether_data::repository::wallet::WalletLookupKey::WalletId(
&wallet.id,
))
.await?
.is_some()
{
return Err(GatewayError::Internal(format!(
"refusing to delete standalone API key {api_key_id}: wallet is still referenced"
)));
}
}
// Treat an already-removed key as an idempotent successful cleanup. The
// wallet owner check above prevents deleting a key while its funded wallet
// remains attached.
let _ = state.delete_standalone_api_key(api_key_id).await?;
Ok(())
}
pub(super) async fn build_admin_create_api_key_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -148,7 +194,16 @@ pub(super) async fn build_admin_create_api_key_response(
};
let plaintext_key = generate_admin_user_api_key_plaintext();
let Some(key_encrypted) = state.encrypt_catalog_secret_with_fallbacks(&plaintext_key) else {
let api_key_id = uuid::Uuid::new_v4().to_string();
let key_hash = hash_admin_user_api_key(&plaintext_key);
let Ok(key_encrypted) = seal_auth_api_key_secret(
state.app(),
&operator_id,
&api_key_id,
&key_hash,
true,
&plaintext_key,
) else {
return Ok((
http::StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({ "detail": "API密钥加密失败" })),
@@ -160,8 +215,8 @@ pub(super) async fn build_admin_create_api_key_response(
.create_standalone_api_key(
aether_data::repository::auth::CreateStandaloneApiKeyRecord {
user_id: operator_id,
api_key_id: uuid::Uuid::new_v4().to_string(),
key_hash: hash_admin_user_api_key(&plaintext_key),
api_key_id,
key_hash,
key_encrypted: Some(key_encrypted),
name: Some(name),
allowed_providers,
@@ -184,11 +239,39 @@ pub(super) async fn build_admin_create_api_key_response(
return Ok(build_admin_api_keys_data_unavailable_response());
};
let wallet = match state
.initialize_auth_api_key_wallet(&created.api_key_id, initial_balance_usd, unlimited_balance)
.await?
.initialize_auth_api_key_wallet_with_outcome(
&created.api_key_id,
initial_balance_usd,
unlimited_balance,
)
.await
{
Some(wallet) => wallet,
None => return Ok(build_admin_api_keys_data_unavailable_response()),
Ok(Some(initialized)) => initialized.wallet,
Ok(None) => {
if let Err(error) =
compensate_standalone_api_key_creation(state, &created.api_key_id).await
{
tracing::error!(
api_key_id = %created.api_key_id,
error = ?error,
"standalone API key wallet provisioning cleanup failed"
);
return Err(error);
}
return Ok(build_admin_api_keys_data_unavailable_response());
}
Err(error) => {
if let Err(cleanup_error) =
compensate_standalone_api_key_creation(state, &created.api_key_id).await
{
tracing::error!(
api_key_id = %created.api_key_id,
error = ?cleanup_error,
"standalone API key wallet provisioning cleanup failed"
);
}
return Err(error);
}
};
let created = if feature_settings.is_some() {
state
@@ -199,30 +282,32 @@ pub(super) async fn build_admin_create_api_key_response(
created
};
Ok(attach_admin_audit_response(
Json(json!({
"id": created.api_key_id,
"key": plaintext_key,
"name": created.name,
"key_display": masked_user_api_key_display(state, created.key_encrypted.as_deref()),
"is_standalone": true,
"is_active": created.is_active,
"rate_limit": created.rate_limit,
"concurrent_limit": created.concurrent_limit,
"allowed_providers": created.allowed_providers,
"allowed_api_formats": created.allowed_api_formats,
"allowed_models": created.allowed_models,
"expires_at": format_optional_unix_secs_iso8601(created.expires_at_unix_secs),
"auto_delete_on_expiry": created.auto_delete_on_expiry,
"feature_settings": created.feature_settings,
"wallet": serialize_admin_system_users_export_wallet(Some(&wallet)),
"message": "独立余额Key创建成功,请妥善保存完整密钥,后续将无法查看",
}))
.into_response(),
"admin_standalone_api_key_created",
"create_standalone_api_key",
"api_key",
&created.api_key_id,
Ok(mark_sensitive_admin_response_no_store(
attach_admin_audit_response(
Json(json!({
"id": created.api_key_id,
"key": plaintext_key,
"name": created.name,
"key_display": masked_user_api_key_display(state, &created),
"is_standalone": true,
"is_active": created.is_active,
"rate_limit": created.rate_limit,
"concurrent_limit": created.concurrent_limit,
"allowed_providers": created.allowed_providers,
"allowed_api_formats": created.allowed_api_formats,
"allowed_models": created.allowed_models,
"expires_at": format_optional_unix_secs_iso8601(created.expires_at_unix_secs),
"auto_delete_on_expiry": created.auto_delete_on_expiry,
"feature_settings": created.feature_settings,
"wallet": serialize_admin_system_users_export_wallet(Some(&wallet)),
"message": "独立余额Key创建成功,请妥善保存完整密钥,后续将无法查看",
}))
.into_response(),
"admin_standalone_api_key_created",
"create_standalone_api_key",
"api_key",
&created.api_key_id,
),
))
}
@@ -405,7 +490,11 @@ pub(super) async fn build_admin_update_api_key_response(
.update_standalone_api_key_basic(
aether_data::repository::auth::UpdateStandaloneApiKeyBasicRecord {
api_key_id: api_key_id.clone(),
key_encrypted: None,
key_encrypted_present: false,
name,
name_present: field_presence.contains("name"),
force_capabilities: None,
rate_limit_present: field_presence.contains("rate_limit"),
rate_limit: payload.rate_limit,
concurrent_limit_present: field_presence.contains("concurrent_limit"),
@@ -540,3 +629,127 @@ pub(super) async fn build_admin_delete_api_key_response(
false => Ok(build_admin_api_keys_not_found_response()),
}
}
#[cfg(test)]
mod tests {
use super::compensate_standalone_api_key_creation;
use crate::data::GatewayDataState;
use crate::handlers::admin::request::AdminAppState;
use crate::state::AppState;
use aether_data::repository::auth::{
AuthApiKeyReadRepository, AuthApiKeyWriteRepository, CreateStandaloneApiKeyRecord,
InMemoryAuthApiKeySnapshotRepository,
};
use aether_data::repository::wallet::{StoredWalletSnapshot, WalletLookupKey};
use std::sync::Arc;
async fn seed_standalone_key(
repository: &InMemoryAuthApiKeySnapshotRepository,
api_key_id: &str,
) {
repository
.create_standalone_api_key(CreateStandaloneApiKeyRecord {
user_id: "admin-user".to_string(),
api_key_id: api_key_id.to_string(),
key_hash: format!("hash-{api_key_id}"),
key_encrypted: None,
name: Some("test-key".to_string()),
allowed_providers: None,
allowed_api_formats: None,
allowed_models: None,
ip_rules: None,
rate_limit: None,
concurrent_limit: None,
force_capabilities: None,
is_active: true,
expires_at_unix_secs: None,
auto_delete_on_expiry: false,
total_requests: 0,
total_tokens: 0,
total_cost_usd: 0.0,
})
.await
.expect("key creation should succeed")
.expect("key should be returned");
}
fn wallet_for_key(api_key_id: &str, balance: f64) -> StoredWalletSnapshot {
StoredWalletSnapshot::new(
format!("wallet-{api_key_id}"),
None,
Some(api_key_id.to_string()),
balance,
0.0,
"finite".to_string(),
"USD".to_string(),
"active".to_string(),
if balance > 0.0 { balance } else { 0.0 },
0.0,
0.0,
0.0,
1,
)
.expect("wallet should build")
}
#[tokio::test]
async fn standalone_key_compensation_removes_unreferenced_wallet_and_key() {
let api_key_id = "compensate-key";
let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default());
seed_standalone_key(&repository, api_key_id).await;
let wallet = wallet_for_key(api_key_id, 0.0);
let state = AppState::new()
.expect("state should build")
.with_auth_wallets_for_tests([wallet])
.with_data_state_for_tests(GatewayDataState::with_auth_api_key_repository_for_tests(
Arc::clone(&repository),
));
let admin_state = AdminAppState::new(&state);
compensate_standalone_api_key_creation(&admin_state, api_key_id)
.await
.expect("compensation should succeed");
assert!(state
.find_wallet(WalletLookupKey::ApiKeyId(api_key_id))
.await
.expect("wallet lookup should succeed")
.is_none());
assert!(repository
.find_export_standalone_api_key_by_id(api_key_id)
.await
.expect("key lookup should succeed")
.is_none());
}
#[tokio::test]
async fn standalone_key_compensation_preserves_funded_wallet_and_key() {
let api_key_id = "funded-compensate-key";
let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default());
seed_standalone_key(&repository, api_key_id).await;
let wallet = wallet_for_key(api_key_id, 5.0);
let state = AppState::new()
.expect("state should build")
.with_auth_wallets_for_tests([wallet])
.with_data_state_for_tests(GatewayDataState::with_auth_api_key_repository_for_tests(
Arc::clone(&repository),
));
let admin_state = AdminAppState::new(&state);
assert!(
compensate_standalone_api_key_creation(&admin_state, api_key_id)
.await
.is_err()
);
assert!(state
.find_wallet(WalletLookupKey::ApiKeyId(api_key_id))
.await
.expect("wallet lookup should succeed")
.is_some());
assert!(repository
.find_export_standalone_api_key_by_id(api_key_id)
.await
.expect("key lookup should succeed")
.is_some());
}
}
@@ -4,9 +4,12 @@ use super::shared::{
build_admin_api_keys_bad_request_response, build_admin_api_keys_data_unavailable_response,
build_admin_api_keys_not_found_response,
};
use super::{decrypt_catalog_secret_with_fallbacks, query_param_bool, query_param_optional_bool};
use super::{query_param_bool, query_param_optional_bool};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::handlers::admin::shared::{
attach_admin_audit_response, mark_sensitive_admin_response_no_store,
};
use crate::handlers::shared::decrypt_or_migrate_auth_api_key_secret;
use crate::GatewayError;
use axum::{
body::Body,
@@ -132,20 +135,24 @@ pub(super) async fn build_admin_api_key_detail_response(
"该密钥没有存储完整密钥信息",
));
};
let Some(key) = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)
else {
return Ok((
http::StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({ "detail": "解密密钥失败" })),
)
.into_response());
let key = match decrypt_or_migrate_auth_api_key_secret(state.app(), &record).await {
Ok(value) => value,
Err(_) => {
return Ok((
http::StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({ "detail": "解密或校验密钥失败" })),
)
.into_response())
}
};
return Ok(attach_admin_audit_response(
Json(json!({ "key": key })).into_response(),
"admin_standalone_api_key_revealed",
"reveal_standalone_api_key",
"api_key",
&api_key_id,
return Ok(mark_sensitive_admin_response_no_store(
attach_admin_audit_response(
Json(json!({ "key": key })).into_response(),
"admin_standalone_api_key_revealed",
"reveal_standalone_api_key",
"api_key",
&api_key_id,
),
));
}
@@ -13,6 +13,7 @@ pub(super) async fn maybe_build_local_admin_api_keys_routes_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_headers: &http::HeaderMap,
remote_addr: &std::net::SocketAddr,
request_body: Option<&axum::body::Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.decision() else {
@@ -68,6 +69,7 @@ pub(super) async fn maybe_build_local_admin_api_keys_routes_response(
state,
request_context,
request_headers,
remote_addr,
request_body,
)
.await?,
@@ -146,8 +146,11 @@ pub(super) fn admin_api_keys_parse_limit(query: Option<&str>) -> Result<usize, S
}
}
fn masked_admin_api_key_display(state: &AdminAppState<'_>, ciphertext: Option<&str>) -> String {
masked_user_api_key_display(state, ciphertext)
fn masked_admin_api_key_display(
state: &AdminAppState<'_>,
record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord,
) -> String {
masked_user_api_key_display(state, record)
}
pub(super) fn build_admin_api_key_list_item_payload(
@@ -159,7 +162,7 @@ pub(super) fn build_admin_api_key_list_item_payload(
"id": record.api_key_id,
"user_id": record.user_id,
"name": record.name,
"key_display": masked_admin_api_key_display(state, record.key_encrypted.as_deref()),
"key_display": masked_admin_api_key_display(state, record),
"is_active": record.is_active,
"is_standalone": true,
"total_requests": record.total_requests,
@@ -190,7 +193,7 @@ pub(super) fn build_admin_api_key_detail_payload(
"id": record.api_key_id,
"user_id": record.user_id,
"name": record.name,
"key_display": masked_admin_api_key_display(state, record.key_encrypted.as_deref()),
"key_display": masked_admin_api_key_display(state, record),
"is_active": record.is_active,
"is_standalone": true,
"total_requests": record.total_requests,
@@ -1,9 +1,10 @@
use super::shared::*;
use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError;
use aether_data::repository::auth_modules::{LdapBindPasswordUpdate, StoredLdapModuleConfig};
use serde::Deserialize;
#[derive(Debug, Deserialize)]
#[derive(Deserialize)]
pub(super) struct AdminLdapConfigUpdateRequest {
server_url: String,
bind_dn: String,
@@ -28,7 +29,7 @@ pub(super) struct AdminLdapConfigUpdateRequest {
connect_timeout: i32,
}
#[derive(Debug, Default, Deserialize)]
#[derive(Default, Deserialize)]
pub(super) struct AdminLdapConfigTestRequest {
#[serde(default)]
server_url: Option<String>,
@@ -56,7 +57,7 @@ pub(super) struct AdminLdapConfigTestRequest {
connect_timeout: Option<i32>,
}
#[derive(Debug, Clone)]
#[derive(Clone)]
pub(super) struct AdminLdapConnectionTestConfig {
server_url: String,
bind_dn: String,
@@ -66,13 +67,26 @@ pub(super) struct AdminLdapConnectionTestConfig {
connect_timeout: i32,
}
pub(super) struct AdminLdapConfigUpdate {
pub(super) expected: Option<StoredLdapModuleConfig>,
pub(super) replacement: StoredLdapModuleConfig,
pub(super) bind_password_update: LdapBindPasswordUpdate,
}
pub(super) async fn build_admin_ldap_update_config(
state: &AdminAppState<'_>,
payload: AdminLdapConfigUpdateRequest,
) -> Result<aether_data::repository::auth_modules::StoredLdapModuleConfig, String> {
) -> Result<AdminLdapConfigUpdate, String> {
let server_url = admin_ldap_trim_required(payload.server_url, "LDAP 服务器地址不能为空")?;
let server_url = admin_ldap_normalize_server_url(&server_url, payload.use_starttls)
.ok_or_else(|| {
"LDAP 服务器地址必须使用 ldaps://,或在启用 StartTLS 时使用 ldap://;不得包含凭据、查询参数或片段"
.to_string()
})?;
let bind_dn = admin_ldap_trim_required(payload.bind_dn, "绑定 DN 不能为空")?;
let base_dn = admin_ldap_trim_required(payload.base_dn, "Base DN 不能为空")?;
admin_ldap_validate_distinguished_name(&bind_dn, "绑定 DN")?;
admin_ldap_validate_distinguished_name(&base_dn, "Base DN")?;
let user_search_filter =
admin_ldap_trim_required(payload.user_search_filter, "搜索过滤器不能为空")?;
admin_ldap_validate_search_filter(&user_search_filter)?;
@@ -80,39 +94,42 @@ pub(super) async fn build_admin_ldap_update_config(
let email_attr = admin_ldap_trim_required(payload.email_attr, "邮箱属性不能为空")?;
let display_name_attr =
admin_ldap_trim_required(payload.display_name_attr, "显示名称属性不能为空")?;
admin_ldap_validate_attribute_description(&username_attr, "用户名属性")?;
admin_ldap_validate_attribute_description(&email_attr, "邮箱属性")?;
admin_ldap_validate_attribute_description(&display_name_attr, "显示名称属性")?;
if !(1..=60).contains(&payload.connect_timeout) {
return Err("连接超时时间必须在 1 到 60 秒之间".to_string());
}
let existing = state
let mut existing = state
.get_ldap_module_config()
.await
.map_err(|err| format!("{err:?}"))?;
let bind_password_update_requested = payload
.bind_password
.as_ref()
.is_some_and(|value| !value.is_empty());
let bind_password = match payload.bind_password {
Some(value) if value.is_empty() => Some(String::new()),
Some(value) => Some(admin_ldap_trim_required(value, "绑定密码不能为空")?),
None => None,
};
if payload.bind_password.is_none() {
if let Some(config) = existing.as_ref() {
crate::handlers::shared::decrypt_or_migrate_ldap_bind_password(state.app(), config)
.await
.map_err(|_| "已保存的 LDAP 绑定密码无法解密".to_string())?;
existing = state
.get_ldap_module_config()
.await
.map_err(|err| format!("{err:?}"))?;
}
}
let requested_bind_password = payload.bind_password;
let is_new_config = existing.is_none();
if is_new_config && bind_password.as_deref().unwrap_or("").is_empty() {
let will_have_password = match requested_bind_password.as_deref() {
Some(value) => !value.trim().is_empty(),
None => existing
.as_ref()
.and_then(|config| config.bind_password_encrypted.as_deref())
.map(str::trim)
.is_some_and(|value: &str| !value.is_empty()),
};
if is_new_config && !will_have_password {
return Err("首次配置 LDAP 时必须设置绑定密码".to_string());
}
let will_have_password = bind_password
.as_ref()
.map(|value| !value.is_empty())
.unwrap_or_else(|| {
existing
.as_ref()
.and_then(|config| config.bind_password_encrypted.as_deref())
.map(str::trim)
.is_some_and(|value: &str| !value.is_empty())
});
if payload.is_exclusive && !payload.is_enabled {
return Err("仅允许 LDAP 登录 需要先启用 LDAP 认证".to_string());
}
@@ -135,31 +152,58 @@ pub(super) async fn build_admin_ldap_update_config(
}
}
let bind_password_encrypted = match bind_password {
Some(value) if value.is_empty() => None,
Some(value) => state.encrypt_catalog_secret_with_fallbacks(&value),
None => existing.and_then(|config| config.bind_password_encrypted),
let replacement = StoredLdapModuleConfig {
server_url,
bind_dn,
// The repository ignores this field for config mutations. Password changes are
// carried exclusively by `bind_password_update`, so Preserve never copies an old
// ciphertext into the replacement record.
bind_password_encrypted: None,
base_dn,
user_search_filter: Some(user_search_filter),
username_attr: Some(username_attr),
email_attr: Some(email_attr),
display_name_attr: Some(display_name_attr),
is_enabled: payload.is_enabled,
is_exclusive: payload.is_exclusive,
use_starttls: payload.use_starttls,
connect_timeout: Some(payload.connect_timeout),
};
if bind_password_update_requested && bind_password_encrypted.is_none() {
return Err("LDAP 绑定密码加密失败,请检查 Rust 数据加密配置".to_string());
}
Ok(
aether_data::repository::auth_modules::StoredLdapModuleConfig {
server_url,
bind_dn,
bind_password_encrypted,
base_dn,
user_search_filter: Some(user_search_filter),
username_attr: Some(username_attr),
email_attr: Some(email_attr),
display_name_attr: Some(display_name_attr),
is_enabled: payload.is_enabled,
is_exclusive: payload.is_exclusive,
use_starttls: payload.use_starttls,
connect_timeout: Some(payload.connect_timeout),
},
)
let bind_password_update = match requested_bind_password {
Some(value) if value.is_empty() => LdapBindPasswordUpdate::Clear,
Some(value) => {
let password = admin_ldap_trim_required(value, "绑定密码不能为空")?;
let ciphertext = state
.encrypt_ldap_bind_password(&replacement, &password)
.ok_or_else(|| "LDAP 绑定密码加密失败,请检查 Rust 数据加密配置".to_string())?;
LdapBindPasswordUpdate::Set(ciphertext)
}
None => {
if let Some(existing_config) = existing.as_ref() {
if existing_config
.bind_password_encrypted
.as_deref()
.is_some_and(|value| !value.trim().is_empty())
&& !crate::handlers::shared::ldap_bind_password_binding_matches(
existing_config,
&replacement,
)
.unwrap_or(false)
{
return Err(
"修改 LDAP 服务器、StartTLS、bind DN 或 Base DN 时必须重新提供绑定密码"
.to_string(),
);
}
}
LdapBindPasswordUpdate::Preserve
}
};
Ok(AdminLdapConfigUpdate {
expected: existing,
replacement,
bind_password_update,
})
}
pub(super) async fn build_admin_ldap_test_config(
@@ -167,7 +211,16 @@ pub(super) async fn build_admin_ldap_test_config(
payload: AdminLdapConfigTestRequest,
) -> Result<Option<AdminLdapConnectionTestConfig>, String> {
if let Some(value) = payload.user_search_filter.as_deref() {
admin_ldap_validate_search_filter(value.trim())?;
admin_ldap_validate_search_filter(value)?;
}
for (value, label) in [
(payload.username_attr.as_deref(), "用户名属性"),
(payload.email_attr.as_deref(), "邮箱属性"),
(payload.display_name_attr.as_deref(), "显示名称属性"),
] {
if let Some(value) = value {
admin_ldap_validate_attribute_description(value.trim(), label)?;
}
}
if let Some(connect_timeout) = payload.connect_timeout {
if !(1..=60).contains(&connect_timeout) {
@@ -199,9 +252,14 @@ pub(super) async fn build_admin_ldap_test_config(
.as_ref()
.and_then(|config| config.connect_timeout)
.unwrap_or_else(admin_ldap_default_connect_timeout);
let mut bind_password = saved
.as_ref()
.and_then(|config| admin_ldap_read_saved_bind_password(state, config));
let mut bind_password = match saved.as_ref() {
Some(config) => {
crate::handlers::shared::decrypt_or_migrate_ldap_bind_password(state.app(), config)
.await
.map_err(|_| "已保存的 LDAP 绑定密码无法解密".to_string())?
}
None => None,
};
if let Some(value) = payload.server_url {
server_url = Some(admin_ldap_trim_required(value, "LDAP 服务器地址不能为空")?);
@@ -239,11 +297,21 @@ pub(super) async fn build_admin_ldap_test_config(
return Ok(None);
}
let server_url = server_url.expect("server_url already checked");
let server_url = admin_ldap_normalize_server_url(&server_url, use_starttls).ok_or_else(|| {
"LDAP 服务器地址必须使用 ldaps://,或在启用 StartTLS 时使用 ldap://;不得包含凭据、查询参数或片段"
.to_string()
})?;
let bind_dn = bind_dn.expect("bind_dn already checked");
let base_dn = base_dn.expect("base_dn already checked");
admin_ldap_validate_distinguished_name(&bind_dn, "绑定 DN")?;
admin_ldap_validate_distinguished_name(&base_dn, "Base DN")?;
Ok(Some(AdminLdapConnectionTestConfig {
server_url: server_url.expect("server_url already checked"),
bind_dn: bind_dn.expect("bind_dn already checked"),
server_url,
bind_dn,
bind_password: bind_password.expect("bind_password already checked"),
base_dn: base_dn.expect("base_dn already checked"),
base_dn,
use_starttls,
connect_timeout,
}))
@@ -270,7 +338,8 @@ pub(super) async fn admin_ldap_test_connection(
}
fn admin_ldap_test_connection_blocking(config: AdminLdapConnectionTestConfig) -> (bool, String) {
let Some(server_url): Option<String> = admin_ldap_normalize_server_url(&config.server_url)
let Some(server_url): Option<String> =
admin_ldap_normalize_server_url(&config.server_url, config.use_starttls)
else {
return (false, ADMIN_LDAP_TEST_FAILURE_MESSAGE.to_string());
};
@@ -302,51 +371,21 @@ fn admin_ldap_trim_required(value: String, detail: &str) -> Result<String, Strin
}
fn admin_ldap_validate_search_filter(value: &str) -> Result<(), String> {
if value.is_empty() {
return Err("搜索过滤器不能为空".to_string());
}
if !value.contains("{username}") {
return Err("搜索过滤器必须包含 {username} 占位符".to_string());
}
let mut depth = 0i32;
let mut max_depth = 0i32;
for ch in value.chars() {
if ch == '(' {
depth += 1;
max_depth = max_depth.max(depth);
} else if ch == ')' {
depth -= 1;
if depth < 0 {
return Err("搜索过滤器括号不匹配".to_string());
}
}
}
if depth != 0 {
return Err("搜索过滤器括号不匹配".to_string());
}
if max_depth > 5 {
return Err("搜索过滤器嵌套层数过深(最多5层)".to_string());
}
if value.len() > 200 {
return Err("搜索过滤器过长(最多200字符)".to_string());
}
Ok(())
}
fn admin_ldap_read_saved_bind_password(
state: &AdminAppState<'_>,
config: &aether_data::repository::auth_modules::StoredLdapModuleConfig,
) -> Option<String> {
config
.bind_password_encrypted
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(|value| {
state
.decrypt_catalog_secret_with_fallbacks(value)
.or_else(|| Some(value.to_string()))
crate::handlers::shared::ldap_search_filter_is_valid(value)
.then_some(())
.ok_or_else(|| {
"搜索过滤器格式无效,必须包含 {username} 且使用唯一、有限的外层括号结构".to_string()
})
.filter(|value| !value.trim().is_empty())
}
fn admin_ldap_validate_distinguished_name(value: &str, label: &str) -> Result<(), String> {
crate::handlers::shared::ldap_distinguished_name_is_valid(value)
.then_some(())
.ok_or_else(|| format!("{label}格式无效或过长"))
}
fn admin_ldap_validate_attribute_description(value: &str, label: &str) -> Result<(), String> {
crate::handlers::shared::ldap_attribute_description_is_valid(value)
.then_some(())
.ok_or_else(|| format!("{label}必须是有效的 LDAP 属性名称"))
}
@@ -6,6 +6,7 @@ use super::shared::*;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::GatewayError;
use aether_data::repository::auth_modules::CompareAndSwapLdapConfigResult;
use axum::{
body::{Body, Bytes},
http,
@@ -61,9 +62,20 @@ pub(super) async fn maybe_build_local_admin_ldap_response(
Ok(config) => config,
Err(detail) => return Ok(Some(admin_ldap_bad_request_response(detail))),
};
let saved = state.upsert_ldap_module_config(&update).await?;
if saved.is_none() {
let saved = state
.compare_and_swap_ldap_module_config(
update.expected.as_ref(),
&update.replacement,
&update.bind_password_update,
)
.await?;
let Some(saved) = saved else {
return Ok(Some(admin_ldap_unavailable_response()));
};
if saved == CompareAndSwapLdapConfigResult::Conflict {
return Ok(Some(admin_ldap_conflict_response(
"LDAP 配置已被其他请求更新,请重新加载后重试",
)));
}
return Ok(Some(
Json(json!({ "message": "LDAP配置更新成功" })).into_response(),
@@ -102,6 +102,14 @@ pub(super) fn admin_ldap_bad_request_response(detail: impl Into<String>) -> Resp
.into_response()
}
pub(super) fn admin_ldap_conflict_response(detail: impl Into<String>) -> Response<Body> {
(
http::StatusCode::CONFLICT,
Json(json!({ "detail": detail.into() })),
)
.into_response()
}
pub(super) fn admin_ldap_unavailable_response() -> Response<Body> {
(
http::StatusCode::SERVICE_UNAVAILABLE,
@@ -110,13 +118,9 @@ pub(super) fn admin_ldap_unavailable_response() -> Response<Body> {
.into_response()
}
pub(super) fn admin_ldap_normalize_server_url(server_url: &str) -> Option<String> {
let server_url = server_url.trim();
if server_url.is_empty() {
return None;
}
if server_url.contains("://") {
return Some(server_url.to_string());
}
Some(format!("ldap://{server_url}"))
pub(super) fn admin_ldap_normalize_server_url(
server_url: &str,
use_starttls: bool,
) -> Option<String> {
crate::handlers::shared::normalize_ldap_transport_server_url(server_url, use_starttls)
}
@@ -1,13 +1,14 @@
use crate::handlers::admin::request::AdminAppState;
use aether_data::repository::oauth_providers::{
EncryptedSecretUpdate, UpsertOAuthProviderConfigRecord,
validate_oauth_frontend_callback_url, validate_oauth_provider_endpoint_config,
validate_oauth_redirect_uri, EncryptedSecretUpdate, UpsertOAuthProviderConfigRecord,
};
use axum::http;
use serde::Deserialize;
use serde_json::json;
use url::Url;
use url::{Host, Url};
#[derive(Debug, Deserialize)]
#[derive(Deserialize)]
pub(crate) struct AdminOAuthProviderUpsertRequest {
pub(super) display_name: String,
pub(super) client_id: String,
@@ -122,7 +123,9 @@ pub(super) fn admin_oauth_is_supported_provider(provider_type: &str) -> bool {
})
}
fn admin_oauth_builtin_allowed_domains(provider_type: &str) -> Option<&'static [&'static str]> {
pub(super) fn admin_oauth_builtin_allowed_domains(
provider_type: &str,
) -> Option<&'static [&'static str]> {
if provider_type.eq_ignore_ascii_case("linuxdo") {
Some(&["linux.do", "connect.linux.do", "connect.linuxdo.org"])
} else {
@@ -130,7 +133,9 @@ fn admin_oauth_builtin_allowed_domains(provider_type: &str) -> Option<&'static [
}
}
fn admin_oauth_custom_allowed_domains(extra_config: Option<&serde_json::Value>) -> Vec<String> {
pub(super) fn admin_oauth_custom_allowed_domains(
extra_config: Option<&serde_json::Value>,
) -> Vec<String> {
extra_config
.and_then(serde_json::Value::as_object)
.and_then(|object| {
@@ -146,42 +151,57 @@ fn admin_oauth_custom_allowed_domains(extra_config: Option<&serde_json::Value>)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.trim_end_matches('.').to_ascii_lowercase())
.filter(|value| value.parse::<std::net::IpAddr>().is_err())
.collect::<Vec<_>>()
})
.unwrap_or_default()
}
fn validate_admin_oauth_frontend_callback_url(url: &str) -> Result<(), String> {
let parsed = Url::parse(url).map_err(|_| "frontend_callback_url 必须是绝对 URL".to_string())?;
if !matches!(parsed.scheme(), "http" | "https") {
return Err("frontend_callback_url scheme 必须是 http/https".to_string());
}
if parsed.host_str().is_none() {
return Err("frontend_callback_url 必须是绝对 URL".to_string());
}
let path = parsed.path().trim_end_matches('/');
if !path.ends_with("/auth/callback") {
return Err("frontend_callback_url 路径必须以 /auth/callback 结尾".to_string());
}
Ok(())
pub(super) fn validate_admin_oauth_url_override(
url: &str,
allowed_domains: &[&str],
) -> Result<(), String> {
validate_admin_oauth_url_override_with_options(url, allowed_domains, false)
}
fn validate_admin_oauth_redirect_uri(url: &str) -> Result<(), String> {
let parsed = Url::parse(url).map_err(|_| "redirect_uri 必须是绝对 URL".to_string())?;
if !matches!(parsed.scheme(), "http" | "https") {
return Err("redirect_uri scheme 必须是 http/https".to_string());
}
if parsed.host_str().is_none() {
return Err("redirect_uri 必须是绝对 URL".to_string());
}
Ok(())
fn validate_admin_oauth_authorization_url_override(
url: &str,
allowed_domains: &[&str],
) -> Result<(), String> {
validate_admin_oauth_url_override_with_options(url, allowed_domains, true)
}
fn validate_admin_oauth_url_override(url: &str, allowed_domains: &[&str]) -> Result<(), String> {
fn validate_admin_oauth_url_override_with_options(
url: &str,
allowed_domains: &[&str],
reject_authorization_parameters: bool,
) -> Result<(), String> {
let parsed = Url::parse(url).map_err(|_| "端点覆盖必须是 https 绝对 URL".to_string())?;
if parsed.scheme() != "https" || parsed.host_str().is_none() {
return Err("端点覆盖必须是 https 绝对 URL".to_string());
}
if matches!(parsed.host(), Some(Host::Ipv4(_)) | Some(Host::Ipv6(_))) {
return Err("端点覆盖必须使用 DNS 主机名,不能使用 IP 字面量".to_string());
}
if !parsed.username().is_empty() || parsed.password().is_some() || parsed.fragment().is_some() {
return Err("端点覆盖不得包含 URL 凭据或 fragment".to_string());
}
if reject_authorization_parameters
&& parsed.query_pairs().any(|(name, _)| {
matches!(
name.to_ascii_lowercase().as_str(),
"response_type"
| "client_id"
| "redirect_uri"
| "state"
| "scope"
| "code_challenge"
| "code_challenge_method"
)
})
{
return Err("authorization endpoint 不得预置 OAuth authorization 参数".to_string());
}
let host = parsed
.host_str()
.map(|value| value.trim().trim_end_matches('.').to_ascii_lowercase())
@@ -207,6 +227,17 @@ fn validate_admin_oauth_url_override_for_domains(
validate_admin_oauth_url_override(url, &allowed)
}
fn validate_admin_oauth_authorization_url_override_for_domains(
url: &str,
allowed_domains: &[String],
) -> Result<(), String> {
let allowed = allowed_domains
.iter()
.map(String::as_str)
.collect::<Vec<_>>();
validate_admin_oauth_authorization_url_override(url, &allowed)
}
pub(super) fn build_admin_oauth_upsert_record(
state: &AdminAppState<'_>,
provider_type: &str,
@@ -236,8 +267,8 @@ pub(super) fn build_admin_oauth_upsert_record(
return Err("frontend_callback_url 不能为空".to_string());
}
validate_admin_oauth_frontend_callback_url(frontend_callback_url)?;
validate_admin_oauth_redirect_uri(redirect_uri)?;
validate_oauth_frontend_callback_url(frontend_callback_url)?;
validate_oauth_redirect_uri(redirect_uri)?;
let is_custom_oidc = admin_oauth_is_custom_provider_type(&provider_type);
let custom_allowed_domains = if is_custom_oidc {
@@ -268,16 +299,26 @@ pub(super) fn build_admin_oauth_upsert_record(
let Some(value) = value.map(str::trim).filter(|value| !value.is_empty()) else {
return Err(format!("custom_oidc 必须配置 {field_name}"));
};
validate_admin_oauth_url_override_for_domains(value, &custom_allowed_domains)?;
if field_name == "authorization_url_override" {
validate_admin_oauth_authorization_url_override_for_domains(
value,
&custom_allowed_domains,
)?;
} else {
validate_admin_oauth_url_override_for_domains(value, &custom_allowed_domains)?;
}
}
}
if let Some(value) = payload.authorization_url_override.as_deref().map(str::trim) {
if !value.is_empty() {
if let Some(allowed_domains) = builtin_allowed_domains {
validate_admin_oauth_url_override(value, allowed_domains)?;
validate_admin_oauth_authorization_url_override(value, allowed_domains)?;
} else {
validate_admin_oauth_url_override_for_domains(value, &custom_allowed_domains)?;
validate_admin_oauth_authorization_url_override_for_domains(
value,
&custom_allowed_domains,
)?;
}
}
}
@@ -322,40 +363,34 @@ pub(super) fn build_admin_oauth_upsert_record(
return Err("scopes 不能为空".to_string());
}
let client_secret_encrypted = match payload.client_secret.as_deref() {
None => EncryptedSecretUpdate::Preserve,
Some(raw) => {
let secret = raw.trim();
if secret == "__CLEAR__" {
EncryptedSecretUpdate::Clear
} else if secret.is_empty() {
EncryptedSecretUpdate::Preserve
} else {
let encrypted = state
.encrypt_catalog_secret_with_fallbacks(secret)
.ok_or_else(|| "gateway 未配置 OAuth provider 加密密钥".to_string())?;
EncryptedSecretUpdate::Set(encrypted)
}
}
};
validate_oauth_provider_endpoint_config(
&provider_type,
payload.authorization_url_override.as_deref(),
payload.token_url_override.as_deref(),
payload.userinfo_url_override.as_deref(),
payload.extra_config.as_ref(),
)?;
Ok(UpsertOAuthProviderConfigRecord {
let authorization_url_override = payload.authorization_url_override.and_then(|value| {
let value = value.trim().to_string();
(!value.is_empty()).then_some(value)
});
let token_url_override = payload.token_url_override.and_then(|value| {
let value = value.trim().to_string();
(!value.is_empty()).then_some(value)
});
let userinfo_url_override = payload.userinfo_url_override.and_then(|value| {
let value = value.trim().to_string();
(!value.is_empty()).then_some(value)
});
let mut record = UpsertOAuthProviderConfigRecord {
provider_type,
display_name: display_name.to_string(),
client_id: client_id.to_string(),
client_secret_encrypted,
authorization_url_override: payload.authorization_url_override.and_then(|value| {
let value = value.trim().to_string();
(!value.is_empty()).then_some(value)
}),
token_url_override: payload.token_url_override.and_then(|value| {
let value = value.trim().to_string();
(!value.is_empty()).then_some(value)
}),
userinfo_url_override: payload.userinfo_url_override.and_then(|value| {
let value = value.trim().to_string();
(!value.is_empty()).then_some(value)
}),
client_secret_encrypted: EncryptedSecretUpdate::Preserve,
authorization_url_override,
token_url_override,
userinfo_url_override,
scopes: payload.scopes.map(|items| {
items
.into_iter()
@@ -372,5 +407,27 @@ pub(super) fn build_admin_oauth_upsert_record(
(!value.is_empty()).then_some(value)
}),
is_enabled: payload.is_enabled,
})
};
record.client_secret_encrypted = match payload.client_secret.as_deref() {
None => EncryptedSecretUpdate::Preserve,
Some(raw) => {
let secret = raw.trim();
if secret == "__CLEAR__" {
EncryptedSecretUpdate::Clear
} else if secret.is_empty() {
EncryptedSecretUpdate::Preserve
} else {
let encrypted =
crate::handlers::shared::seal_identity_oauth_provider_client_secret(
state.as_ref(),
&record,
secret,
)
.map_err(str::to_string)?;
EncryptedSecretUpdate::Set(encrypted)
}
}
};
Ok(record)
}
@@ -1,8 +1,9 @@
use super::oauth_config::{
admin_oauth_builtin_allowed_domains, admin_oauth_custom_allowed_domains,
admin_oauth_is_supported_provider, admin_oauth_provider_type_from_path,
admin_oauth_test_provider_type_from_path, build_admin_oauth_provider_payload,
build_admin_oauth_supported_types_payload, build_admin_oauth_upsert_record,
AdminOAuthProviderUpsertRequest,
validate_admin_oauth_url_override, AdminOAuthProviderUpsertRequest,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{attach_admin_audit_response, build_proxy_error_response};
@@ -14,9 +15,11 @@ use axum::{
Json,
};
use serde_json::json;
use std::net::{IpAddr, SocketAddr};
use std::time::Duration;
const ADMIN_OAUTH_TEST_TIMEOUT_SECS: u64 = 10;
const ADMIN_OAUTH_TEST_MAX_REDIRECTS: usize = 3;
const LINUXDO_AUTHORIZATION_URL: &str = "https://connect.linux.do/oauth2/authorize";
const LINUXDO_TOKEN_URL: &str = "https://connect.linux.do/oauth2/token";
@@ -36,27 +39,204 @@ fn admin_oauth_secret_status(has_secret: bool) -> &'static str {
}
}
async fn admin_oauth_endpoint_reachable(client: &reqwest::Client, url: &str) -> bool {
let Ok(parsed) = reqwest::Url::parse(url) else {
async fn admin_oauth_endpoint_reachable(
url: &str,
allowed_domains: &[&str],
allow_benchmarking_ip: bool,
) -> bool {
let Ok(mut current) = reqwest::Url::parse(url) else {
return false;
};
if !matches!(parsed.scheme(), "http" | "https") || parsed.host_str().is_none() {
for redirects in 0..=ADMIN_OAUTH_TEST_MAX_REDIRECTS {
if validate_admin_oauth_url_override(current.as_str(), allowed_domains).is_err() {
return false;
}
let Ok((host, addrs)) =
resolve_public_admin_oauth_endpoint_with_policy(&current, allow_benchmarking_ip).await
else {
return false;
};
let mut builder = reqwest::Client::builder()
.timeout(Duration::from_secs(ADMIN_OAUTH_TEST_TIMEOUT_SECS))
.redirect(reqwest::redirect::Policy::none())
.no_proxy();
if host.parse::<IpAddr>().is_err() {
builder = builder.resolve_to_addrs(host.as_str(), &addrs);
}
let Ok(client) = builder.build() else {
return false;
};
let Ok(response) = client
.get(current.clone())
.header(reqwest::header::ACCEPT, "*/*")
.header(
reqwest::header::USER_AGENT,
"Aether OAuth configuration tester",
)
.send()
.await
else {
return false;
};
if !response.status().is_redirection() {
return response.status().as_u16() < 500;
}
if redirects == ADMIN_OAUTH_TEST_MAX_REDIRECTS {
return false;
}
let Some(location) = response
.headers()
.get(reqwest::header::LOCATION)
.and_then(|value| value.to_str().ok())
else {
return false;
};
let Ok(next) = current.join(location) else {
return false;
};
current = next;
}
false
}
async fn resolve_public_admin_oauth_endpoint(
url: &reqwest::Url,
) -> Result<(String, Vec<SocketAddr>), ()> {
resolve_public_admin_oauth_endpoint_with_policy(url, false).await
}
async fn resolve_public_admin_oauth_endpoint_with_policy(
url: &reqwest::Url,
allow_benchmarking_ip: bool,
) -> Result<(String, Vec<SocketAddr>), ()> {
if url.scheme() != "https"
|| !url.username().is_empty()
|| url.password().is_some()
|| url.host_str().is_none()
{
return Err(());
}
let host = url.host_str().ok_or(())?;
let port = url.port_or_known_default().ok_or(())?;
let addrs = if let Ok(ip) = host.parse::<IpAddr>() {
vec![SocketAddr::new(ip, port)]
} else {
aether_http::lookup_host_with_limits(host, port, aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT)
.await
.map_err(|_| ())?
};
if validate_public_admin_oauth_resolved_addrs(url, &addrs, allow_benchmarking_ip).is_err() {
return Err(());
}
Ok((host.to_string(), addrs))
}
fn validate_public_admin_oauth_resolved_addrs(
url: &reqwest::Url,
addrs: &[SocketAddr],
allow_benchmarking_ip: bool,
) -> Result<(), ()> {
if addrs.is_empty()
|| addrs.iter().any(|addr| {
aether_http::is_private_or_reserved_ip(addr.ip())
&& !(allow_benchmarking_ip
&& is_fixed_linuxdo_oauth_origin(url)
&& aether_http::is_ipv4_benchmarking_fake_ip(addr.ip()))
})
{
return Err(());
}
Ok(())
}
fn is_fixed_linuxdo_oauth_origin(url: &reqwest::Url) -> bool {
url.scheme() == "https"
&& url.host_str().is_some_and(|host| {
host.trim_end_matches('.')
.eq_ignore_ascii_case("connect.linux.do")
})
&& url.port_or_known_default() == Some(443)
&& url.username().is_empty()
&& url.password().is_none()
&& url.query().is_none()
&& url.fragment().is_none()
}
fn admin_oauth_test_allowed_domains(
provider_type: &str,
payload: &serde_json::Value,
persisted_config: Option<&aether_data::repository::oauth_providers::StoredOAuthProviderConfig>,
) -> Vec<String> {
if let Some(domains) = admin_oauth_builtin_allowed_domains(provider_type) {
return domains.iter().map(|domain| (*domain).to_string()).collect();
}
let payload_extra = payload.get("extra_config");
let domains = admin_oauth_custom_allowed_domains(payload_extra);
if domains.is_empty() {
admin_oauth_custom_allowed_domains(
persisted_config.and_then(|provider| provider.extra_config.as_ref()),
)
} else {
domains
}
}
fn management_token_may_configure_frontend_callback(
request_context: &AdminRequestContext<'_>,
existing: Option<&aether_data::repository::oauth_providers::StoredOAuthProviderConfig>,
requested_callback: &str,
) -> bool {
let Some(principal) = request_context
.decision()
.and_then(|decision| decision.admin_principal.as_ref())
else {
return false;
};
if principal.management_token_id.is_none() {
return true;
}
let callback_changed =
existing.is_none_or(|provider| provider.frontend_callback_url != requested_callback.trim());
if !callback_changed {
return true;
}
match client
.get(parsed)
.header(reqwest::header::ACCEPT, "*/*")
.header(
reqwest::header::USER_AGENT,
"Aether OAuth configuration tester",
// A missing permission list is the legacy full-access token representation.
principal
.management_token_permissions
.as_ref()
.is_none_or(|permissions| {
permissions
.iter()
.any(|permission| permission == "admin:oauth:admin")
})
}
fn oauth_frontend_callback_permission_denied_response(
request_context: &AdminRequestContext<'_>,
) -> Response<Body> {
let actor_id = request_context
.decision()
.and_then(|decision| decision.admin_principal.as_ref())
.and_then(|principal| principal.management_token_id.as_deref())
.unwrap_or("unknown");
attach_admin_audit_response(
(
http::StatusCode::FORBIDDEN,
Json(json!({
"detail": "management token permission denied",
"required_permission": "admin:oauth:admin",
"route_family": "oauth_manage",
"route_kind": "upsert_provider",
"request_path": request_context.path(),
})),
)
.send()
.await
{
Ok(response) => response.status().as_u16() < 500,
Err(_) => false,
}
.into_response(),
"admin_oauth_frontend_callback_permission_denied",
"permission_denied",
"oauth_frontend_callback",
actor_id,
)
}
async fn build_admin_oauth_test_payload(
@@ -117,34 +297,38 @@ async fn build_admin_oauth_test_payload(
}));
};
let proxy_snapshot = state.app().resolve_system_proxy_snapshot().await;
let mut client_builder = reqwest::Client::builder()
.timeout(Duration::from_secs(ADMIN_OAUTH_TEST_TIMEOUT_SECS))
.redirect(reqwest::redirect::Policy::limited(3));
if let Some(proxy_url) = proxy_snapshot.as_ref().and_then(|p| p.url.as_deref()) {
if let Ok(proxy) = reqwest::Proxy::all(proxy_url) {
client_builder = client_builder.proxy(proxy);
}
}
let client = client_builder.build();
let Ok(client) = client else {
let allowed_domains =
admin_oauth_test_allowed_domains(provider_type, payload, persisted_config.as_ref());
let allowed_domain_refs = allowed_domains
.iter()
.map(String::as_str)
.collect::<Vec<_>>();
if allowed_domain_refs.is_empty()
|| validate_admin_oauth_url_override(&authorization_url, &allowed_domain_refs).is_err()
|| validate_admin_oauth_url_override(&token_url, &allowed_domain_refs).is_err()
{
return Ok(json!({
"authorization_url_reachable": false,
"token_url_reachable": false,
"secret_status": admin_oauth_secret_status(has_secret),
"details": "OAuth 配置测试 HTTP client 初始化失败",
"details": "OAuth 端点必须使用 https 且位于 provider 域名白名单中",
}));
};
}
let allow_benchmarking_ip = provider_type.eq_ignore_ascii_case("linuxdo");
let (authorization_url_reachable, token_url_reachable) = tokio::join!(
admin_oauth_endpoint_reachable(&client, &authorization_url),
admin_oauth_endpoint_reachable(&client, &token_url),
admin_oauth_endpoint_reachable(
&authorization_url,
&allowed_domain_refs,
allow_benchmarking_ip,
),
admin_oauth_endpoint_reachable(&token_url, &allowed_domain_refs, allow_benchmarking_ip),
);
let details = if authorization_url_reachable && token_url_reachable {
"OAuth 端点可达;client_secret 仅在授权回调兑换 code 时校验"
} else {
"OAuth 端点不可达或返回不可用状态;请检查端点 URL、网络和代理配置"
"OAuth 端点不可达或返回不可用状态;请检查端点 URL 和网络配置"
};
Ok(json!({
@@ -155,6 +339,57 @@ async fn build_admin_oauth_test_payload(
}))
}
#[cfg(test)]
mod tests {
use super::{
is_fixed_linuxdo_oauth_origin, resolve_public_admin_oauth_endpoint,
validate_public_admin_oauth_resolved_addrs,
};
use std::net::SocketAddr;
#[tokio::test]
async fn oauth_test_endpoint_rejects_loopback_https_targets_before_connecting() {
let url = reqwest::Url::parse("https://127.0.0.1/oauth/token").expect("URL");
assert!(resolve_public_admin_oauth_endpoint(&url).await.is_err());
}
#[test]
fn linuxdo_builtin_origin_allows_only_benchmarking_addresses() {
let fixed = reqwest::Url::parse("https://connect.linux.do/oauth2/token")
.expect("LinuxDo URL should parse");
let fake = SocketAddr::from(([198, 18, 75, 234], 443));
assert!(is_fixed_linuxdo_oauth_origin(&fixed));
assert!(validate_public_admin_oauth_resolved_addrs(&fixed, &[fake], true).is_ok());
assert!(validate_public_admin_oauth_resolved_addrs(&fixed, &[fake], false).is_err());
assert!(validate_public_admin_oauth_resolved_addrs(
&fixed,
&[fake, SocketAddr::from(([127, 0, 0, 1], 443))],
true,
)
.is_err());
}
#[test]
fn custom_or_non_default_oauth_origins_reject_benchmarking_addresses() {
let fake = SocketAddr::from(([198, 18, 75, 234], 443));
for raw_url in [
"https://oauth.example.test/token",
"https://connect.linux.do:8443/oauth2/token",
"https://connect.linuxdo.org/oauth2/token",
"https://connect.linux.do.evil.test/oauth2/token",
"https://connect.linux.do/oauth2/token?tenant=unexpected",
] {
let url = reqwest::Url::parse(raw_url).expect("test URL should parse");
assert!(
!is_fixed_linuxdo_oauth_origin(&url),
"must not trust {raw_url}"
);
assert!(validate_public_admin_oauth_resolved_addrs(&url, &[fake], true).is_err());
}
}
}
pub(crate) async fn maybe_build_local_admin_oauth_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -263,34 +498,16 @@ pub(crate) async fn maybe_build_local_admin_oauth_response(
}
};
let existing = state.get_oauth_provider_config(&provider_type).await?;
let ldap_exclusive = state.get_ldap_module_config().await?.is_some_and(|config| {
config.is_enabled
&& config.is_exclusive
&& config
.bind_password_encrypted
.as_deref()
.map(str::trim)
.is_some_and(|value| !value.is_empty())
});
if existing
.as_ref()
.is_some_and(|provider| provider.is_enabled && !payload.is_enabled)
{
let affected_count = state
.count_locked_users_if_oauth_provider_disabled(&provider_type, ldap_exclusive)
.await?;
if affected_count > 0 && !payload.force {
return Ok(Some(build_proxy_error_response(
http::StatusCode::CONFLICT,
"confirmation_required",
format!("禁用该 Provider 会导致 {affected_count} 个用户无法登录"),
Some(json!({
"affected_count": affected_count,
"action": "disable_oauth_provider",
})),
)));
}
if !management_token_may_configure_frontend_callback(
request_context,
existing.as_ref(),
&payload.frontend_callback_url,
) {
return Ok(Some(oauth_frontend_callback_permission_denied_response(
request_context,
)));
}
let force_disable = payload.force;
let record = match build_admin_oauth_upsert_record(state, &provider_type, payload) {
Ok(record) => record,
Err(message) => {
@@ -302,9 +519,63 @@ pub(crate) async fn maybe_build_local_admin_oauth_response(
)));
}
};
let Some(provider) = state.upsert_oauth_provider_config(&record).await? else {
if let Some(existing) = existing.as_ref() {
if existing
.client_secret_encrypted
.as_deref()
.is_some_and(|value| !value.trim().is_empty())
&& matches!(
record.client_secret_encrypted,
aether_data::repository::oauth_providers::EncryptedSecretUpdate::Preserve
)
{
match crate::handlers::shared::identity_oauth_provider_secret_binding_matches(
existing, &record,
) {
Ok(true) => {}
Ok(false) => {
return Ok(Some(build_proxy_error_response(
http::StatusCode::BAD_REQUEST,
"invalid_request",
"修改 OAuth Provider 的 Client ID、端点或 redirect_uri 时必须重新提供 client_secret",
None,
)));
}
Err(_) => {
return Ok(Some(build_proxy_error_response(
http::StatusCode::BAD_REQUEST,
"invalid_request",
"OAuth Provider 密钥绑定校验失败,请重新提供 client_secret",
None,
)));
}
}
}
}
let Some(outcome) = state
.upsert_oauth_provider_config_with_force_disable(&record, force_disable)
.await?
else {
return Ok(None);
};
let provider = match outcome {
aether_data::repository::oauth_providers::UpsertOAuthProviderConfigOutcome::Upserted(
provider,
) => provider,
aether_data::repository::oauth_providers::UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation {
affected_count,
} => {
return Ok(Some(build_proxy_error_response(
http::StatusCode::CONFLICT,
"confirmation_required",
format!("禁用该 Provider 会导致 {affected_count} 个用户无法登录"),
Some(json!({
"affected_count": affected_count,
"action": "disable_oauth_provider",
})),
)));
}
};
return Ok(Some(
Json(build_admin_oauth_provider_payload(&provider)).into_response(),
));
@@ -322,7 +593,7 @@ pub(crate) async fn maybe_build_local_admin_oauth_response(
None,
)));
};
let Some(existing) = state.get_oauth_provider_config(&provider_type).await? else {
let Some(_existing) = state.get_oauth_provider_config(&provider_type).await? else {
return Ok(Some(build_proxy_error_response(
http::StatusCode::BAD_REQUEST,
"invalid_request",
@@ -330,32 +601,27 @@ pub(crate) async fn maybe_build_local_admin_oauth_response(
None,
)));
};
if existing.is_enabled {
let ldap_exclusive = state.get_ldap_module_config().await?.is_some_and(|config| {
config.is_enabled
&& config.is_exclusive
&& config
.bind_password_encrypted
.as_deref()
.map(str::trim)
.is_some_and(|value| !value.is_empty())
});
let affected_count = state
.count_locked_users_if_oauth_provider_disabled(&provider_type, ldap_exclusive)
.await?;
if affected_count > 0 {
let _mutation_guard = crate::oauth::lock_identity_oauth_mutation().await;
if state.has_oauth_links_for_provider(&provider_type).await? {
return Ok(Some(build_proxy_error_response(
http::StatusCode::CONFLICT,
"provider_has_bindings",
"Provider 仍有用户绑定,必须先解除全部绑定",
None,
)));
}
let deleted = state
.delete_oauth_provider_config_if_unlinked(&provider_type)
.await?;
if !deleted {
if state.has_oauth_links_for_provider(&provider_type).await? {
return Ok(Some(build_proxy_error_response(
http::StatusCode::BAD_REQUEST,
"invalid_request",
format!(
"删除该 Provider 会导致部分用户无法登录(数量: {affected_count}),已阻止操作"
),
http::StatusCode::CONFLICT,
"provider_has_bindings",
"Provider 仍有用户绑定,必须先解除全部绑定",
None,
)));
}
}
let deleted = state.delete_oauth_provider_config(&provider_type).await?;
if !deleted {
return Ok(Some(build_proxy_error_response(
http::StatusCode::BAD_REQUEST,
"invalid_request",
@@ -18,6 +18,7 @@ pub(crate) async fn maybe_build_local_admin_auth_response(
&request.state(),
&request.request_context(),
request.request_headers(),
request.remote_addr(),
request.request_body(),
)
.await?
@@ -5,6 +5,7 @@ use super::super::{
use crate::handlers::admin::request::AdminAppState;
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
use crate::GatewayError;
use aether_data::repository::wallet::stored_timestamp_unix_secs;
use axum::{
body::{Body, Bytes},
http,
@@ -53,7 +54,7 @@ pub(super) fn build_admin_billing_collector_payload_from_record(
"default_value": record.default_value,
"priority": record.priority,
"is_enabled": record.is_enabled,
"created_at": unix_secs_to_rfc3339(record.created_at_unix_ms),
"created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(record.created_at_unix_ms)),
"updated_at": unix_secs_to_rfc3339(record.updated_at_unix_secs),
})
}
@@ -174,17 +175,7 @@ pub(super) async fn parse_admin_billing_collector_request(
));
}
Ok(false) => {}
Err(err) => {
let detail = match err {
GatewayError::Internal(message) => message,
other => format!("{other:?}"),
};
return Err((
http::StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({ "detail": detail })),
)
.into_response());
}
Err(_err) => return Err(build_admin_billing_internal_error_response()),
}
}
@@ -202,6 +193,16 @@ pub(super) async fn parse_admin_billing_collector_request(
})
}
fn build_admin_billing_internal_error_response() -> Response<Body> {
(
http::StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({
"detail": "计费采集器服务暂不可用,请稍后重试"
})),
)
.into_response()
}
pub(in super::super) fn admin_billing_parse_page(query: Option<&str>) -> Result<u32, String> {
super::super::admin_billing_parse_page(query)
}
@@ -20,6 +20,7 @@ mod routes;
mod rules;
mod wallets;
pub(in crate::handlers::admin) use self::payments::admin_payment_gateway_response_projection;
pub(super) use self::payments::maybe_build_local_admin_payments_response;
pub(super) use self::routes::maybe_build_local_admin_billing_routes_response;
pub(super) use self::wallets::maybe_build_local_admin_wallets_response;
@@ -3,12 +3,16 @@ use super::{
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::shared::{
normalize_payment_callback_base_url, normalize_payment_currency, normalize_payment_https_url,
payment_gateway_allow_user_refund, payment_gateway_channels_config_json,
payment_gateway_channels_json, payment_gateway_config_json, payment_gateway_refund_enabled,
payment_gateway_secret_keys_json,
payment_gateway_secret_is_legacy_unbound, payment_gateway_secret_keys_json,
PaymentGatewaySecretBinding,
};
use crate::{GatewayError, LocalMutationOutcome};
use aether_data_contracts::repository::billing::PaymentGatewayConfigWriteInput;
use aether_data_contracts::repository::billing::{
PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigWriteInput,
};
use axum::{
body::Body,
http,
@@ -18,7 +22,9 @@ use axum::{
use serde::Deserialize;
use serde_json::{json, Value};
#[derive(Debug, Deserialize)]
const PAYMENT_GATEWAY_CONFIG_CAS_MAX_ATTEMPTS: usize = 8;
#[derive(Deserialize)]
struct PaymentGatewayConfigRequest {
#[serde(default)]
enabled: bool,
@@ -60,6 +66,14 @@ fn default_min_recharge_usd() -> f64 {
1.0
}
fn build_payment_gateway_conflict_response(detail: impl Into<String>) -> Response<Body> {
(
http::StatusCode::CONFLICT,
Json(json!({ "detail": detail.into() })),
)
.into_response()
}
fn default_channels() -> Value {
json!([
{"channel": "alipay", "display_name": "支付宝", "fee_rate": 0.0},
@@ -121,6 +135,18 @@ fn admin_payment_gateway_provider_from_path(path: &str) -> Option<String> {
Some(provider)
}
fn resolve_admin_payment_gateway_provider(path: &str, route_kind: &str) -> Option<String> {
match route_kind {
"get_epay_gateway" | "update_epay_gateway" | "test_epay_gateway" => {
Some("epay".to_string())
}
"get_payment_gateway" | "update_payment_gateway" | "test_payment_gateway" => {
admin_payment_gateway_provider_from_path(path)
}
_ => None,
}
}
fn default_provider_channels(provider: &str) -> Value {
match provider {
"epay" => default_channels(),
@@ -263,29 +289,97 @@ fn normalize_config_object(config: Value) -> Result<Value, String> {
Err("config must be an object".to_string())
}
fn merge_gateway_secret_maps(
existing_plaintext: Option<&str>,
updates: serde_json::Map<String, Value>,
) -> Result<serde_json::Map<String, Value>, &'static str> {
let mut merged = match existing_plaintext {
Some(plaintext) => serde_json::from_str::<Value>(plaintext)
.ok()
.and_then(|value| value.as_object().cloned())
.ok_or("existing gateway secrets have invalid format")?,
None => serde_json::Map::new(),
};
merged.extend(updates);
Ok(merged)
}
/// A legacy gateway ciphertext has no authenticated destination (or only the
/// provider in v2). Reusing it while changing endpoint/merchant would carry
/// an unknown credential into a different payment account. Require the
/// administrator to provide a replacement secret in that case.
fn legacy_secret_reuse_requires_reentry(
existing: Option<&aether_data_contracts::repository::billing::PaymentGatewayConfigRecord>,
requested_binding: &PaymentGatewaySecretBinding,
) -> bool {
let Some(record) = existing else {
return false;
};
let Some(ciphertext) = record.merchant_key_encrypted.as_deref() else {
return false;
};
if !payment_gateway_secret_is_legacy_unbound(ciphertext) {
return false;
}
// An invalid historical binding cannot establish that the legacy value
// belongs to the requested destination, so fail closed as well.
PaymentGatewaySecretBinding::from_record(record)
.map(|stored_binding| stored_binding != requested_binding.clone())
.unwrap_or(true)
}
fn encrypted_gateway_secret(
state: &AdminAppState<'_>,
provider: &str,
binding: &PaymentGatewaySecretBinding,
payload: &PaymentGatewayConfigRequest,
) -> Result<Option<String>, Response<Body>> {
existing: Option<&aether_data_contracts::repository::billing::PaymentGatewayConfigRecord>,
) -> Result<(Option<String>, Vec<Value>), Response<Body>> {
let provider = binding.provider.as_str();
let decrypt_existing = || {
if legacy_secret_reuse_requires_reentry(existing, binding) {
return Err(build_admin_payments_bad_request_response(
"endpoint_url or merchant_id changed; re-enter the gateway secret",
));
}
existing
.and_then(|record| record.merchant_key_encrypted.as_deref())
.map(|ciphertext| {
crate::handlers::shared::open_payment_gateway_secret(
state.app(), binding, ciphertext,
)
.map(|projection| projection.plaintext)
.map_err(|_| {
build_admin_payments_backend_unavailable_response(
"existing gateway secrets are not valid for the requested destination; re-enter the secret",
)
})
})
.transpose()
};
let secret_plaintext = if provider == "epay" {
payload
let supplied = payload
.merchant_key
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.map(ToOwned::to_owned);
if supplied.is_none() {
decrypt_existing()?;
}
supplied
} else {
let Some(secrets) = payload.secrets.as_object() else {
return if payload.secrets.is_null() {
Ok(None)
decrypt_existing()?;
Ok((None, existing_gateway_secret_keys(existing)))
} else {
Err(build_admin_payments_bad_request_response(
"secrets must be an object",
))
};
};
let filtered = secrets
let updates = secrets
.iter()
.filter_map(|(key, value)| {
let value = value.as_str()?.trim();
@@ -293,39 +387,60 @@ fn encrypted_gateway_secret(
.then(|| (key.trim().to_string(), Value::String(value.to_string())))
})
.collect::<serde_json::Map<_, _>>();
if filtered.is_empty() {
None
} else {
Some(Value::Object(filtered).to_string())
if updates.is_empty() {
decrypt_existing()?;
return Ok((None, existing_gateway_secret_keys(existing)));
}
let existing_plaintext = decrypt_existing()?;
let merged = match merge_gateway_secret_maps(existing_plaintext.as_deref(), updates) {
Ok(value) => value,
Err(detail) => {
return Err(build_admin_payments_backend_unavailable_response(detail));
}
};
Some(Value::Object(merged).to_string())
};
let Some(secret_plaintext) = secret_plaintext else {
return Ok(None);
return Ok((None, existing_gateway_secret_keys(existing)));
};
state
.encrypt_catalog_secret_with_fallbacks(&secret_plaintext)
.ok_or_else(|| {
build_admin_payments_backend_unavailable_response("encryption key is not configured")
})
.map(Some)
let encrypted = crate::handlers::shared::seal_payment_gateway_secret(
state.app(),
binding,
&secret_plaintext,
)
.map_err(build_admin_payments_backend_unavailable_response)?;
let secret_keys = if provider == "epay" {
Vec::new()
} else {
let mut keys = serde_json::from_str::<Value>(&secret_plaintext)
.ok()
.and_then(|value| value.as_object().cloned())
.unwrap_or_default()
.into_iter()
.map(|(key, _)| Value::String(key))
.collect::<Vec<_>>();
keys.sort_by(|left, right| left.as_str().cmp(&right.as_str()));
keys
};
Ok((Some(encrypted), secret_keys))
}
async fn existing_gateway_secret_keys(
state: &AdminAppState<'_>,
provider: &str,
) -> Result<Vec<Value>, GatewayError> {
let Some(record) = state.app().find_payment_gateway_config(provider).await? else {
return Ok(Vec::new());
fn existing_gateway_secret_keys(
record: Option<&aether_data_contracts::repository::billing::PaymentGatewayConfigRecord>,
) -> Vec<Value> {
let Some(record) = record else {
return Vec::new();
};
let (_, _, secret_keys, _, _) = split_gateway_channels_config(&record);
Ok(secret_keys
let (_, _, secret_keys, _, _) = split_gateway_channels_config(record);
secret_keys
.as_array()
.cloned()
.unwrap_or_default()
.into_iter()
.filter(|value| value.as_str().is_some_and(|item| !item.trim().is_empty()))
.collect())
.collect()
}
pub(super) async fn maybe_build_local_admin_payment_gateways_response(
@@ -336,8 +451,14 @@ pub(super) async fn maybe_build_local_admin_payment_gateways_response(
) -> Result<Option<Response<Body>>, GatewayError> {
match route_kind {
Some("get_epay_gateway") | Some("get_payment_gateway") => {
let provider = admin_payment_gateway_provider_from_path(request_context.path())
.unwrap_or_else(|| "epay".to_string());
let Some(provider) = resolve_admin_payment_gateway_provider(
request_context.path(),
route_kind.expect("matched payment gateway route kind"),
) else {
return Ok(Some(build_admin_payments_bad_request_response(
"unsupported payment gateway provider",
)));
};
let record = state.app().find_payment_gateway_config(&provider).await?;
let payload = record
.map(gateway_config_payload)
@@ -345,8 +466,14 @@ pub(super) async fn maybe_build_local_admin_payment_gateways_response(
Ok(Some(Json(payload).into_response()))
}
Some("update_epay_gateway") | Some("update_payment_gateway") => {
let provider = admin_payment_gateway_provider_from_path(request_context.path())
.unwrap_or_else(|| "epay".to_string());
let Some(provider) = resolve_admin_payment_gateway_provider(
request_context.path(),
route_kind.expect("matched payment gateway route kind"),
) else {
return Ok(Some(build_admin_payments_bad_request_response(
"unsupported payment gateway provider",
)));
};
let Some(body) = request_body else {
return Ok(Some(build_admin_payments_bad_request_response(
"缺少请求体",
@@ -371,109 +498,167 @@ pub(super) async fn maybe_build_local_admin_payment_gateways_response(
)));
}
let merchant_key_encrypted = match encrypted_gateway_secret(state, &provider, &payload)
{
Ok(value) => value,
Err(response) => return Ok(Some(response)),
};
let endpoint_url = if provider == "epay" {
match normalize_text(payload.endpoint_url, "endpoint_url", 512) {
Ok(value) => value,
match normalize_text(payload.endpoint_url.clone(), "endpoint_url", 512) {
Ok(value) => match normalize_payment_https_url(&value, "endpoint_url") {
Ok(value) => value,
Err(detail) => {
return Ok(Some(build_admin_payments_bad_request_response(detail)))
}
},
Err(detail) => {
return Ok(Some(build_admin_payments_bad_request_response(detail)))
}
}
} else {
match normalize_optional_text(Some(payload.endpoint_url), 512) {
Ok(value) => value.unwrap_or_default(),
match normalize_optional_text(Some(payload.endpoint_url.clone()), 512) {
Ok(Some(value)) => match normalize_payment_https_url(&value, "endpoint_url") {
Ok(value) => value,
Err(detail) => {
return Ok(Some(build_admin_payments_bad_request_response(detail)))
}
},
Ok(None) => String::new(),
Err(detail) => {
return Ok(Some(build_admin_payments_bad_request_response(detail)))
}
}
};
let callback_base_url = match normalize_optional_text(payload.callback_base_url, 512) {
Ok(value) => value,
Err(detail) => return Ok(Some(build_admin_payments_bad_request_response(detail))),
};
let callback_base_url =
match normalize_optional_text(payload.callback_base_url.clone(), 512) {
Ok(Some(value)) => match normalize_payment_callback_base_url(&value) {
Ok(value) => Some(value),
Err(detail) => {
return Ok(Some(build_admin_payments_bad_request_response(detail)))
}
},
Ok(None) => None,
Err(detail) => {
return Ok(Some(build_admin_payments_bad_request_response(detail)))
}
};
let merchant_id = if provider == "epay" {
match normalize_text(payload.merchant_id, "merchant_id", 128) {
match normalize_text(payload.merchant_id.clone(), "merchant_id", 128) {
Ok(value) => value,
Err(detail) => {
return Ok(Some(build_admin_payments_bad_request_response(detail)))
}
}
} else {
match normalize_optional_text(Some(payload.merchant_id), 128) {
match normalize_optional_text(Some(payload.merchant_id.clone()), 128) {
Ok(value) => value.unwrap_or_default(),
Err(detail) => {
return Ok(Some(build_admin_payments_bad_request_response(detail)))
}
}
};
let pay_currency = match normalize_text(payload.pay_currency, "pay_currency", 16) {
let pay_currency =
match normalize_payment_currency(&payload.pay_currency, "pay_currency") {
Ok(value) => value,
Err(detail) => {
return Ok(Some(build_admin_payments_bad_request_response(detail)))
}
};
let config = match normalize_config_object(payload.config.clone()) {
Ok(value) => value,
Err(detail) => return Ok(Some(build_admin_payments_bad_request_response(detail))),
};
let config = match normalize_config_object(payload.config) {
Ok(value) => value,
Err(detail) => return Ok(Some(build_admin_payments_bad_request_response(detail))),
};
let submitted_secret_keys = payload
.secrets
.as_object()
.map(|secrets| {
secrets
.iter()
.filter(|(_, value)| {
value.as_str().is_some_and(|value| !value.trim().is_empty())
})
.map(|(key, _)| Value::String(key.clone()))
.collect::<Vec<_>>()
})
.unwrap_or_default();
let secret_keys = if provider == "epay" || !submitted_secret_keys.is_empty() {
submitted_secret_keys
} else {
existing_gateway_secret_keys(state, &provider).await?
};
let channels = match normalize_gateway_channels(&provider, payload.channels) {
let channels = match normalize_gateway_channels(&provider, payload.channels.clone()) {
Ok(value) => value,
Err(detail) => return Ok(Some(build_admin_payments_bad_request_response(detail))),
};
let refund_enabled = payload.refund_enabled;
let allow_user_refund = refund_enabled && payload.allow_user_refund;
let channels_json = payment_gateway_channels_config_json(
channels,
config,
Value::Array(secret_keys),
refund_enabled,
allow_user_refund,
);
let input = PaymentGatewayConfigWriteInput {
provider: provider.clone(),
enabled: payload.enabled,
endpoint_url,
callback_base_url,
merchant_id,
preserve_existing_secret: merchant_key_encrypted.is_none(),
merchant_key_encrypted,
pay_currency,
usd_exchange_rate: payload.usd_exchange_rate,
min_recharge_usd: payload.min_recharge_usd,
channels_json,
};
match state.app().upsert_payment_gateway_config(&input).await? {
LocalMutationOutcome::Applied(record) => {
Ok(Some(Json(gateway_config_payload(record)).into_response()))
let binding =
match PaymentGatewaySecretBinding::new(&provider, &endpoint_url, &merchant_id) {
Ok(value) => value,
Err(detail) => {
return Ok(Some(build_admin_payments_bad_request_response(detail)))
}
};
let mut existing_record = state.app().find_payment_gateway_config(&provider).await?;
let expected_existing = existing_record.is_some();
for _ in 0..PAYMENT_GATEWAY_CONFIG_CAS_MAX_ATTEMPTS {
if expected_existing && existing_record.is_none() {
return Ok(Some(build_payment_gateway_conflict_response(
"payment gateway config was removed concurrently",
)));
}
let (merchant_key_encrypted, secret_keys) = match encrypted_gateway_secret(
state,
&binding,
&payload,
existing_record.as_ref(),
) {
Ok(value) => value,
Err(response) => return Ok(Some(response)),
};
let channels_json = payment_gateway_channels_config_json(
channels.clone(),
config.clone(),
Value::Array(secret_keys),
refund_enabled,
allow_user_refund,
);
let mutation = PaymentGatewayConfigCasWriteInput {
input: PaymentGatewayConfigWriteInput {
provider: provider.clone(),
enabled: payload.enabled,
endpoint_url: endpoint_url.clone(),
callback_base_url: callback_base_url.clone(),
merchant_id: merchant_id.clone(),
preserve_existing_secret: merchant_key_encrypted.is_none(),
merchant_key_encrypted,
pay_currency: pay_currency.clone(),
usd_exchange_rate: payload.usd_exchange_rate,
min_recharge_usd: payload.min_recharge_usd,
channels_json,
},
expected_existing,
expected_merchant_key_encrypted: existing_record
.as_ref()
.and_then(|record| record.merchant_key_encrypted.clone()),
};
match state
.app()
.compare_and_swap_payment_gateway_config(&mutation)
.await?
{
LocalMutationOutcome::Applied(record) => {
return Ok(Some(Json(gateway_config_payload(record)).into_response()));
}
LocalMutationOutcome::NotFound if !expected_existing => {
return Ok(Some(build_payment_gateway_conflict_response(
"payment gateway config was created concurrently",
)));
}
LocalMutationOutcome::NotFound => {
existing_record =
state.app().find_payment_gateway_config(&provider).await?;
}
LocalMutationOutcome::Invalid(detail) => {
return Ok(Some(build_admin_payments_bad_request_response(detail)));
}
LocalMutationOutcome::Unavailable => {
return Ok(Some(build_admin_payments_backend_unavailable_response(
"payment gateway config backend unavailable",
)));
}
}
_ => Ok(Some(build_admin_payments_backend_unavailable_response(
"payment gateway config backend unavailable",
))),
}
Ok(Some(build_payment_gateway_conflict_response(
"payment gateway config changed too frequently; retry the request",
)))
}
Some("test_epay_gateway") | Some("test_payment_gateway") => {
let provider = admin_payment_gateway_provider_from_path(request_context.path())
.unwrap_or_else(|| "epay".to_string());
let Some(provider) = resolve_admin_payment_gateway_provider(
request_context.path(),
route_kind.expect("matched payment gateway route kind"),
) else {
return Ok(Some(build_admin_payments_bad_request_response(
"unsupported payment gateway provider",
)));
};
let status = state.app().find_payment_gateway_config(&provider).await?;
let ok = status
.as_ref()
@@ -493,3 +678,164 @@ pub(super) async fn maybe_build_local_admin_payment_gateways_response(
_ => Ok(None),
}
}
#[cfg(test)]
mod tests {
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
use aether_data_contracts::repository::billing::PaymentGatewayConfigRecord;
use serde_json::{json, Value};
use super::{
legacy_secret_reuse_requires_reentry, merge_gateway_secret_maps,
resolve_admin_payment_gateway_provider,
};
use crate::handlers::shared::PaymentGatewaySecretBinding;
fn gateway_record(
endpoint_url: &str,
merchant_id: &str,
merchant_key_encrypted: Option<String>,
) -> PaymentGatewayConfigRecord {
PaymentGatewayConfigRecord {
provider: "stripe".to_string(),
enabled: true,
endpoint_url: endpoint_url.to_string(),
callback_base_url: None,
merchant_id: merchant_id.to_string(),
merchant_key_encrypted,
pay_currency: "USD".to_string(),
usd_exchange_rate: 1.0,
min_recharge_usd: 1.0,
channels_json: json!({}),
created_at_unix_secs: 1,
updated_at_unix_secs: 1,
}
}
#[test]
fn legacy_secret_reuse_requires_reentry_after_binding_change() {
let legacy = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "legacy-secret")
.expect("legacy secret should encrypt");
let old_record = gateway_record("https://api.stripe.com", "merchant-old", Some(legacy));
let changed_binding =
PaymentGatewaySecretBinding::new("stripe", "https://api.stripe.com", "merchant-new")
.expect("changed binding should be valid");
assert!(legacy_secret_reuse_requires_reentry(
Some(&old_record),
&changed_binding,
));
let v2_record = gateway_record(
"https://api.stripe.com",
"merchant-old",
Some("aether-payment-gateway-secret-v2:legacy".to_string()),
);
assert!(legacy_secret_reuse_requires_reentry(
Some(&v2_record),
&changed_binding,
));
}
#[test]
fn legacy_secret_reuse_is_allowed_only_for_same_binding_or_bound_v3() {
let legacy = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "legacy-secret")
.expect("legacy secret should encrypt");
let old_record =
gateway_record("https://API.STRIPE.COM:443/", "merchant-old", Some(legacy));
let same_binding = PaymentGatewaySecretBinding::new(
"stripe",
"https://api.stripe.com:443/",
" merchant-old ",
)
.expect("same binding should be valid");
assert!(!legacy_secret_reuse_requires_reentry(
Some(&old_record),
&same_binding,
));
let v3_record = gateway_record(
"https://api.stripe.com",
"merchant-old",
Some("aether-payment-gateway-secret-v3:bound".to_string()),
);
let changed_binding =
PaymentGatewaySecretBinding::new("stripe", "https://api.stripe.com", "merchant-new")
.expect("changed binding should be valid");
assert!(!legacy_secret_reuse_requires_reentry(
Some(&v3_record),
&changed_binding,
));
}
#[test]
fn stripe_secret_rotation_preserves_omitted_secret_fields() {
let existing = json!({
"secret_key": "old-secret-key",
"webhook_secret": "old-webhook"
})
.to_string();
let updates = json!({"webhook_secret": "new-webhook"})
.as_object()
.cloned()
.expect("updates should be an object");
let merged = Value::Object(
merge_gateway_secret_maps(Some(&existing), updates)
.expect("valid secret maps should merge"),
);
assert_eq!(merged["secret_key"], "old-secret-key");
assert_eq!(merged["webhook_secret"], "new-webhook");
}
#[test]
fn wxpay_secret_rotation_preserves_omitted_secret_fields() {
let existing = json!({
"private_key": "old-private",
"api_v3_key": "old-api-v3-key",
"public_key": "old-public"
})
.to_string();
let updates = json!({"api_v3_key": "new-api-v3-key"})
.as_object()
.cloned()
.expect("updates should be an object");
let merged = Value::Object(
merge_gateway_secret_maps(Some(&existing), updates)
.expect("valid secret maps should merge"),
);
assert_eq!(merged["private_key"], "old-private");
assert_eq!(merged["api_v3_key"], "new-api-v3-key");
assert_eq!(merged["public_key"], "old-public");
}
#[test]
fn generic_gateway_routes_never_fall_back_to_epay() {
assert_eq!(
resolve_admin_payment_gateway_provider(
"/api/admin/payments/gateways/stripe",
"update_payment_gateway",
)
.as_deref(),
Some("stripe")
);
assert!(resolve_admin_payment_gateway_provider(
"/api/admin/payments/gateways/unsupported",
"update_payment_gateway",
)
.is_none());
assert!(resolve_admin_payment_gateway_provider(
"/api/admin/payments/gateways/stripe/extra",
"get_payment_gateway",
)
.is_none());
assert_eq!(
resolve_admin_payment_gateway_provider(
"/api/admin/payments/epay",
"update_epay_gateway",
)
.as_deref(),
Some("epay")
);
}
}
@@ -11,6 +11,7 @@ mod redeem_codes;
mod routes;
mod shared;
pub(in crate::handlers::admin) use self::shared::admin_payment_gateway_response_projection;
use self::shared::{
admin_payment_operator_id, admin_payment_order_id_from_detail_path,
admin_payment_order_id_from_suffix_path, build_admin_payment_callback_payload_from_record,
@@ -19,7 +20,8 @@ use self::shared::{
build_admin_payments_bad_request_response, build_admin_payments_data_unavailable_response,
normalize_admin_payment_currency, normalize_admin_payment_optional_string,
normalize_admin_payment_positive_number, parse_admin_payments_limit,
parse_admin_payments_offset, AdminPaymentOrderCreditRequest,
parse_admin_payments_offset, prepare_admin_payment_gateway_response_for_storage,
AdminPaymentOrderCreditRequest,
};
pub(crate) async fn maybe_build_local_admin_payments_response(
@@ -5,7 +5,8 @@ use super::{
build_admin_payments_backend_unavailable_response, build_admin_payments_bad_request_response,
normalize_admin_payment_currency, normalize_admin_payment_optional_string,
normalize_admin_payment_positive_number, parse_admin_payments_limit,
parse_admin_payments_offset, AdminPaymentOrderCreditRequest,
parse_admin_payments_offset, prepare_admin_payment_gateway_response_for_storage,
AdminPaymentOrderCreditRequest,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{attach_admin_audit_response, query_param_value};
@@ -124,7 +125,9 @@ async fn close_direct_gateway_order_before_terminal_mark(
)))
}
};
if order.status != "pending" || !matches!(order.payment_method.as_str(), "alipay" | "wxpay") {
if order.status != "pending"
|| !matches!(order.payment_method.as_str(), "alipay" | "wxpay" | "stripe")
{
return Ok(None);
}
crate::handlers::shared::close_direct_gateway_order(state.app(), &order)
@@ -231,6 +234,8 @@ async fn build_admin_payment_credit_order_response(
"gateway_response 必须为对象",
));
}
let gateway_response =
prepare_admin_payment_gateway_response_for_storage(payload.gateway_response);
let operator_id = admin_payment_operator_id(request_context);
match state
.admin_credit_payment_order(
@@ -239,7 +244,7 @@ async fn build_admin_payment_credit_order_response(
pay_amount,
pay_currency.as_deref(),
exchange_rate,
payload.gateway_response,
gateway_response,
operator_id.as_deref(),
)
.await?
@@ -8,6 +8,7 @@ use crate::handlers::admin::shared::{
attach_admin_audit_response, query_param_value, unix_secs_to_rfc3339,
};
use crate::GatewayError;
use aether_data::repository::wallet::stored_timestamp_unix_secs;
use axum::{
body::Body,
http,
@@ -122,7 +123,7 @@ fn build_batch_payload(
"description": batch.description,
"created_by": batch.created_by,
"expires_at": batch.expires_at_unix_secs.and_then(unix_secs_to_rfc3339),
"created_at": unix_secs_to_rfc3339(batch.created_at_unix_ms),
"created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(batch.created_at_unix_ms)),
"updated_at": unix_secs_to_rfc3339(batch.updated_at_unix_secs),
})
}
@@ -146,7 +147,7 @@ fn build_code_payload(
"redeemed_at": code.redeemed_at_unix_secs.and_then(unix_secs_to_rfc3339),
"disabled_by": code.disabled_by,
"expires_at": code.expires_at_unix_secs.and_then(unix_secs_to_rfc3339),
"created_at": unix_secs_to_rfc3339(code.created_at_unix_ms),
"created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(code.created_at_unix_ms)),
"updated_at": unix_secs_to_rfc3339(code.updated_at_unix_secs),
})
}
@@ -1,17 +1,19 @@
use crate::handlers::admin::request::AdminRequestContext;
use crate::handlers::admin::shared::{query_param_value, unix_secs_to_rfc3339};
use crate::handlers::shared::normalize_payment_currency;
use crate::GatewayAdminPaymentCallbackView;
use aether_data::repository::wallet::stored_timestamp_unix_secs;
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use serde_json::{json, Value};
const ADMIN_PAYMENTS_DATA_UNAVAILABLE_DETAIL: &str = "Admin payments data unavailable";
#[derive(Debug, Default, serde::Deserialize)]
#[derive(Default, serde::Deserialize)]
pub(super) struct AdminPaymentOrderCreditRequest {
#[serde(default)]
pub(super) gateway_order_id: Option<String>,
@@ -25,6 +27,22 @@ pub(super) struct AdminPaymentOrderCreditRequest {
pub(super) gateway_response: Option<serde_json::Value>,
}
impl std::fmt::Debug for AdminPaymentOrderCreditRequest {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("AdminPaymentOrderCreditRequest")
.field("gateway_order_id", &self.gateway_order_id)
.field("pay_amount", &self.pay_amount)
.field("pay_currency", &self.pay_currency)
.field("exchange_rate", &self.exchange_rate)
.field(
"gateway_response",
&self.gateway_response.as_ref().map(|_| "[REDACTED]"),
)
.finish()
}
}
pub(super) fn build_admin_payments_data_unavailable_response() -> Response<Body> {
(
http::StatusCode::SERVICE_UNAVAILABLE,
@@ -153,11 +171,9 @@ pub(super) fn normalize_admin_payment_currency(
let Some(value) = normalize_admin_payment_optional_string(value, "pay_currency", 3)? else {
return Ok(None);
};
let normalized = value.to_ascii_uppercase();
if normalized.len() != 3 {
return Err("pay_currency 必须是 3 位货币代码".to_string());
}
Ok(Some(normalized))
normalize_payment_currency(&value, "pay_currency")
.map(Some)
.map_err(|_| "pay_currency 必须是 3 位 ASCII 货币代码".to_string())
}
pub(super) fn normalize_admin_payment_positive_number(
@@ -187,13 +203,184 @@ pub(super) fn admin_payment_effective_status(
expires_at_unix_secs: Option<u64>,
) -> String {
let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
if status == "pending" && expires_at_unix_secs.is_some_and(|value| value < now_unix_secs) {
if status == "pending" && expires_at_unix_secs.is_some_and(|value| value <= now_unix_secs) {
"expired".to_string()
} else {
status.to_string()
}
}
fn admin_payment_bounded_string(value: &Value, max_chars: usize) -> Option<Value> {
let value = value.as_str()?.trim();
(!value.is_empty() && value.chars().count() <= max_chars)
.then(|| Value::String(value.to_string()))
}
fn admin_payment_identifier(value: &Value, max_chars: usize) -> Option<Value> {
let value = value.as_str()?.trim();
(!value.is_empty()
&& value.chars().count() <= max_chars
&& value
.chars()
.all(|character| character.is_ascii_alphanumeric() || matches!(character, '_' | '-')))
.then(|| Value::String(value.to_string()))
}
fn admin_payment_gateway_response_field(key: &str, value: &Value) -> Option<Value> {
match key {
"gateway" | "submit_method" | "payment_channel" => admin_payment_identifier(value, 64),
"pay_currency" => admin_payment_identifier(value, 16),
"display_name" | "provider_label" => admin_payment_bounded_string(value, 128),
"gateway_order_id" | "intent_id" => admin_payment_bounded_string(value, 256),
"expires_at" => admin_payment_bounded_string(value, 64),
"pay_amount" | "base_pay_amount" | "fee_rate" | "fee_amount" => {
value.is_number().then(|| value.clone())
}
"manual_credit" => value.as_bool().map(Value::Bool),
"payment_method_types" => {
let values = value.as_array()?;
if values.len() > 16 {
return None;
}
values
.iter()
.map(|value| admin_payment_identifier(value, 64))
.collect::<Option<Vec<_>>>()
.map(Value::Array)
}
_ => None,
}
}
pub(in crate::handlers::admin) fn admin_payment_gateway_response_projection(
value: Option<&Value>,
) -> Value {
let Some(object) = value.and_then(Value::as_object) else {
return Value::Null;
};
Value::Object(
object
.iter()
.filter_map(|(key, value)| {
admin_payment_gateway_response_field(key, value).map(|value| (key.clone(), value))
})
.collect(),
)
}
pub(super) fn prepare_admin_payment_gateway_response_for_storage(
value: Option<Value>,
) -> Option<Value> {
value.map(|value| admin_payment_gateway_response_projection(Some(&value)))
}
#[derive(Default)]
struct AdminPaymentJsonShape {
objects: u64,
arrays: u64,
strings: u64,
numbers: u64,
booleans: u64,
nulls: u64,
object_fields: u64,
array_items: u64,
max_depth: u64,
}
impl AdminPaymentJsonShape {
fn observe(&mut self, value: &Value, depth: u64) {
self.max_depth = self.max_depth.max(depth);
match value {
Value::Object(object) => {
self.objects = self.objects.saturating_add(1);
self.object_fields = self
.object_fields
.saturating_add(u64::try_from(object.len()).unwrap_or(u64::MAX));
for value in object.values() {
self.observe(value, depth.saturating_add(1));
}
}
Value::Array(values) => {
self.arrays = self.arrays.saturating_add(1);
self.array_items = self
.array_items
.saturating_add(u64::try_from(values.len()).unwrap_or(u64::MAX));
for value in values {
self.observe(value, depth.saturating_add(1));
}
}
Value::String(_) => self.strings = self.strings.saturating_add(1),
Value::Number(_) => self.numbers = self.numbers.saturating_add(1),
Value::Bool(_) => self.booleans = self.booleans.saturating_add(1),
Value::Null => self.nulls = self.nulls.saturating_add(1),
}
}
}
fn admin_payment_json_kind(value: &Value) -> &'static str {
match value {
Value::Null => "null",
Value::Bool(_) => "boolean",
Value::Number(_) => "number",
Value::String(_) => "string",
Value::Array(_) => "array",
Value::Object(_) => "object",
}
}
fn admin_payment_payload_summary(value: Option<&Value>) -> Value {
let Some(value) = value else {
return Value::Null;
};
let mut shape = AdminPaymentJsonShape::default();
shape.observe(value, 1);
json!({
"kind": admin_payment_json_kind(value),
"serialized_bytes": serde_json::to_vec(value).map_or(0, |encoded| encoded.len()),
"objects": shape.objects,
"arrays": shape.arrays,
"strings": shape.strings,
"numbers": shape.numbers,
"booleans": shape.booleans,
"nulls": shape.nulls,
"object_fields": shape.object_fields,
"array_items": shape.array_items,
"max_depth": shape.max_depth,
})
}
fn admin_payment_callback_error_projection(value: Option<&str>) -> Option<String> {
const SAFE_ERRORS: &[&str] = &[
"callback amount mismatch",
"callback key reused with different payment payload",
"invalid callback signature",
"invalid payment callback numeric or identity fields",
"payment channel mismatch",
"payment currency mismatch",
"payment gateway order belongs to another payment order",
"payment gateway order identifier mismatch",
"payment gateway order mismatch",
"payment method mismatch",
"payment order expired",
"payment order not found",
"payment order number mismatch",
"payment order user missing",
"payment provider mismatch",
"plan purchase limit reached",
"wallet is not active",
"wallet not found",
];
let value = value?.trim();
if SAFE_ERRORS.contains(&value) {
return Some(value.to_string());
}
if value.starts_with("payment order is not creditable:") {
return Some("payment order is not creditable".to_string());
}
Some("payment callback processing failed".to_string())
}
pub(super) fn build_admin_payment_order_payload(
record: &crate::AdminWalletPaymentOrderRecord,
) -> serde_json::Value {
@@ -210,9 +397,10 @@ pub(super) fn build_admin_payment_order_payload(
"refundable_amount_usd": record.refundable_amount_usd,
"payment_method": record.payment_method,
"gateway_order_id": record.gateway_order_id,
"gateway_response": record.gateway_response,
"gateway_response": admin_payment_gateway_response_projection(record.gateway_response.as_ref()),
"has_gateway_response": record.gateway_response.is_some(),
"status": admin_payment_effective_status(&record.status, record.expires_at_unix_secs),
"created_at": unix_secs_to_rfc3339(record.created_at_unix_ms),
"created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(record.created_at_unix_ms)),
"paid_at": record.paid_at_unix_secs.and_then(unix_secs_to_rfc3339),
"credited_at": record.credited_at_unix_secs.and_then(unix_secs_to_rfc3339),
"expires_at": record.expires_at_unix_secs.and_then(unix_secs_to_rfc3339),
@@ -232,9 +420,196 @@ pub(super) fn build_admin_payment_callback_payload_from_record(
"payload_hash": record.payload_hash,
"signature_valid": record.signature_valid,
"status": record.status,
"payload": record.payload,
"error_message": record.error_message,
"created_at": unix_secs_to_rfc3339(record.created_at_unix_ms),
"payload": Value::Null,
"has_payload": record.payload.is_some(),
"payload_summary": admin_payment_payload_summary(record.payload.as_ref()),
"error_message": admin_payment_callback_error_projection(record.error_message.as_deref()),
"has_error_message": record.error_message.is_some(),
"created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(record.created_at_unix_ms)),
"processed_at": record.processed_at_unix_secs.and_then(unix_secs_to_rfc3339),
})
}
#[cfg(test)]
mod tests {
use super::{
build_admin_payment_callback_payload_from_record, build_admin_payment_order_payload,
prepare_admin_payment_gateway_response_for_storage,
};
use crate::{AdminWalletPaymentOrderRecord, GatewayAdminPaymentCallbackView};
use serde_json::json;
#[test]
fn admin_payment_order_projection_excludes_replayable_gateway_fields() {
let record = AdminWalletPaymentOrderRecord {
id: "order-1".to_string(),
order_no: "merchant-order-1".to_string(),
wallet_id: "wallet-1".to_string(),
user_id: Some("user-1".to_string()),
amount_usd: 10.0,
pay_amount: Some(72.0),
pay_currency: Some("CNY".to_string()),
exchange_rate: Some(7.2),
refunded_amount_usd: 0.0,
refundable_amount_usd: 0.0,
payment_method: "stripe".to_string(),
gateway_order_id: Some("pi_1".to_string()),
status: "pending".to_string(),
gateway_response: Some(json!({
"gateway": "stripe",
"intent_id": "pi_1",
"client_secret": "pi_1_secret_replayable",
"payment_url": "https://pay.example/checkout?token=secret",
"payment_params": {"sign": "signed-secret"},
"customer_email": "[email protected]"
})),
created_at_unix_ms: 1,
paid_at_unix_secs: None,
credited_at_unix_secs: None,
expires_at_unix_secs: None,
};
let payload = build_admin_payment_order_payload(&record);
assert_eq!(payload["has_gateway_response"], true);
assert_eq!(
payload.pointer("/gateway_response/gateway"),
Some(&json!("stripe"))
);
assert_eq!(
payload.pointer("/gateway_response/intent_id"),
Some(&json!("pi_1"))
);
for key in [
"client_secret",
"payment_url",
"payment_params",
"customer_email",
] {
assert!(payload
.pointer(&format!("/gateway_response/{key}"))
.is_none());
}
}
#[test]
fn admin_payment_order_projection_rejects_nested_or_mistyped_safe_fields() {
let mut record = AdminWalletPaymentOrderRecord {
id: "order-1".to_string(),
order_no: "merchant-order-1".to_string(),
wallet_id: "wallet-1".to_string(),
user_id: Some("user-1".to_string()),
amount_usd: 10.0,
pay_amount: Some(72.0),
pay_currency: Some("CNY".to_string()),
exchange_rate: Some(7.2),
refunded_amount_usd: 0.0,
refundable_amount_usd: 0.0,
payment_method: "stripe".to_string(),
gateway_order_id: Some("pi_1".to_string()),
status: "pending".to_string(),
gateway_response: None,
created_at_unix_ms: 1,
paid_at_unix_secs: None,
credited_at_unix_secs: None,
expires_at_unix_secs: None,
};
record.gateway_response = Some(json!({
"gateway": {"client_secret": "secret-in-nested-object"},
"intent_id": ["pi_1", "secret-in-array"],
"payment_method_types": ["card", {"secret": "nested"}],
"manual_credit": "secret-in-string",
}));
let encoded = build_admin_payment_order_payload(&record).to_string();
assert!(!encoded.contains("secret-in-nested-object"));
assert!(!encoded.contains("secret-in-array"));
assert!(!encoded.contains("nested"));
assert!(!encoded.contains("secret-in-string"));
}
#[test]
fn admin_payment_gateway_response_is_projected_before_storage() {
let projected = prepare_admin_payment_gateway_response_for_storage(Some(json!({
"gateway": "stripe",
"intent_id": "pi_1",
"client_secret": "pi_1_secret_replayable",
"customer": {"email": "[email protected]"},
"payment_params": {"authorization": "Bearer secret"},
})))
.expect("provided gateway response should remain present");
assert_eq!(projected, json!({"gateway": "stripe", "intent_id": "pi_1"}));
let encoded = projected.to_string();
for forbidden in [
"client_secret",
"replayable",
"customer",
"[email protected]",
"authorization",
"Bearer secret",
] {
assert!(!encoded.contains(forbidden), "persisted {forbidden}");
}
}
#[test]
fn admin_payment_callback_projection_does_not_return_raw_payload() {
let record = GatewayAdminPaymentCallbackView {
id: "callback-1".to_string(),
payment_order_id: Some("order-1".to_string()),
payment_method: "stripe".to_string(),
callback_key: "stripe:event-1".to_string(),
order_no: Some("merchant-order-1".to_string()),
gateway_order_id: Some("pi_1".to_string()),
payload_hash: Some("hash-1".to_string()),
signature_valid: true,
status: "processed".to_string(),
payload: Some(json!({
"data": {"object": {"client_secret": "secret", "customer_email": "[email protected]"}}
})),
error_message: None,
created_at_unix_ms: 1,
processed_at_unix_secs: Some(1),
};
let payload = build_admin_payment_callback_payload_from_record(&record);
assert_eq!(payload["has_payload"], true);
assert!(payload["payload"].is_null());
assert_eq!(payload["payload_summary"]["kind"], "object");
assert_eq!(payload["payload_summary"]["objects"], 3);
assert_eq!(payload["payload_summary"]["strings"], 2);
assert_eq!(payload["payload_summary"]["max_depth"], 4);
let encoded = payload.to_string();
assert!(!encoded.contains("customer_email"));
assert!(!encoded.contains("[email protected]"));
assert!(!encoded.contains("client_secret"));
assert!(!encoded.contains("secret"));
}
#[test]
fn admin_payment_callback_projection_does_not_return_unknown_historical_errors() {
let record = GatewayAdminPaymentCallbackView {
id: "callback-1".to_string(),
payment_order_id: None,
payment_method: "stripe".to_string(),
callback_key: "stripe:event-1".to_string(),
order_no: None,
gateway_order_id: None,
payload_hash: None,
signature_valid: false,
status: "failed".to_string(),
payload: None,
error_message: Some("upstream rejected sk_live_secret_value".to_string()),
created_at_unix_ms: 1,
processed_at_unix_secs: Some(1),
};
let payload = build_admin_payment_callback_payload_from_record(&record);
assert_eq!(payload["has_error_message"], true);
assert_eq!(
payload["error_message"],
"payment callback processing failed"
);
assert!(!payload.to_string().contains("sk_live_secret_value"));
}
}
@@ -3,8 +3,12 @@ use super::{
build_admin_billing_data_unavailable_response, build_admin_billing_not_found_response,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::shared::normalize_payment_currency;
use crate::{GatewayError, LocalMutationOutcome};
use aether_data_contracts::repository::billing::{BillingPlanRecord, BillingPlanWriteInput};
use aether_data_contracts::repository::billing::{
checked_plan_duration_days, parse_usage_policy_entitlements,
validate_entitlement_replacement_groups, BillingPlanRecord, BillingPlanWriteInput,
};
use axum::{
body::{Body, Bytes},
http,
@@ -172,11 +176,16 @@ fn validate_entitlements(value: &serde_json::Value) -> Result<(), String> {
}
}
}
"usage_policy" => {}
_ => return Err(format!("unsupported entitlement type: {kind}")),
}
}
validate_entitlement_replacement_groups(value).map_err(|error| error.to_string())?;
parse_usage_policy_entitlements(value).map_err(|error| error.to_string())?;
if !entitlements_include_package_rights(items) {
return Err("套餐至少需要包含每日额度或会员分组;钱包充值请使用充值功能".to_string());
return Err(
"套餐至少需要包含每日额度、会员分组或使用限制;钱包充值请使用充值功能".to_string(),
);
}
Ok(())
}
@@ -185,7 +194,7 @@ fn entitlements_include_package_rights(items: &[serde_json::Value]) -> bool {
items.iter().any(|item| {
matches!(
item.get("type").and_then(|value| value.as_str()),
Some("daily_quota" | "membership_group")
Some("daily_quota" | "membership_group" | "usage_policy")
)
})
}
@@ -204,6 +213,7 @@ fn normalize_plan_input(payload: BillingPlanRequest) -> Result<BillingPlanWriteI
if !matches!(duration_unit.as_str(), "day" | "month" | "year" | "custom") {
return Err("duration_unit must be day/month/year/custom".to_string());
}
checked_plan_duration_days(&duration_unit, payload.duration_value)?;
let purchase_limit_scope =
normalize_text(payload.purchase_limit_scope, "purchase_limit_scope", 32)?;
if !matches!(
@@ -213,11 +223,12 @@ fn normalize_plan_input(payload: BillingPlanRequest) -> Result<BillingPlanWriteI
return Err("purchase_limit_scope must be active_period/lifetime/unlimited".to_string());
}
validate_entitlements(&payload.entitlements)?;
let price_currency = normalize_payment_currency(&payload.price_currency, "price_currency")?;
Ok(BillingPlanWriteInput {
title: normalize_text(payload.title, "title", 128)?,
description: normalize_optional_text(payload.description, 2048)?,
price_amount: payload.price_amount,
price_currency: normalize_text(payload.price_currency, "price_currency", 16)?,
price_currency,
duration_unit,
duration_value: payload.duration_value,
enabled: payload.enabled,
@@ -9,6 +9,7 @@ use super::{
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
use crate::GatewayError;
use aether_data::repository::wallet::stored_timestamp_unix_secs;
use axum::{
body::{Body, Bytes},
http,
@@ -53,7 +54,7 @@ fn build_admin_billing_rule_payload_from_record(
"variables": record.variables,
"dimension_mappings": record.dimension_mappings,
"is_enabled": record.is_enabled,
"created_at": unix_secs_to_rfc3339(record.created_at_unix_ms),
"created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(record.created_at_unix_ms)),
"updated_at": unix_secs_to_rfc3339(record.updated_at_unix_secs),
})
}
@@ -10,6 +10,7 @@ use super::super::shared::{
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{attach_admin_audit_response, unix_secs_to_rfc3339};
use crate::GatewayError;
use aether_data::repository::wallet::stored_timestamp_unix_secs;
use axum::{
body::Body,
response::{IntoResponse, Response},
@@ -98,7 +99,7 @@ pub(in super::super) async fn build_admin_wallet_adjust_response(
transaction.link_id.as_deref(),
transaction.operator_id.as_deref(),
transaction.description.as_deref(),
unix_secs_to_rfc3339(transaction.created_at_unix_ms),
unix_secs_to_rfc3339(stored_timestamp_unix_secs(transaction.created_at_unix_ms)),
);
let response = Json(json!({
"wallet": wallet_payload,
@@ -13,12 +13,55 @@ use crate::handlers::shared::{
use crate::GatewayError;
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::{json, Value};
use tracing::warn;
fn is_safe_gateway_refund_id(value: &str) -> bool {
!value.is_empty()
&& value.len() <= 128
&& value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.'))
}
fn gateway_refund_mode_allowed(refund_mode: &str) -> bool {
refund_mode.trim().eq_ignore_ascii_case("original_channel")
}
fn stored_refund_to_gateway(
refund: aether_data::repository::wallet::StoredAdminWalletRefund,
) -> crate::AdminWalletRefundRecord {
crate::AdminWalletRefundRecord {
id: refund.id,
refund_no: refund.refund_no,
wallet_id: refund.wallet_id,
user_id: refund.user_id,
payment_order_id: refund.payment_order_id,
source_type: refund.source_type,
source_id: refund.source_id,
refund_mode: refund.refund_mode,
amount_usd: refund.amount_usd,
status: refund.status,
reason: refund.reason,
failure_reason: refund.failure_reason,
gateway_refund_id: refund.gateway_refund_id,
payout_method: refund.payout_method,
payout_reference: refund.payout_reference,
payout_proof: refund.payout_proof,
requested_by: refund.requested_by,
approved_by: refund.approved_by,
processed_by: refund.processed_by,
created_at_unix_ms: refund.created_at_unix_ms,
updated_at_unix_secs: refund.updated_at_unix_secs,
processed_at_unix_secs: refund.processed_at_unix_secs,
completed_at_unix_secs: refund.completed_at_unix_secs,
}
}
fn merge_gateway_refund_proof(
proof: Option<Value>,
gateway_refund: Option<&crate::handlers::shared::DirectGatewayRefundResult>,
@@ -29,14 +72,7 @@ fn merge_gateway_refund_proof(
let mut object = proof
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
object.insert(
"gateway_refund".to_string(),
json!({
"id": gateway_refund.gateway_refund_id,
"status": gateway_refund.status,
"payload": gateway_refund.payload,
}),
);
object.insert("gateway_refund".to_string(), gateway_refund.proof.clone());
Some(Value::Object(object))
}
@@ -64,7 +100,12 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response(
"gateway_refund_id",
128,
) {
Ok(value) => value,
Ok(value) if value.as_deref().is_none_or(is_safe_gateway_refund_id) => value,
Ok(_) => {
return Ok(build_admin_wallets_bad_request_response(
"gateway_refund_id 格式无效",
))
}
Err(detail) => return Ok(build_admin_wallets_bad_request_response(detail)),
};
let payout_reference = match normalize_admin_wallet_optional_text(
@@ -107,9 +148,55 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response(
else {
return Ok(build_admin_wallet_refund_not_found_response());
};
let refund_before_complete = stored_refund_to_gateway(refund_before_complete);
if !refund_before_complete.amount_usd.is_finite() || refund_before_complete.amount_usd <= 0.0 {
return Ok(build_admin_wallets_bad_request_response("退款金额无效"));
}
if refund_before_complete.status == "succeeded" {
if let Some(order_id) = refund_before_complete.payment_order_id.as_deref() {
if let Err(err) = state
.app()
.reverse_referral_rewards_for_order(order_id, refund_before_complete.amount_usd)
.await
{
warn!(
error = ?err,
order_id = %order_id,
refund_id = %refund_before_complete.id,
"failed to reconcile referral rewards for completed refund"
);
return Ok(build_admin_wallets_data_unavailable_response());
}
}
let response = Json(json!({
"refund": build_admin_wallet_refund_payload(
&wallet,
&owner,
&refund_before_complete,
),
}))
.into_response();
return Ok(attach_admin_audit_response(
response,
"admin_wallet_refund_completed",
"complete_wallet_refund",
"wallet_refund",
&refund_id,
));
}
let mut gateway_refund_id = gateway_refund_id;
let mut payout_proof = payload.payout_proof;
if payload.gateway_refund {
// A line-item refund in `offline_payout` mode has no provider-side
// settlement contract. Calling a gateway before recording evidence
// would let `/fail` concurrently release the local reservation and
// leave an external refund with no durable proof. Keep the mode
// constraint at the boundary, before any network request.
if !gateway_refund_mode_allowed(&refund_before_complete.refund_mode) {
return Ok(build_admin_wallets_bad_request_response(
"只有原支付渠道退款可以调用支付网关",
));
}
let Some(payment_order_id) = refund_before_complete.payment_order_id.as_deref() else {
return Ok(build_admin_wallets_bad_request_response(
"网关原路退款需要退款申请关联支付订单",
@@ -120,8 +207,10 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response(
crate::AdminWalletMutationOutcome::NotFound => {
return Ok(build_admin_wallets_bad_request_response("支付订单不存在"))
}
crate::AdminWalletMutationOutcome::Invalid(detail) => {
return Ok(build_admin_wallets_bad_request_response(detail))
crate::AdminWalletMutationOutcome::Invalid(_) => {
return Ok(build_admin_wallets_bad_request_response(
"支付订单状态或数据无效",
))
}
crate::AdminWalletMutationOutcome::Unavailable => {
return Ok(build_admin_wallets_data_unavailable_response())
@@ -150,14 +239,96 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response(
{
Ok(Some(result)) => {
gateway_refund_id = Some(result.gateway_refund_id.clone());
if result.is_pending() {
let persisted = match state
.app()
.update_admin_wallet_refund_gateway(
aether_data::repository::wallet::UpdateAdminWalletRefundGatewayInput {
wallet_id: wallet_id.clone(),
refund_id: refund_id.clone(),
gateway_refund_id: result.gateway_refund_id.clone(),
payout_proof: Some(result.proof.clone()),
},
)
.await?
{
Some(aether_data::repository::wallet::WalletMutationOutcome::Applied(
refund,
)) => refund,
Some(aether_data::repository::wallet::WalletMutationOutcome::NotFound) => {
return Ok(build_admin_wallet_refund_not_found_response())
}
Some(aether_data::repository::wallet::WalletMutationOutcome::Invalid(
detail,
)) => return Ok(build_admin_wallets_bad_request_response(detail)),
None => return Ok(build_admin_wallets_data_unavailable_response()),
};
let persisted = stored_refund_to_gateway(persisted);
let response = (
http::StatusCode::ACCEPTED,
Json(json!({
"refund": build_admin_wallet_refund_payload(&wallet, &owner, &persisted),
"gateway_refund": {
"id": result.gateway_refund_id,
"status": result.status,
},
})),
)
.into_response();
return Ok(attach_admin_audit_response(
response,
"admin_wallet_refund_pending",
"complete_wallet_refund",
"wallet_refund",
&refund_id,
));
}
if !result.is_succeeded() {
return Ok(build_admin_wallets_bad_request_response("上游退款未成功"));
}
payout_proof = merge_gateway_refund_proof(payout_proof, Some(&result));
// Persist the provider evidence before releasing the local refund reservation.
// If the local completion transaction fails after a successful gateway call,
// a retry can reuse the idempotent gateway identifier instead of issuing a
// second refund with no durable proof of the first one.
match state
.app()
.update_admin_wallet_refund_gateway(
aether_data::repository::wallet::UpdateAdminWalletRefundGatewayInput {
wallet_id: wallet_id.clone(),
refund_id: refund_id.clone(),
gateway_refund_id: result.gateway_refund_id.clone(),
payout_proof: payout_proof.clone(),
},
)
.await?
{
Some(aether_data::repository::wallet::WalletMutationOutcome::Applied(_)) => {}
Some(aether_data::repository::wallet::WalletMutationOutcome::NotFound) => {
return Ok(build_admin_wallet_refund_not_found_response())
}
Some(aether_data::repository::wallet::WalletMutationOutcome::Invalid(
detail,
)) => return Ok(build_admin_wallets_bad_request_response(detail)),
None => return Ok(build_admin_wallets_data_unavailable_response()),
}
}
Ok(None) => {
return Ok(build_admin_wallets_bad_request_response(
"该支付方式不支持官方直连退款,请使用线下完成",
))
}
Err(detail) => return Ok(build_admin_wallets_bad_request_response(detail)),
Err(detail) => {
warn!(
error = %detail,
refund_id = %refund_id,
"direct payment gateway refund failed"
);
return Ok(build_admin_wallets_bad_request_response(
"支付网关退款请求失败",
));
}
}
}
match state
@@ -183,6 +354,7 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response(
refund_id = %refund.id,
"failed to reverse referral rewards for completed refund"
);
return Ok(build_admin_wallets_data_unavailable_response());
}
}
let response = Json(json!({
@@ -204,7 +376,7 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response(
let detail = if detail == "refund status must be processing before completion" {
"只有 processing 状态的退款可以标记完成".to_string()
} else {
detail
"退款状态或参数无效".to_string()
};
Ok(build_admin_wallets_bad_request_response(detail))
}
@@ -213,3 +385,70 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response(
}
}
}
#[cfg(test)]
mod tests {
use super::{
gateway_refund_mode_allowed, is_safe_gateway_refund_id, merge_gateway_refund_proof,
};
use crate::handlers::shared::DirectGatewayRefundResult;
use serde_json::json;
#[test]
fn gateway_refund_merge_replaces_legacy_raw_payload() {
let existing = json!({
"channel": "manual",
"gateway_refund": {
"payload": {
"authorization": "Bearer legacy-secret",
"payer": {"openid": "openid-secret"}
}
}
});
let result = DirectGatewayRefundResult {
gateway_refund_id: "refund-1".to_string(),
status: "success".to_string(),
proof: json!({
"gateway": "wxpay",
"id": "refund-1",
"status": "success",
"order_no": "order-1",
"refund_no": "request-1",
"amount": 8.5,
"currency": "CNY",
"processed_at": "2026-08-27T12:00:00Z"
}),
};
let merged = merge_gateway_refund_proof(Some(existing), Some(&result))
.expect("gateway proof should be merged");
assert_eq!(merged["channel"], "manual");
assert_eq!(merged["gateway_refund"], result.proof);
let encoded = merged.to_string();
assert!(!encoded.contains("legacy-secret"));
assert!(!encoded.contains("openid-secret"));
assert!(!encoded.contains("payload"));
}
#[test]
fn manual_gateway_refund_ids_use_the_same_strict_identifier_policy() {
assert!(is_safe_gateway_refund_id("refund_123-ABC"));
for value in [
"Authorization: Bearer secret",
"https://internal.example/refund?token=secret",
"refund id",
"payer/openid",
] {
assert!(!is_safe_gateway_refund_id(value));
}
assert!(!is_safe_gateway_refund_id(&"a".repeat(129)));
}
#[test]
fn gateway_refunds_are_limited_to_original_channel_mode() {
assert!(gateway_refund_mode_allowed("original_channel"));
assert!(gateway_refund_mode_allowed(" Original_Channel "));
assert!(!gateway_refund_mode_allowed("offline_payout"));
assert!(!gateway_refund_mode_allowed(""));
}
}
@@ -10,6 +10,7 @@ use super::super::shared::{
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{attach_admin_audit_response, unix_secs_to_rfc3339};
use crate::GatewayError;
use aether_data::repository::wallet::stored_timestamp_unix_secs;
use axum::{
body::Body,
response::{IntoResponse, Response},
@@ -86,7 +87,7 @@ pub(in super::super) async fn build_admin_wallet_fail_refund_response(
transaction.link_id.as_deref(),
transaction.operator_id.as_deref(),
transaction.description.as_deref(),
unix_secs_to_rfc3339(transaction.created_at_unix_ms),
unix_secs_to_rfc3339(stored_timestamp_unix_secs(transaction.created_at_unix_ms)),
)
})
.unwrap_or(serde_json::Value::Null),
@@ -9,6 +9,7 @@ use super::super::shared::{
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{attach_admin_audit_response, unix_secs_to_rfc3339};
use crate::GatewayError;
use aether_data::repository::wallet::stored_timestamp_unix_secs;
use axum::{
body::Body,
response::{IntoResponse, Response},
@@ -71,7 +72,7 @@ pub(in super::super) async fn build_admin_wallet_process_refund_response(
transaction.link_id.as_deref(),
transaction.operator_id.as_deref(),
transaction.description.as_deref(),
unix_secs_to_rfc3339(transaction.created_at_unix_ms),
unix_secs_to_rfc3339(stored_timestamp_unix_secs(transaction.created_at_unix_ms)),
),
}))
.into_response();
@@ -10,6 +10,7 @@ use super::super::shared::{
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{attach_admin_audit_response, unix_secs_to_rfc3339};
use crate::GatewayError;
use aether_data::repository::wallet::stored_timestamp_unix_secs;
use axum::{
body::Body,
response::{IntoResponse, Response},
@@ -89,7 +90,7 @@ pub(in super::super) async fn build_admin_wallet_recharge_response(
payment_order.amount_usd,
payment_order.payment_method,
payment_order.status,
unix_secs_to_rfc3339(payment_order.created_at_unix_ms),
unix_secs_to_rfc3339(stored_timestamp_unix_secs(payment_order.created_at_unix_ms)),
payment_order
.credited_at_unix_secs
.and_then(unix_secs_to_rfc3339),
@@ -6,6 +6,7 @@ use super::super::shared::{
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{query_param_value, unix_secs_to_rfc3339};
use crate::GatewayError;
use aether_data::repository::wallet::stored_timestamp_unix_secs;
use axum::{
body::Body,
response::{IntoResponse, Response},
@@ -69,7 +70,10 @@ pub(in super::super) async fn build_admin_wallet_ledger_response(
"operator_name": entry.operator_name,
"operator_email": entry.operator_email,
"description": entry.description,
"created_at": entry.created_at_unix_ms.and_then(unix_secs_to_rfc3339),
"created_at": entry
.created_at_unix_ms
.map(stored_timestamp_unix_secs)
.and_then(unix_secs_to_rfc3339),
})
})
.collect::<Vec<_>>();
@@ -6,6 +6,7 @@ use super::super::shared::{
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{query_param_value, unix_secs_to_rfc3339};
use crate::GatewayError;
use aether_data::repository::wallet::stored_timestamp_unix_secs;
use axum::{
body::Body,
response::{IntoResponse, Response},
@@ -61,7 +62,10 @@ pub(in super::super) async fn build_admin_wallet_list_response(
"total_consumed": wallet.total_consumed,
"total_refunded": wallet.total_refunded,
"total_adjusted": wallet.total_adjusted,
"created_at": wallet.created_at_unix_ms.and_then(unix_secs_to_rfc3339),
"created_at": wallet
.created_at_unix_ms
.map(stored_timestamp_unix_secs)
.and_then(unix_secs_to_rfc3339),
"updated_at": wallet.updated_at_unix_secs.and_then(unix_secs_to_rfc3339),
});
enrich_admin_wallet_package_summary(
@@ -1,12 +1,13 @@
use super::super::shared::{
build_admin_wallets_bad_request_response, parse_admin_wallets_limit,
parse_admin_wallets_offset, parse_admin_wallets_owner_type_filter,
admin_wallet_payout_proof_projection, build_admin_wallets_bad_request_response,
parse_admin_wallets_limit, parse_admin_wallets_offset, parse_admin_wallets_owner_type_filter,
resolve_admin_wallet_owner_summary, wallet_owner_summary_from_fields,
ADMIN_WALLETS_API_KEY_REFUND_DETAIL,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{query_param_value, unix_secs_to_rfc3339};
use crate::GatewayError;
use aether_data::repository::wallet::stored_timestamp_unix_secs;
use axum::{
body::Body,
response::{IntoResponse, Response},
@@ -40,6 +41,7 @@ pub(in super::super) async fn build_admin_wallet_refund_requests_response(
.await?;
let mut items = Vec::with_capacity(refunds.len());
for refund in refunds {
let payout_proof = admin_wallet_payout_proof_projection(refund.payout_proof.as_ref());
let mut owner = wallet_owner_summary_from_fields(
refund.wallet_user_id.as_deref(),
refund.wallet_user_name.clone(),
@@ -75,11 +77,14 @@ pub(in super::super) async fn build_admin_wallet_refund_requests_response(
"gateway_refund_id": refund.gateway_refund_id,
"payout_method": refund.payout_method,
"payout_reference": refund.payout_reference,
"payout_proof": refund.payout_proof,
"payout_proof": payout_proof,
"requested_by": refund.requested_by,
"approved_by": refund.approved_by,
"processed_by": refund.processed_by,
"created_at": refund.created_at_unix_ms.and_then(unix_secs_to_rfc3339),
"created_at": refund
.created_at_unix_ms
.map(stored_timestamp_unix_secs)
.and_then(unix_secs_to_rfc3339),
"updated_at": refund.updated_at_unix_secs.and_then(unix_secs_to_rfc3339),
"processed_at": refund.processed_at_unix_secs.and_then(unix_secs_to_rfc3339),
"completed_at": refund.completed_at_unix_secs.and_then(unix_secs_to_rfc3339),
@@ -6,6 +6,7 @@ use super::super::shared::{
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
use crate::GatewayError;
use aether_data::repository::wallet::stored_timestamp_unix_secs;
use axum::{
body::Body,
response::{IntoResponse, Response},
@@ -81,7 +82,10 @@ pub(in super::super) async fn build_admin_wallet_transactions_response(
"operator_name": operator_name,
"operator_email": operator_email,
"description": transaction.description,
"created_at": transaction.created_at_unix_ms.and_then(unix_secs_to_rfc3339),
"created_at": transaction
.created_at_unix_ms
.map(stored_timestamp_unix_secs)
.and_then(unix_secs_to_rfc3339),
}));
}
@@ -51,11 +51,12 @@ pub(in super::super) fn normalize_admin_wallet_optional_text(
pub(in super::super) fn normalize_admin_wallet_payment_method(
value: String,
) -> Result<String, String> {
let normalized = value.trim();
if normalized.is_empty() {
return Err("payment_method 不能为空".to_string());
let normalized = aether_data::repository::wallet::canonicalize_payment_method(&value)
.map_err(|detail| format!("payment_method 无效: {detail}"))?;
if normalized.chars().count() > 30 {
return Err("payment_method 长度不能超过 30".to_string());
}
Ok(normalized.chars().take(30).collect())
Ok(normalized)
}
pub(in super::super) fn normalize_admin_wallet_balance_type(
@@ -2,7 +2,8 @@ use crate::handlers::admin::request::AdminAppState;
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
use crate::handlers::shared::round_to;
use crate::GatewayError;
use serde_json::json;
use aether_data::repository::wallet::stored_timestamp_unix_secs;
use serde_json::{json, Map, Value};
#[derive(Clone)]
pub(in super::super) struct AdminWalletOwnerSummary {
@@ -10,6 +11,10 @@ pub(in super::super) struct AdminWalletOwnerSummary {
pub(in super::super) owner_name: Option<String>,
}
fn api_key_display_prefix(api_key_id: &str) -> String {
api_key_id.chars().take(8).collect()
}
pub(in super::super) fn build_admin_wallet_payment_order_payload(
order_id: String,
order_no: String,
@@ -92,7 +97,7 @@ pub(in super::super) fn wallet_owner_summary_from_fields(
owner_type: "api_key",
owner_name: api_key_name
.filter(|value| !value.trim().is_empty())
.or_else(|| Some(format!("Key-{}", &api_key_id[..api_key_id.len().min(8)]))),
.or_else(|| Some(format!("Key-{}", api_key_display_prefix(api_key_id)))),
};
}
AdminWalletOwnerSummary {
@@ -121,7 +126,7 @@ pub(in super::super) async fn resolve_admin_wallet_owner_summary(
.find(|snapshot| snapshot.api_key_id == api_key_id)
.and_then(|snapshot| snapshot.api_key_name)
.filter(|value| !value.trim().is_empty())
.or_else(|| Some(format!("Key-{}", &api_key_id[..api_key_id.len().min(8)])));
.or_else(|| Some(format!("Key-{}", api_key_display_prefix(api_key_id))));
Ok(AdminWalletOwnerSummary {
owner_type: "api_key",
owner_name,
@@ -234,6 +239,104 @@ pub(in super::super) async fn enrich_admin_wallet_package_summary(
Ok(())
}
fn admin_refund_proof_identifier(value: Option<&Value>, max_bytes: usize) -> Option<String> {
let value = value?.as_str()?.trim();
if value.is_empty()
|| value.len() > max_bytes
|| !value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.'))
{
return None;
}
Some(value.to_string())
}
fn admin_gateway_refund_proof_projection(value: &Value) -> Option<Value> {
let source = value.as_object()?;
let mut projected = Map::new();
if let Some(gateway) = source
.get("gateway")
.and_then(Value::as_str)
.map(str::trim)
.map(str::to_ascii_lowercase)
.filter(|value| matches!(value.as_str(), "alipay" | "wxpay"))
{
projected.insert("gateway".to_string(), json!(gateway));
}
for (key, max_bytes) in [
("id", 128usize),
("order_no", 64usize),
("refund_no", 64usize),
] {
if let Some(value) = admin_refund_proof_identifier(source.get(key), max_bytes) {
projected.insert(key.to_string(), json!(value));
}
}
if let Some(status) = source
.get("status")
.and_then(Value::as_str)
.map(str::trim)
.map(str::to_ascii_lowercase)
.and_then(|value| match value.as_str() {
"success" | "succeeded" => Some("success"),
"pending" | "processing" => Some("processing"),
"failed" | "closed" | "abnormal" => Some("failed"),
_ => None,
})
{
projected.insert("status".to_string(), json!(status));
}
if let Some(amount) = source
.get("amount")
.and_then(Value::as_f64)
.filter(|value| value.is_finite() && *value > 0.0)
{
projected.insert("amount".to_string(), json!(amount));
}
if let Some(currency) = source
.get("currency")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| value.len() == 3 && value.bytes().all(|byte| byte.is_ascii_alphabetic()))
.map(str::to_ascii_uppercase)
{
projected.insert("currency".to_string(), json!(currency));
}
if let Some(processed_at) = source
.get("processed_at")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| chrono::DateTime::parse_from_rfc3339(value).is_ok())
{
projected.insert("processed_at".to_string(), json!(processed_at));
}
(!projected.is_empty()).then_some(Value::Object(projected))
}
pub(in super::super) fn admin_wallet_payout_proof_projection(
payout_proof: Option<&Value>,
) -> Option<Value> {
let source = payout_proof?.as_object()?;
let mut projected = source.clone();
if source.contains_key("gateway_refund") {
match source
.get("gateway_refund")
.and_then(admin_gateway_refund_proof_projection)
{
Some(gateway_refund) => {
projected.insert("gateway_refund".to_string(), gateway_refund);
}
None => {
projected.remove("gateway_refund");
}
}
}
Some(Value::Object(projected))
}
pub(in super::super) fn build_admin_wallet_refund_payload(
wallet: &aether_data::repository::wallet::StoredWalletSnapshot,
owner: &AdminWalletOwnerSummary,
@@ -258,13 +361,94 @@ pub(in super::super) fn build_admin_wallet_refund_payload(
"gateway_refund_id": refund.gateway_refund_id.clone(),
"payout_method": refund.payout_method.clone(),
"payout_reference": refund.payout_reference.clone(),
"payout_proof": refund.payout_proof.clone(),
"payout_proof": admin_wallet_payout_proof_projection(refund.payout_proof.as_ref()),
"requested_by": refund.requested_by.clone(),
"approved_by": refund.approved_by.clone(),
"processed_by": refund.processed_by.clone(),
"created_at": unix_secs_to_rfc3339(refund.created_at_unix_ms),
"created_at": unix_secs_to_rfc3339(stored_timestamp_unix_secs(refund.created_at_unix_ms)),
"updated_at": unix_secs_to_rfc3339(refund.updated_at_unix_secs),
"processed_at": refund.processed_at_unix_secs.and_then(unix_secs_to_rfc3339),
"completed_at": refund.completed_at_unix_secs.and_then(unix_secs_to_rfc3339),
})
}
#[cfg(test)]
mod tests {
use super::admin_wallet_payout_proof_projection;
use serde_json::json;
#[test]
fn admin_refund_proof_projection_removes_historical_gateway_payloads() {
let proof = json!({
"channel": "manual",
"gateway_refund": {
"gateway": "WXPAY",
"id": "refund-1",
"status": "SUCCESS",
"order_no": "order-1",
"refund_no": "request-1",
"amount": 8.5,
"currency": "cny",
"processed_at": "2026-08-27T12:00:00Z",
"payload": {
"authorization": "Bearer payment-secret",
"url": "https://internal.example/refund?token=secret",
"payer": {"openid": "openid-secret"},
"credential": "gateway-credential"
},
"message": "upstream secret message"
}
});
let projection = admin_wallet_payout_proof_projection(Some(&proof))
.expect("object payout proof should be projected");
assert_eq!(projection["channel"], "manual");
assert_eq!(projection["gateway_refund"]["gateway"], "wxpay");
assert_eq!(projection["gateway_refund"]["status"], "success");
assert_eq!(
projection["gateway_refund"]
.as_object()
.expect("gateway proof should be an object")
.len(),
8
);
let encoded = projection.to_string();
for sensitive in [
"payment-secret",
"?token=secret",
"openid-secret",
"gateway-credential",
"upstream secret message",
"payload",
"authorization",
"payer",
"openid",
"credential",
"message",
] {
assert!(!encoded.contains(sensitive));
}
}
#[test]
fn admin_refund_proof_projection_drops_mistyped_gateway_fields() {
let proof = json!({
"operator": "finance",
"gateway_refund": {
"gateway": {"credential": "secret"},
"id": ["refund-1"],
"status": "unknown-secret-status",
"order_no": "https://example.test/?token=secret",
"refund_no": 123,
"amount": "8.5",
"currency": "CNY?token=secret",
"processed_at": "Bearer secret"
}
});
let projection = admin_wallet_payout_proof_projection(Some(&proof))
.expect("manual payout proof should remain available");
assert_eq!(projection, json!({"operator": "finance"}));
assert!(!projection.to_string().contains("secret"));
}
}
@@ -1,13 +1,14 @@
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{
attach_admin_audit_response, query_param_value, unix_secs_to_rfc3339,
attach_admin_audit_response, mark_sensitive_admin_response_no_store, query_param_value,
unix_secs_to_rfc3339,
};
use crate::task_runtime::{
self, set_cancel_signal, TASK_KEY_PROVIDER_DELETE, TASK_KEY_PROVIDER_OAUTH_BATCH_IMPORT,
};
use crate::GatewayError;
use aether_data_contracts::repository::background_tasks::{
BackgroundTaskKind, BackgroundTaskListQuery, BackgroundTaskStatus,
BackgroundTaskKind, BackgroundTaskListQuery, BackgroundTaskStatus, StoredBackgroundTaskRun,
};
use axum::{
body::{Body, Bytes},
@@ -21,6 +22,30 @@ const DEFAULT_PAGE_SIZE: usize = 20;
const MAX_PAGE_SIZE: usize = 100;
const DEFAULT_EVENTS_PAGE_SIZE: usize = 50;
fn build_background_task_list_item(run: &StoredBackgroundTaskRun) -> serde_json::Value {
json!({
"id": run.id,
"task_key": run.task_key,
"kind": run.kind.as_database(),
"trigger": run.trigger,
"status": run.status.as_database(),
"attempt": run.attempt,
"max_attempts": run.max_attempts,
"owner_instance": run.owner_instance,
"progress_percent": run.progress_percent,
"progress_message": run.progress_message,
"has_payload": run.payload_json.is_some(),
"has_result": run.result_json.is_some(),
"has_error": run.error_message.is_some(),
"cancel_requested": run.cancel_requested,
"created_by": run.created_by,
"created_at": unix_secs_to_rfc3339(run.created_at_unix_secs),
"started_at": run.started_at_unix_secs.and_then(unix_secs_to_rfc3339),
"finished_at": run.finished_at_unix_secs.and_then(unix_secs_to_rfc3339),
"updated_at": unix_secs_to_rfc3339(run.updated_at_unix_secs),
})
}
pub(super) async fn maybe_build_local_admin_background_tasks_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -71,29 +96,7 @@ pub(super) async fn maybe_build_local_admin_background_tasks_response(
let items = response
.items
.iter()
.map(|run| {
json!({
"id": run.id,
"task_key": run.task_key,
"kind": run.kind.as_database(),
"trigger": run.trigger,
"status": run.status.as_database(),
"attempt": run.attempt,
"max_attempts": run.max_attempts,
"owner_instance": run.owner_instance,
"progress_percent": run.progress_percent,
"progress_message": run.progress_message,
"payload": run.payload_json,
"result": run.result_json,
"error_message": run.error_message,
"cancel_requested": run.cancel_requested,
"created_by": run.created_by,
"created_at": unix_secs_to_rfc3339(run.created_at_unix_secs),
"started_at": run.started_at_unix_secs.and_then(unix_secs_to_rfc3339),
"finished_at": run.finished_at_unix_secs.and_then(unix_secs_to_rfc3339),
"updated_at": unix_secs_to_rfc3339(run.updated_at_unix_secs),
})
})
.map(build_background_task_list_item)
.collect::<Vec<_>>();
let definitions = task_runtime::task_definitions()
.iter()
@@ -153,8 +156,9 @@ pub(super) async fn maybe_build_local_admin_background_tasks_response(
.into_response(),
));
};
return Ok(Some(attach_admin_audit_response(
Json(json!({
return Ok(Some(mark_sensitive_admin_response_no_store(
attach_admin_audit_response(
Json(json!({
"id": run.id,
"task_key": run.task_key,
"kind": run.kind.as_database(),
@@ -174,12 +178,13 @@ pub(super) async fn maybe_build_local_admin_background_tasks_response(
"started_at": run.started_at_unix_secs.and_then(unix_secs_to_rfc3339),
"finished_at": run.finished_at_unix_secs.and_then(unix_secs_to_rfc3339),
"updated_at": unix_secs_to_rfc3339(run.updated_at_unix_secs),
}))
.into_response(),
"admin_task_detail_viewed",
"view_task_detail",
"background_task",
run_id,
}))
.into_response(),
"admin_task_detail_viewed",
"view_task_detail",
"background_task",
run_id,
),
)));
}
Some("events") if request_context.method() == http::Method::GET => {
@@ -205,7 +210,7 @@ pub(super) async fn maybe_build_local_admin_background_tasks_response(
let events = state
.list_background_task_events(run_id, offset, page_size)
.await?;
return Ok(Some(
return Ok(Some(mark_sensitive_admin_response_no_store(
Json(json!({
"items": events.into_iter().map(|event| {
json!({
@@ -221,7 +226,7 @@ pub(super) async fn maybe_build_local_admin_background_tasks_response(
"page_size": page_size,
}))
.into_response(),
));
)));
}
Some("cancel") if request_context.method() == http::Method::POST => {
let Some(run_id) = nested_task_id_from_path(request_context.path(), "/cancel") else {
@@ -371,3 +376,55 @@ fn parse_json_payload(request_body: Option<&Bytes>) -> Result<serde_json::Value,
serde_json::from_slice::<serde_json::Value>(body)
.map_err(|err| GatewayError::Internal(format!("invalid json body: {err}")))
}
#[cfg(test)]
mod tests {
use super::build_background_task_list_item;
use aether_data_contracts::repository::background_tasks::{
BackgroundTaskKind, BackgroundTaskStatus, StoredBackgroundTaskRun,
};
use serde_json::json;
#[test]
fn task_list_item_exposes_only_safe_diagnostic_presence_flags() {
let run = StoredBackgroundTaskRun {
id: "run-1".to_string(),
task_key: "provider.oauth.import".to_string(),
kind: BackgroundTaskKind::OnDemand,
trigger: "manual".to_string(),
status: BackgroundTaskStatus::Failed,
attempt: 1,
max_attempts: 3,
owner_instance: Some("gateway-1".to_string()),
progress_percent: 100,
progress_message: Some("task failed".to_string()),
payload_json: Some(json!({"refresh_token": "secret-refresh-token"})),
result_json: Some(json!({"access_token": "secret-access-token"})),
error_message: Some("upstream error containing secret-api-key".to_string()),
cancel_requested: false,
created_by: Some("admin".to_string()),
created_at_unix_secs: 1,
started_at_unix_secs: Some(2),
finished_at_unix_secs: Some(3),
updated_at_unix_secs: 3,
};
let item = build_background_task_list_item(&run);
assert_eq!(item["status"], "failed");
assert_eq!(item["has_payload"], true);
assert_eq!(item["has_result"], true);
assert_eq!(item["has_error"], true);
assert!(item.get("payload").is_none());
assert!(item.get("result").is_none());
assert!(item.get("error_message").is_none());
let serialized = item.to_string();
for secret in [
"secret-refresh-token",
"secret-access-token",
"secret-api-key",
] {
assert!(!serialized.contains(secret));
}
}
}
@@ -173,6 +173,7 @@ pub(super) async fn maybe_build_local_admin_gemini_files_read_response(
let mappings = state
.list_gemini_file_mappings(
&aether_data::repository::gemini_file_mappings::GeminiFileMappingListQuery {
user_id: None,
include_expired: page.include_expired,
search: page.search.clone(),
offset: (page.page - 1).saturating_mul(page.page_size),
@@ -1,4 +1,11 @@
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::shared::{
find_multipart_boundary, find_multipart_boundary_after_crlf, parse_multipart_boundary,
MAX_MULTIPART_PARTS, MAX_MULTIPART_PART_HEADER_BYTES,
};
use aether_data_contracts::repository::gemini_file_mappings::{
GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS, GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS,
};
use axum::body::Bytes;
use base64::Engine as _;
@@ -21,7 +28,8 @@ pub(super) fn admin_gemini_files_parse_upload_request(
.map(str::trim)
.filter(|value| !value.is_empty())
.ok_or_else(|| "Content-Type 缺失".to_string())?;
let boundary = admin_gemini_files_multipart_boundary(content_type)?;
let boundary = parse_multipart_boundary(content_type)
.ok_or_else(|| "multipart boundary 缺失或无效".to_string())?;
let body = request_body
.filter(|body| !body.is_empty())
.ok_or_else(|| "上传文件不能为空".to_string())?;
@@ -35,47 +43,32 @@ pub(super) fn admin_gemini_files_parse_upload_request(
})
}
fn admin_gemini_files_multipart_boundary(content_type: &str) -> Result<String, String> {
let normalized = content_type.trim();
if !normalized
.to_ascii_lowercase()
.starts_with("multipart/form-data")
{
return Err("Content-Type 必须是 multipart/form-data".to_string());
}
for part in normalized.split(';').skip(1) {
let Some((key, value)) = part.trim().split_once('=') else {
continue;
};
if !key.trim().eq_ignore_ascii_case("boundary") {
continue;
}
let boundary = value.trim().trim_matches('"').trim();
if !boundary.is_empty() {
return Ok(boundary.to_string());
}
}
Err("multipart boundary 缺失".to_string())
}
fn admin_gemini_files_extract_file_part(
body: &[u8],
boundary: &str,
) -> Result<(String, String, Vec<u8>), String> {
let boundary_marker = format!("--{boundary}");
let next_boundary_marker = format!("\r\n--{boundary}");
let boundary_bytes = boundary_marker.as_bytes();
let next_boundary_bytes = next_boundary_marker.as_bytes();
let mut cursor = 0usize;
let mut part_count = 0usize;
let mut file_part = None;
while cursor < body.len() {
if !body[cursor..].starts_with(boundary_bytes) {
if find_multipart_boundary(&body[cursor..], boundary_bytes) != Some(0) {
return Err("multipart body 格式无效".to_string());
}
cursor += boundary_bytes.len();
if body[cursor..].starts_with(b"--") {
let closing_suffix = body.get(cursor + 2..).unwrap_or_default();
if !(closing_suffix.is_empty() || closing_suffix.starts_with(b"\r\n")) {
return Err("multipart 结束边界格式无效".to_string());
}
break;
}
part_count = part_count.saturating_add(1);
if part_count > MAX_MULTIPART_PARTS {
return Err("multipart part 数量超过上限".to_string());
}
if !body[cursor..].starts_with(b"\r\n") {
return Err("multipart body 缺少头部分隔符".to_string());
}
@@ -85,34 +78,60 @@ fn admin_gemini_files_extract_file_part(
return Err("multipart part 缺少头部".to_string());
};
let headers_end = cursor + headers_end_rel;
if headers_end_rel > MAX_MULTIPART_PART_HEADER_BYTES {
return Err("multipart part 头部超过大小上限".to_string());
}
let headers_text = std::str::from_utf8(&body[cursor..headers_end])
.map_err(|_| "multipart part 头部编码无效".to_string())?;
cursor = headers_end + 4;
let Some(next_boundary_rel) =
admin_gemini_files_find_subslice(&body[cursor..], next_boundary_bytes)
find_multipart_boundary_after_crlf(&body[cursor..], boundary_bytes)
else {
return Err("multipart body 缺少结束边界".to_string());
};
let content_end = cursor + next_boundary_rel;
let content = &body[cursor..content_end];
cursor = content_end + 2;
// The CRLF immediately before the delimiter belongs to the
// multipart framing, not to the uploaded file bytes.
let content = body[cursor..content_end]
.strip_suffix(b"\r\n")
.unwrap_or(&body[cursor..content_end]);
cursor = content_end;
let Some((field_name, file_name, mime_type)) =
admin_gemini_files_parse_part_headers(headers_text)
else {
continue;
return Err("multipart part 头部无效".to_string());
};
if field_name != "file" {
continue;
}
return Ok((
if file_part.is_some() {
return Err("multipart body 包含多个 file 字段".to_string());
}
file_part = Some((
file_name.unwrap_or_else(|| "uploaded-file".to_string()),
mime_type.unwrap_or_else(|| "application/octet-stream".to_string()),
content.to_vec(),
));
}
Err("multipart body 中缺少 file 字段".to_string())
let (display_name, mime_type, content) =
file_part.ok_or_else(|| "multipart body 中缺少 file 字段".to_string())?;
if display_name
.chars()
.nth(GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS)
.is_some()
{
return Err("上传文件名超过长度上限".to_string());
}
if mime_type
.chars()
.nth(GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS)
.is_some()
{
return Err("上传文件 Content-Type 超过长度上限".to_string());
}
Ok((display_name, mime_type, content))
}
fn admin_gemini_files_parse_part_headers(
@@ -121,27 +140,29 @@ fn admin_gemini_files_parse_part_headers(
let mut field_name = None;
let mut file_name = None;
let mut mime_type = None;
let mut disposition_seen = false;
let mut content_type_seen = false;
for line in headers_text.split("\r\n") {
let Some((header_name, header_value)) = line.split_once(':') else {
continue;
};
let (header_name, header_value) = line.split_once(':')?;
let header_name = header_name.trim();
let header_value = header_value.trim();
if header_name.eq_ignore_ascii_case("content-disposition") {
for part in header_value.split(';').skip(1) {
let Some((key, value)) = part.trim().split_once('=') else {
continue;
};
let key = key.trim();
let value = value.trim().trim_matches('"').trim();
if key.eq_ignore_ascii_case("name") && !value.is_empty() {
field_name = Some(value.to_string());
} else if key.eq_ignore_ascii_case("filename") && !value.is_empty() {
file_name = Some(value.to_string());
}
if disposition_seen {
return None;
}
} else if header_name.eq_ignore_ascii_case("content-type") && !header_value.is_empty() {
disposition_seen = true;
let (name, filename) = admin_gemini_files_parse_content_disposition(header_value)?;
field_name = Some(name);
file_name = filename;
} else if header_name.eq_ignore_ascii_case("content-type") {
if content_type_seen
|| header_value.is_empty()
|| header_value.chars().any(char::is_control)
{
return None;
}
content_type_seen = true;
mime_type = Some(header_value.to_string());
}
}
@@ -149,6 +170,150 @@ fn admin_gemini_files_parse_part_headers(
field_name.map(|field_name| (field_name, file_name, mime_type))
}
fn admin_gemini_files_parse_content_disposition(value: &str) -> Option<(String, Option<String>)> {
let segments = admin_gemini_files_split_header_parameters(value)?;
if !segments.first()?.trim().eq_ignore_ascii_case("form-data") {
return None;
}
let mut seen_keys = Vec::new();
let mut name = None;
let mut filename = None;
for segment in segments.into_iter().skip(1) {
let segment = segment.trim();
if segment.is_empty() {
return None;
}
let (raw_key, raw_value) = segment.split_once('=')?;
let key = raw_key.trim();
if key.is_empty()
|| !key
.as_bytes()
.iter()
.copied()
.all(admin_gemini_files_is_token_byte)
{
return None;
}
if seen_keys
.iter()
.any(|seen: &String| seen.eq_ignore_ascii_case(key))
{
return None;
}
seen_keys.push(key.to_ascii_lowercase());
let parsed_value = admin_gemini_files_parse_parameter_value(raw_value.trim())?;
if key.eq_ignore_ascii_case("name") {
if parsed_value.is_empty() {
return None;
}
name = Some(parsed_value);
} else if key.eq_ignore_ascii_case("filename") {
if !parsed_value.is_empty() {
filename = Some(parsed_value);
}
}
}
Some((name?, filename))
}
fn admin_gemini_files_split_header_parameters(value: &str) -> Option<Vec<&str>> {
let mut segments = Vec::new();
let mut start = 0usize;
let mut in_quotes = false;
let mut escaped = false;
for (index, byte) in value.as_bytes().iter().copied().enumerate() {
if in_quotes {
if escaped {
escaped = false;
} else if byte == b'\\' {
escaped = true;
} else if byte == b'"' {
in_quotes = false;
}
} else if byte == b'"' {
in_quotes = true;
} else if byte == b';' {
segments.push(&value[start..index]);
start = index + 1;
}
}
if in_quotes || escaped {
return None;
}
segments.push(&value[start..]);
Some(segments)
}
fn admin_gemini_files_parse_parameter_value(value: &str) -> Option<String> {
if value.is_empty() {
return None;
}
if value.starts_with('"') {
if value.len() < 2 || !value.ends_with('"') {
return None;
}
let inner = &value[1..value.len() - 1];
let mut parsed = String::with_capacity(inner.len());
let mut escaped = false;
for character in inner.chars() {
if escaped {
if character.is_control() {
return None;
}
parsed.push(character);
escaped = false;
} else if character == '\\' {
escaped = true;
} else {
if character == '"' || character.is_control() {
return None;
}
parsed.push(character);
}
}
if escaped {
return None;
}
return Some(parsed);
}
value
.as_bytes()
.iter()
.copied()
.all(admin_gemini_files_is_token_byte)
.then(|| value.to_string())
}
fn admin_gemini_files_is_token_byte(byte: u8) -> bool {
matches!(
byte,
b'0'..=b'9'
| b'A'..=b'Z'
| b'a'..=b'z'
| b'!'
| b'#'
| b'$'
| b'%'
| b'&'
| b'\''
| b'*'
| b'+'
| b'-'
| b'.'
| b'^'
| b'_'
| b'`'
| b'|'
| b'~'
)
}
fn admin_gemini_files_find_subslice(haystack: &[u8], needle: &[u8]) -> Option<usize> {
if haystack.is_empty() || needle.is_empty() || haystack.len() < needle.len() {
return None;
@@ -157,3 +322,224 @@ fn admin_gemini_files_find_subslice(haystack: &[u8], needle: &[u8]) -> Option<us
.windows(needle.len())
.position(|window| window == needle)
}
#[cfg(test)]
mod tests {
use aether_data_contracts::repository::gemini_file_mappings::{
GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS, GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS,
};
use super::{
admin_gemini_files_extract_file_part, MAX_MULTIPART_PARTS, MAX_MULTIPART_PART_HEADER_BYTES,
};
#[test]
fn multipart_upload_rejects_metadata_beyond_storage_limits() {
let boundary = "metadata-limits";
let oversized_filename = "f".repeat(GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS + 1);
let oversized_mime_type = "m".repeat(GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS + 1);
for (filename, mime_type, expected) in [
(
oversized_filename.as_str(),
"application/octet-stream",
"上传文件名超过长度上限",
),
(
"payload.bin",
oversized_mime_type.as_str(),
"上传文件 Content-Type 超过长度上限",
),
] {
let body = format!(
"--{boundary}\r\nContent-Disposition: form-data; name=\"file\"; filename=\"{filename}\"\r\nContent-Type: {mime_type}\r\n\r\nfile-body\r\n--{boundary}--\r\n"
);
assert_eq!(
admin_gemini_files_extract_file_part(body.as_bytes(), boundary),
Err(expected.to_string())
);
}
}
#[test]
fn multipart_upload_rejects_excessive_part_count() {
let boundary = "bounded-parts";
let mut body = Vec::new();
for index in 0..(MAX_MULTIPART_PARTS + 1) {
body.extend_from_slice(
format!(
"--{boundary}\r\nContent-Disposition: form-data; name=\"field-{index}\"\r\n\r\nvalue\r\n"
)
.as_bytes(),
);
}
body.extend_from_slice(format!("--{boundary}--\r\n").as_bytes());
assert_eq!(
admin_gemini_files_extract_file_part(&body, boundary),
Err("multipart part 数量超过上限".to_string())
);
}
#[test]
fn multipart_upload_rejects_oversized_part_headers() {
let boundary = "bounded-header";
let mut body =
format!("--{boundary}\r\nContent-Disposition: form-data; name=\"file\"; x=\"")
.into_bytes();
body.extend(std::iter::repeat_n(b'x', MAX_MULTIPART_PART_HEADER_BYTES));
body.extend_from_slice(format!("\"\r\n\r\nfile-body\r\n--{boundary}--\r\n").as_bytes());
assert_eq!(
admin_gemini_files_extract_file_part(&body, boundary),
Err("multipart part 头部超过大小上限".to_string())
);
}
#[test]
fn multipart_upload_preserves_boundary_like_payload() {
let boundary = "payload-boundary";
let body = format!(
concat!(
"--{boundary}\r\n",
"Content-Disposition: form-data; name=\"file\"; filename=\"payload.bin\"\r\n",
"Content-Type: application/octet-stream\r\n\r\n",
"prefix\r\n--{boundary}X\r\nsuffix--{boundary}\r\n",
"--{boundary}--\r\n"
),
boundary = boundary,
);
let (_, mime_type, content) =
admin_gemini_files_extract_file_part(body.as_bytes(), boundary).expect("file part");
assert_eq!(mime_type, "application/octet-stream");
assert_eq!(
content,
format!("prefix\r\n--{boundary}X\r\nsuffix--{boundary}").into_bytes()
);
}
#[test]
fn multipart_upload_rejects_invalid_suffix_without_closing_boundary() {
let boundary = "invalid-suffix";
let body = format!(
concat!(
"--{boundary}\r\n",
"Content-Disposition: form-data; name=\"file\"\r\n\r\n",
"file-body\r\n--{boundary}X\r\n"
),
boundary = boundary,
);
assert_eq!(
admin_gemini_files_extract_file_part(body.as_bytes(), boundary),
Err("multipart body 缺少结束边界".to_string())
);
}
#[test]
fn multipart_upload_rejects_garbage_after_closing_boundary() {
let boundary = "closing-suffix";
let body = format!(
concat!(
"--{boundary}\r\n",
"Content-Disposition: form-data; name=\"file\"\r\n\r\n",
"file-body\r\n",
"--{boundary}--junk"
),
boundary = boundary,
);
assert!(admin_gemini_files_extract_file_part(body.as_bytes(), boundary).is_err());
}
#[test]
fn multipart_upload_validates_parts_after_file_before_returning() {
let boundary = "trailing-invalid";
let body = format!(
concat!(
"--{boundary}\r\n",
"Content-Disposition: form-data; name=\"file\"\r\n\r\n",
"file-body\r\n",
"--{boundary}\r\n",
"Content-Disposition: form-data; name=\"metadata\"\r\n\r\n",
"metadata\r\n--{boundary}X\r\n"
),
boundary = boundary,
);
assert_eq!(
admin_gemini_files_extract_file_part(body.as_bytes(), boundary),
Err("multipart body 缺少结束边界".to_string())
);
}
#[test]
fn multipart_upload_parses_quoted_parameters_without_filename_confusion() {
let boundary = "quoted-parameters";
let body = format!(
concat!(
"--{boundary}\r\n",
"Content-Disposition: form-data; filename=\"prefix; name=\\\"decoy\\\".bin\"; name=\"file\"\r\n",
"Content-Type: application/octet-stream\r\n\r\n",
"file-body\r\n",
"--{boundary}--\r\n"
),
boundary = boundary,
);
let (filename, _, content) =
admin_gemini_files_extract_file_part(body.as_bytes(), boundary).expect("file part");
assert_eq!(filename, "prefix; name=\"decoy\".bin");
assert_eq!(content, b"file-body");
}
#[test]
fn multipart_upload_rejects_filename_embedded_name_and_duplicate_parameters() {
let boundary = "ambiguous-parameters";
for content_disposition in [
"form-data; filename=\"name=\\\"file\\\"\"",
"form-data; name=\"file\"; name=\"metadata\"",
"form-data; name=\"file\"; filename=\"unterminated",
] {
let body = format!(
concat!(
"--{boundary}\r\n",
"Content-Disposition: {content_disposition}\r\n\r\n",
"file-body\r\n",
"--{boundary}--\r\n"
),
boundary = boundary,
content_disposition = content_disposition,
);
assert_eq!(
admin_gemini_files_extract_file_part(body.as_bytes(), boundary),
Err("multipart part 头部无效".to_string()),
"header should be rejected: {content_disposition}"
);
}
}
#[test]
fn multipart_upload_rejects_duplicate_file_parts() {
let boundary = "duplicate-file";
let body = format!(
concat!(
"--{boundary}\r\n",
"Content-Disposition: form-data; name=\"file\"\r\n\r\n",
"first\r\n",
"--{boundary}\r\n",
"Content-Disposition: form-data; name=\"file\"\r\n\r\n",
"second\r\n",
"--{boundary}--\r\n"
),
boundary = boundary,
);
assert_eq!(
admin_gemini_files_extract_file_part(body.as_bytes(), boundary),
Err("multipart body 包含多个 file 字段".to_string())
);
}
}
@@ -1,16 +1,24 @@
use super::super::admin_gemini_files_key_capable;
use super::request::AdminGeminiFilesUploadRequest;
use crate::execution_runtime::transport::{
apply_upstream_response_body_limit, decode_base64_body_with_limit,
json_value_fits_serialized_limit,
};
use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError;
use aether_contracts::{ExecutionPlan, ExecutionResult, RequestBody};
use aether_data_contracts::repository::gemini_file_mappings::{
GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS, GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
};
use axum::http;
use base64::Engine;
use serde_json::json;
use std::collections::{BTreeMap, BTreeSet};
const MAX_GEMINI_UPLOAD_RESPONSE_JSON_BYTES: usize = 8 * 1024 * 1024;
#[derive(Debug)]
struct AdminGeminiFilesUploadExecutionSuccess {
file_name: String,
@@ -138,7 +146,7 @@ async fn admin_gemini_files_upload_single_key(
let transport = state
.read_provider_transport_snapshot(&key.provider_id, &endpoint.id, &key.id)
.await
.map_err(|err| format!("{err:?}"))?
.map_err(|_| "无法读取 Key 传输配置".to_string())?
.ok_or_else(|| "无法读取 Key 传输配置".to_string())?;
if !state.supports_local_gemini_transport_with_network(&transport, "gemini:generate_content") {
return Err("Key 传输配置不支持 Gemini Files 上传".to_string());
@@ -188,7 +196,7 @@ async fn admin_gemini_files_upload_single_key(
.build_gemini_files_passthrough_url(&transport.endpoint.base_url, upload_path, upload_query)
.ok_or_else(|| "无法构建 Gemini Files 上传地址".to_string())?;
let plan = ExecutionPlan {
let mut plan = ExecutionPlan {
request_id: format!("{trace_id}:admin-gemini-upload:{}", key.id),
candidate_id: None,
provider_name: Some(transport.provider.name.clone()),
@@ -215,10 +223,11 @@ async fn admin_gemini_files_upload_single_key(
transport_profile: state.resolve_transport_profile(&transport),
timeouts: state.resolve_transport_execution_timeouts(&transport),
};
apply_upstream_response_body_limit(&mut plan, MAX_GEMINI_UPLOAD_RESPONSE_JSON_BYTES);
let result = admin_gemini_files_execute_upload_plan(state, trace_id, &plan)
.await
.map_err(|error| format!("{error:?}"))?;
.map_err(|_| "Gemini Files 上传执行失败".to_string())?;
if result.status_code >= 400 {
return Err(admin_gemini_files_execution_error_message(&result));
}
@@ -241,7 +250,7 @@ async fn admin_gemini_files_upload_single_key(
.or(Some(upload.mime_type.as_str())),
)
.await
.map_err(|err| format!("上传成功但本地映射写入失败: {err:?}"))?;
.map_err(|_| "上传成功但本地映射写入失败".to_string())?;
Ok(success)
}
@@ -261,7 +270,8 @@ fn admin_gemini_files_execution_json_body(result: &ExecutionResult) -> Option<se
.as_ref()
.and_then(|body| body.json_body.as_ref())
{
return Some(body_json.clone());
return json_value_fits_serialized_limit(body_json, MAX_GEMINI_UPLOAD_RESPONSE_JSON_BYTES)
.then(|| body_json.clone());
}
let content_type = result
.headers
@@ -278,9 +288,12 @@ fn admin_gemini_files_execution_json_body(result: &ExecutionResult) -> Option<se
.body
.as_ref()
.and_then(|body| body.body_bytes_b64.as_deref())?;
let decoded = base64::engine::general_purpose::STANDARD
.decode(body_bytes_b64)
.ok()?;
let decoded = decode_base64_body_with_limit(
body_bytes_b64,
crate::headers::max_internal_buffered_body_bytes()
.min(MAX_GEMINI_UPLOAD_RESPONSE_JSON_BYTES),
)
.ok()?;
serde_json::from_slice(&decoded).ok()
}
@@ -296,13 +309,20 @@ fn admin_gemini_files_upload_success_from_body(
.get("name")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())?;
.filter(|value| !value.is_empty())
.filter(|value| aether_usage_runtime::normalize_gemini_file_name(value).is_some())?;
let display_name = file_object
.get("displayName")
.or_else(|| file_object.get("display_name"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.filter(|value| {
value
.chars()
.nth(GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS)
.is_none()
})
.map(ToOwned::to_owned)
.or_else(|| Some(upload.display_name.clone()));
let mime_type = file_object
@@ -311,6 +331,12 @@ fn admin_gemini_files_upload_success_from_body(
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.filter(|value| {
value
.chars()
.nth(GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS)
.is_none()
})
.map(ToOwned::to_owned)
.or_else(|| Some(upload.mime_type.clone()));
Some(AdminGeminiFilesUploadExecutionSuccess {
@@ -330,7 +356,7 @@ fn admin_gemini_files_execution_error_message(result: &ExecutionResult) -> Strin
.map(str::trim)
.filter(|value| !value.is_empty())
{
return message.to_string();
return bound_gemini_files_error_message(message);
}
if let Some(message) = body_json
.get("message")
@@ -338,7 +364,7 @@ fn admin_gemini_files_execution_error_message(result: &ExecutionResult) -> Strin
.map(str::trim)
.filter(|value| !value.is_empty())
{
return message.to_string();
return bound_gemini_files_error_message(message);
}
}
if let Some(error) = result
@@ -347,7 +373,112 @@ fn admin_gemini_files_execution_error_message(result: &ExecutionResult) -> Strin
.map(|error| error.message.trim())
.filter(|value| !value.is_empty())
{
return error.to_string();
return bound_gemini_files_error_message(error);
}
format!("上传失败,状态码 {}", result.status_code)
}
fn bound_gemini_files_error_message(value: &str) -> String {
let value = value.trim();
let end = value.floor_char_boundary(value.len().min(crate::MAX_ERROR_BODY_BYTES));
value[..end].to_string()
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use aether_contracts::{ExecutionResult, ResponseBody};
use aether_data_contracts::repository::gemini_file_mappings::{
GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS, GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS,
GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS,
};
use super::super::request::AdminGeminiFilesUploadRequest;
use super::{
admin_gemini_files_execution_error_message, admin_gemini_files_execution_json_body,
admin_gemini_files_upload_success_from_body,
};
fn sample_upload() -> AdminGeminiFilesUploadRequest {
AdminGeminiFilesUploadRequest {
display_name: "fallback.bin".to_string(),
mime_type: "application/octet-stream".to_string(),
body_bytes: vec![1],
body_bytes_b64: "AQ==".to_string(),
}
}
#[test]
fn gemini_upload_result_rejects_file_name_beyond_storage_limit() {
let body = serde_json::json!({
"file": {"name": "n".repeat(GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS + 1)}
});
assert!(admin_gemini_files_upload_success_from_body(&body, &sample_upload()).is_none());
}
#[test]
fn gemini_upload_result_ignores_oversized_optional_metadata() {
let body = serde_json::json!({
"file": {
"name": "files/safe",
"displayName": "d".repeat(GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS + 1),
"mimeType": "m".repeat(GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS + 1),
}
});
let success = admin_gemini_files_upload_success_from_body(&body, &sample_upload())
.expect("valid file name should remain usable");
assert_eq!(success.display_name.as_deref(), Some("fallback.bin"));
assert_eq!(
success.mime_type.as_deref(),
Some("application/octet-stream")
);
}
#[test]
fn gemini_upload_result_rejects_oversized_base64_before_decode() {
let encoded_limit =
crate::execution_runtime::transport::maximum_base64_len_for_decoded_limit(
super::MAX_GEMINI_UPLOAD_RESPONSE_JSON_BYTES,
);
let result = ExecutionResult {
request_id: "gemini-upload-oversized".to_string(),
candidate_id: None,
status_code: 200,
headers: BTreeMap::from([("content-type".to_string(), "application/json".to_string())]),
response_observation: None,
body: Some(ResponseBody {
json_body: None,
body_bytes_b64: Some("A".repeat(encoded_limit + 1)),
}),
telemetry: None,
error: None,
};
assert!(admin_gemini_files_execution_json_body(&result).is_none());
}
#[test]
fn gemini_upload_error_message_is_bounded_without_splitting_utf8() {
let message = format!("{}界", "x".repeat(crate::MAX_ERROR_BODY_BYTES));
let result = ExecutionResult {
request_id: "gemini-upload-oversized-error".to_string(),
candidate_id: None,
status_code: 500,
headers: BTreeMap::new(),
response_observation: None,
body: Some(ResponseBody {
json_body: Some(serde_json::json!({"error": {"message": message}})),
body_bytes_b64: None,
}),
telemetry: None,
error: None,
};
let detail = admin_gemini_files_execution_error_message(&result);
assert_eq!(detail.len(), crate::MAX_ERROR_BODY_BYTES);
assert!(detail.bytes().all(|byte| byte == b'x'));
}
}
@@ -22,6 +22,20 @@ pub(super) fn admin_video_task_status_name(status: VideoTaskStatus) -> &'static
}
}
pub(super) fn admin_video_task_error_projection(task: &StoredVideoTask) -> Option<String> {
if task.error_message.as_deref().is_none_or(str::is_empty)
&& task.error_code.as_deref().is_none_or(str::is_empty)
{
return None;
}
task.error_code
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.or_else(|| Some("provider_error".to_string()))
}
pub(super) fn admin_video_task_timestamp(unix_secs: Option<u64>) -> Option<String> {
unix_secs.and_then(|value| {
chrono::DateTime::<Utc>::from_timestamp(value as i64, 0)
@@ -91,7 +105,7 @@ pub(super) fn build_admin_video_task_list_item(
"aspect_ratio": task.aspect_ratio,
"video_url": task.video_url,
"error_code": task.error_code,
"error_message": task.error_message,
"error_message": admin_video_task_error_projection(task),
"poll_count": task.poll_count,
"max_poll_count": task.max_poll_count,
"created_at": admin_video_task_timestamp(Some(task.created_at_unix_ms)),
@@ -127,3 +141,65 @@ pub(super) fn current_admin_video_task_unix_secs() -> u64 {
.unwrap_or_default()
.as_secs()
}
#[cfg(test)]
mod tests {
use super::{admin_video_task_error_projection, build_admin_video_task_list_item};
use aether_data_contracts::repository::video_tasks::{StoredVideoTask, VideoTaskStatus};
use std::collections::BTreeMap;
fn failed_task() -> StoredVideoTask {
StoredVideoTask::new(
"task-1".to_string(),
None,
"request-1".to_string(),
Some("user-1".to_string()),
None,
Some("alice".to_string()),
None,
None,
Some("provider-1".to_string()),
Some("endpoint-1".to_string()),
Some("key-1".to_string()),
Some("openai:video".to_string()),
Some("openai:video".to_string()),
false,
Some("video-model".to_string()),
None,
None,
None,
None,
None,
None,
VideoTaskStatus::Failed,
100,
None,
0,
10,
None,
1,
10,
1,
None,
Some(2),
2,
Some("authentication_error".to_string()),
Some("Authorization: Bearer live-secret at https://api.example?key=secret".to_string()),
None,
None,
)
.expect("task")
}
#[test]
fn video_task_payload_does_not_return_historical_raw_error_text() {
let task = failed_task();
assert_eq!(
admin_video_task_error_projection(&task).as_deref(),
Some("authentication_error")
);
let payload = build_admin_video_task_list_item(&task, &BTreeMap::new());
assert_eq!(payload["error_message"], "authentication_error");
assert!(!payload.to_string().contains("live-secret"));
}
}
@@ -15,9 +15,10 @@ use axum::{
use serde_json::json;
use super::builders::{
admin_video_task_detail_id_from_path, admin_video_task_nested_id_from_path,
admin_video_task_status_name, admin_video_task_timestamp, build_admin_video_task_list_item,
build_admin_video_task_provider_names, current_admin_video_task_unix_secs,
admin_video_task_detail_id_from_path, admin_video_task_error_projection,
admin_video_task_nested_id_from_path, admin_video_task_status_name, admin_video_task_timestamp,
build_admin_video_task_list_item, build_admin_video_task_provider_names,
current_admin_video_task_unix_secs,
};
pub(super) async fn maybe_build_local_admin_video_tasks_response(
@@ -259,7 +260,10 @@ pub(super) async fn maybe_build_local_admin_video_tasks_response(
payload.insert("stored_video_path".to_string(), serde_json::Value::Null);
payload.insert("storage_provider".to_string(), serde_json::Value::Null);
payload.insert("error_code".to_string(), json!(task.error_code));
payload.insert("error_message".to_string(), json!(task.error_message));
payload.insert(
"error_message".to_string(),
json!(admin_video_task_error_projection(&task)),
);
payload.insert("retry_count".to_string(), json!(task.retry_count));
payload.insert("max_retries".to_string(), serde_json::Value::Null);
payload.insert(
@@ -41,6 +41,7 @@ pub(crate) use self::provider::oauth::runtime::{
refresh_provider_oauth_account_state_after_update,
};
pub(crate) use self::provider::ops::providers::actions::admin_provider_ops_local_action_response;
pub(crate) use self::provider::ops::providers::admin_provider_ops_credential_snapshot;
pub(crate) use self::provider::ops::providers::store_admin_provider_ops_balance_cache;
pub(crate) use self::provider::pool::config::admin_provider_pool_config;
pub(crate) use self::provider::pool_admin::maybe_build_local_admin_pool_response;
@@ -53,7 +54,7 @@ pub(crate) use self::provider::{
};
pub(crate) use self::request::{
AdminAppState, AdminGatewayProviderTransportSnapshot, AdminLocalOAuthRefreshError,
AdminRequestContext, AdminRouteRequest, AdminRouteResponse, AdminRouteResult,
AdminRequestContext, AdminRouteRequest, AdminRouteResponse, AdminRouteResult, SystemExportMode,
};
pub(crate) use self::routes::maybe_build_local_admin_response;
#[cfg(test)]
@@ -61,3 +62,7 @@ pub(crate) use self::system::{
clear_proxy_node_references_with_cache_failure_for_tests,
override_proxy_connectivity_probe_url_for_tests,
};
pub(crate) use self::system::{
execute_admin_system_import_exclusively, release_admin_system_import_lease,
try_acquire_admin_system_import_lease, AdminSystemImportLockError,
};
@@ -10,6 +10,7 @@ use axum::http;
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
use std::collections::BTreeMap;
use std::net::{IpAddr, SocketAddr};
use std::time::Duration;
use tracing::warn;
@@ -19,10 +20,18 @@ const ADMIN_EXTERNAL_MODELS_CACHE_VERSION: u8 = 2;
const ADMIN_EXTERNAL_MODELS_CACHE_TTL_SECS: u64 = 15 * 60;
const ADMIN_EXTERNAL_MODELS_SOURCE_URL_ENV: &str = "AETHER_GATEWAY_EXTERNAL_MODELS_URL";
const ADMIN_EXTERNAL_MODELS_SOURCE_URL_DEFAULT: &str = "https://models.dev/api.json";
const ADMIN_EXTERNAL_MODELS_OFFICIAL_HOST: &str = "models.dev";
const ADMIN_EXTERNAL_MODELS_OFFICIAL_PATH: &str = "/api.json";
pub(in crate::handlers::admin) const ADMIN_EXTERNAL_MODELS_PROXY_NODE_CONFIG_KEY: &str =
"external_models_proxy_node_id";
const ADMIN_EXTERNAL_MODELS_CONNECT_TIMEOUT_MS: u64 = 10_000;
const ADMIN_EXTERNAL_MODELS_TOTAL_TIMEOUT_MS: u64 = 300_000;
const ADMIN_EXTERNAL_MODELS_TOTAL_TIMEOUT_MS: u64 = 30_000;
const ADMIN_EXTERNAL_MODELS_RESPONSE_LIMIT_BYTES: usize = 8 * 1024 * 1024;
// Keep the cache envelope bounded independently of the upstream body limit.
// Normalization adds a small amount of metadata, while a corrupted/shared
// runtime KV value must never be allowed to drive an unbounded serde
// allocation during cache reads.
const ADMIN_EXTERNAL_MODELS_CACHE_MAX_BYTES: usize = 16 * 1024 * 1024;
pub(crate) const ADMIN_EXTERNAL_MODELS_CONFIG_MUTATION_LOCK_KEY: &str =
"admin:external_models_proxy_node_config:mutation";
const ADMIN_EXTERNAL_MODELS_CONFIG_MUTATION_LOCK_TTL: Duration = Duration::from_secs(10 * 60);
@@ -34,6 +43,13 @@ struct AdminExternalModelsCacheEnvelope {
payload: Value,
}
#[derive(Debug)]
struct ResolvedAdminExternalModelsSource {
url: url::Url,
host: String,
addresses: Vec<SocketAddr>,
}
#[cfg(test)]
pub(crate) struct AdminExternalModelsSourceUrlEnvGuard {
previous: Option<String>,
@@ -76,13 +92,139 @@ fn admin_external_models_source_url() -> String {
.unwrap_or_else(|| ADMIN_EXTERNAL_MODELS_SOURCE_URL_DEFAULT.to_string())
}
fn parse_admin_external_models_source_url(
raw_url: &str,
allow_insecure_test_target: bool,
) -> Result<(url::Url, String, u16), GatewayError> {
let url = url::Url::parse(raw_url)
.map_err(|_| GatewayError::Internal("external models source URL is invalid".to_string()))?;
let allowed_scheme =
url.scheme() == "https" || (allow_insecure_test_target && url.scheme() == "http");
if !allowed_scheme
|| !url.username().is_empty()
|| url.password().is_some()
|| url.query().is_some()
|| url.fragment().is_some()
{
return Err(GatewayError::Internal(
"external models source must be an HTTPS URL without credentials, query, or fragment"
.to_string(),
));
}
let host = url.host_str().map(ToOwned::to_owned).ok_or_else(|| {
GatewayError::Internal("external models source is missing a host".to_string())
})?;
let port = url.port_or_known_default().ok_or_else(|| {
GatewayError::Internal("external models source is missing a port".to_string())
})?;
Ok((url, host, port))
}
fn validate_admin_external_models_source_addresses(
url: &url::Url,
addresses: &[SocketAddr],
allow_insecure_test_target: bool,
) -> Result<(), GatewayError> {
if addresses.is_empty() {
return Err(GatewayError::Internal(
"external models source DNS resolution returned no addresses".to_string(),
));
}
if !allow_insecure_test_target
&& addresses.iter().any(|address| {
aether_http::is_private_or_reserved_ip(address.ip())
&& !(is_official_external_models_catalog_url(url)
&& aether_http::is_ipv4_benchmarking_fake_ip(address.ip()))
})
{
return Err(GatewayError::Internal(
"external models source resolves to a private or reserved address".to_string(),
));
}
Ok(())
}
fn is_official_external_models_catalog_url(url: &url::Url) -> bool {
url.scheme() == "https"
&& url
.host_str()
.is_some_and(|host| host.eq_ignore_ascii_case(ADMIN_EXTERNAL_MODELS_OFFICIAL_HOST))
&& url.port_or_known_default() == Some(443)
&& url.path() == ADMIN_EXTERNAL_MODELS_OFFICIAL_PATH
&& url.username().is_empty()
&& url.password().is_none()
&& url.query().is_none()
&& url.fragment().is_none()
}
async fn resolve_admin_external_models_source(
raw_url: &str,
allow_insecure_test_target: bool,
) -> Result<ResolvedAdminExternalModelsSource, GatewayError> {
let (url, host, port) =
parse_admin_external_models_source_url(raw_url, allow_insecure_test_target)?;
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
vec![SocketAddr::new(ip, port)]
} else {
aether_http::lookup_host_with_limits(
host.as_str(),
port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
)
.await
.map_err(|_| {
GatewayError::Internal("external models source DNS resolution failed".to_string())
})?
};
validate_admin_external_models_source_addresses(&url, &addresses, allow_insecure_test_target)?;
Ok(ResolvedAdminExternalModelsSource {
url,
host,
addresses,
})
}
fn build_admin_external_models_direct_client(
source: &ResolvedAdminExternalModelsSource,
) -> Result<reqwest::Client, GatewayError> {
let mut builder = aether_http::apply_http_client_config(
reqwest::Client::builder()
.no_proxy()
.redirect(reqwest::redirect::Policy::none()),
&aether_http::HttpClientConfig {
connect_timeout_ms: Some(ADMIN_EXTERNAL_MODELS_CONNECT_TIMEOUT_MS),
request_timeout_ms: Some(ADMIN_EXTERNAL_MODELS_TOTAL_TIMEOUT_MS),
http2_adaptive_window: true,
..aether_http::HttpClientConfig::default()
},
);
if source.host.parse::<IpAddr>().is_err() {
builder = builder.resolve_to_addrs(&source.host, &source.addresses);
}
builder.build().map_err(|_| {
GatewayError::Internal("external models HTTP client initialization failed".to_string())
})
}
fn normalize_admin_external_models_payload(payload: serde_json::Value) -> serde_json::Value {
mark_external_models_official_providers(&payload).unwrap_or(payload)
}
fn classify_admin_external_models_transport_error(message: &str) -> &'static str {
let message = message.to_ascii_lowercase();
if message.contains("timed out") || message.contains("timeout") {
if message.contains("dns resolution")
|| message.contains("dns lookup")
|| (message.contains("resolve") && message.contains("host"))
{
"dns_resolution"
} else if message.contains("private or reserved")
|| message.contains("ssrf")
|| message.contains("address policy")
{
"ssrf_blocked"
} else if message.contains("invalid") && message.contains("url") {
"invalid_url"
} else if message.contains("timed out") || message.contains("timeout") {
"timeout"
} else if message.contains("relay") || message.contains("tunnel") {
"relay"
@@ -96,6 +238,8 @@ fn classify_admin_external_models_transport_error(message: &str) -> &'static str
"response_decode"
} else if message.contains("connect") || message.contains("dns") || message.contains("tcp") {
"connect"
} else if message.contains("source returned http") || message.contains("status ") {
"upstream_http"
} else if message.contains("header") || message.contains("method") || message.contains("build")
{
"request_build"
@@ -116,6 +260,11 @@ async fn store_admin_external_models_cache(
};
let serialized =
serde_json::to_string(&envelope).map_err(|err| GatewayError::Internal(err.to_string()))?;
if serialized.len() > ADMIN_EXTERNAL_MODELS_CACHE_MAX_BYTES {
return Err(GatewayError::Internal(
"external models cache envelope exceeds the allowed size".to_string(),
));
}
state
.as_ref()
.runtime_kv_setex(
@@ -127,6 +276,13 @@ async fn store_admin_external_models_cache(
Ok(())
}
fn parse_admin_external_models_cache(raw: &str) -> Option<AdminExternalModelsCacheEnvelope> {
if raw.len() > ADMIN_EXTERNAL_MODELS_CACHE_MAX_BYTES {
return None;
}
serde_json::from_str::<AdminExternalModelsCacheEnvelope>(raw).ok()
}
fn normalize_admin_external_models_proxy_node_id(
value: Option<&Value>,
) -> Result<Option<String>, GatewayError> {
@@ -345,8 +501,15 @@ async fn fetch_admin_external_models_from_source(
request_id: &str,
proxy_node_id: Option<&str>,
) -> Result<serde_json::Value, GatewayError> {
let url = admin_external_models_source_url();
let source_url = admin_external_models_source_url();
// A proxy/tunnel resolves the target in its own network namespace. Do not
// resolve it locally first: local DNS may be unavailable, may intentionally
// return synthetic addresses, or may not be able to see an internal target
// that the configured proxy can reach. URL shape is still validated below,
// and the direct path retains the local DNS validation/pinning guard.
let parsed_source = parse_admin_external_models_source_url(&source_url, cfg!(test))?;
if let Some(node_id) = proxy_node_id {
let (source_url, _host, _port) = parsed_source;
let Some(proxy) = state.resolve_admin_proxy_node_snapshot(Some(node_id)).await else {
warn!(
request_id = %request_id,
@@ -374,7 +537,7 @@ async fn fetch_admin_external_models_from_source(
),
(
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.to_string(),
"true".to_string(),
"false".to_string(),
),
]);
let plan = ExecutionPlan {
@@ -385,7 +548,7 @@ async fn fetch_admin_external_models_from_source(
endpoint_id: String::new(),
key_id: String::new(),
method: http::Method::GET.as_str().to_string(),
url,
url: source_url.to_string(),
headers,
content_type: None,
content_encoding: None,
@@ -409,8 +572,12 @@ async fn fetch_admin_external_models_from_source(
..ExecutionTimeouts::default()
}),
};
let bounded_plan = crate::execution_runtime::transport::with_upstream_response_body_limit(
&plan,
ADMIN_EXTERNAL_MODELS_RESPONSE_LIMIT_BYTES,
);
let result = match state
.execute_execution_runtime_sync_plan(Some(request_id), &plan)
.execute_execution_runtime_sync_plan(Some(request_id), &bounded_plan)
.await
{
Ok(result) => result,
@@ -444,19 +611,65 @@ async fn fetch_admin_external_models_from_source(
return Ok(normalize_admin_external_models_payload(payload));
}
let response = state
.http_client()
.get(&url)
let (url, host, port) = parsed_source;
let addresses = if let Ok(ip) = host.parse::<IpAddr>() {
vec![SocketAddr::new(ip, port)]
} else {
aether_http::lookup_host_with_limits(
host.as_str(),
port,
aether_http::DEFAULT_DNS_LOOKUP_TIMEOUT,
)
.await
.map_err(|_| {
GatewayError::Internal("external models source DNS resolution failed".to_string())
})?
};
validate_admin_external_models_source_addresses(&url, &addresses, cfg!(test))?;
let source = ResolvedAdminExternalModelsSource {
url,
host,
addresses,
};
let client = build_admin_external_models_direct_client(&source)?;
let response = client
.get(source.url)
.header(reqwest::header::ACCEPT, "application/json")
.header(
reqwest::header::USER_AGENT,
"aether-gateway/external-models",
)
.send()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let response = response
.error_for_status()
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let payload = response
.json::<serde_json::Value>()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
.map_err(|err| {
let error_message = err.to_string();
let transport_error_kind =
classify_admin_external_models_transport_error(&error_message);
warn!(
request_id = %request_id,
transport_error_kind,
"external models direct request failed"
);
GatewayError::Internal("external models source request failed".to_string())
})?;
if !response.status().is_success() {
return Err(GatewayError::Internal(format!(
"external models source returned HTTP {}",
response.status().as_u16()
)));
}
let body = aether_http::read_response_bytes_with_limit(
response,
ADMIN_EXTERNAL_MODELS_RESPONSE_LIMIT_BYTES,
)
.await
.map_err(|_| {
GatewayError::Internal("external models source response read failed".to_string())
})?;
let payload = serde_json::from_slice::<serde_json::Value>(&body).map_err(|_| {
GatewayError::Internal("external models source returned invalid JSON".to_string())
})?;
Ok(normalize_admin_external_models_payload(payload))
}
@@ -470,8 +683,8 @@ pub(crate) async fn read_admin_external_models_cache(
.runtime_kv_get(ADMIN_EXTERNAL_MODELS_CACHE_KEY)
.await?
{
match serde_json::from_str::<AdminExternalModelsCacheEnvelope>(&raw) {
Ok(envelope)
match parse_admin_external_models_cache(&raw) {
Some(envelope)
if envelope.schema_version == ADMIN_EXTERNAL_MODELS_CACHE_VERSION
&& envelope.proxy_node_id == proxy_node_id =>
{
@@ -479,25 +692,36 @@ pub(crate) async fn read_admin_external_models_cache(
envelope.payload,
)));
}
Ok(_) => {}
Err(err) => {
warn!(error = %err, "failed to parse cached external models payload");
}
Some(_) => {}
None => warn!("failed to parse cached external models payload"),
}
}
match fetch_admin_external_models_from_source(state, request_id, proxy_node_id.as_deref()).await
{
Ok(payload) => {
if let Err(err) =
store_admin_external_models_cache(state, proxy_node_id.as_deref(), &payload).await
if store_admin_external_models_cache(state, proxy_node_id.as_deref(), &payload)
.await
.is_err()
{
warn!(error = ?err, "failed to store fetched external models cache");
warn!("failed to store fetched external models cache");
}
Ok(Some(payload))
}
Err(err) => {
warn!(error = ?err, "failed to fetch external models catalog");
Err(error) => {
// Keep the client-facing response generic, but leave an actionable,
// low-cardinality diagnostic for operators. The underlying error
// is intentionally not logged here because a future transport
// implementation could include URL, proxy, or credential details.
let error_message = error.into_message();
let transport_error_kind =
classify_admin_external_models_transport_error(&error_message);
warn!(
request_id = %request_id,
proxy_mode = if proxy_node_id.is_some() { "configured_node" } else { "direct" },
transport_error_kind,
"failed to fetch external models catalog"
);
Ok(None)
}
}
@@ -518,7 +742,10 @@ mod tests {
use super::{
admin_external_models_source_url, classify_admin_external_models_transport_error,
normalize_admin_external_models_payload, normalize_admin_external_models_proxy_node_id,
read_admin_external_models_cache, set_admin_external_models_source_url_for_tests,
parse_admin_external_models_cache, parse_admin_external_models_source_url,
read_admin_external_models_cache, resolve_admin_external_models_source,
set_admin_external_models_source_url_for_tests,
validate_admin_external_models_source_addresses, ADMIN_EXTERNAL_MODELS_CACHE_MAX_BYTES,
};
use crate::handlers::admin::request::AdminAppState;
use crate::tests::{start_server, AppState};
@@ -564,12 +791,22 @@ mod tests {
fn classifies_external_models_transport_errors_without_exposing_details() {
for (message, expected) in [
("request timeout after 300000ms", "timeout"),
(
"external models source DNS resolution failed",
"dns_resolution",
),
(
"external models source resolves to a private or reserved address",
"ssrf_blocked",
),
("external models source URL is invalid", "invalid_url"),
("hub relay request failed", "relay"),
("invalid proxy configuration", "proxy_config"),
("upstream response body exceeds limit", "response_too_large"),
("upstream response is not valid JSON", "invalid_json"),
("failed to decode content-encoding gzip", "response_decode"),
("tcp connect error", "connect"),
("external models source returned HTTP 503", "upstream_http"),
("invalid upstream header value", "request_build"),
("opaque execution failure", "unknown_transport"),
] {
@@ -590,6 +827,102 @@ mod tests {
);
}
#[test]
fn production_external_models_source_requires_safe_https_url_shape() {
assert!(
parse_admin_external_models_source_url("https://models.dev/api.json", false).is_ok()
);
for source_url in [
"http://models.dev/api.json",
"file:///etc/passwd",
"https://user:[email protected]/api.json",
"https://models.dev/api.json?next=http://169.254.169.254",
"https://models.dev/api.json#fragment",
] {
assert!(
parse_admin_external_models_source_url(source_url, false).is_err(),
"source URL should be rejected: {source_url}"
);
}
}
#[test]
fn external_models_cache_parser_rejects_oversized_runtime_values() {
let oversized = "x".repeat(ADMIN_EXTERNAL_MODELS_CACHE_MAX_BYTES + 1);
assert!(parse_admin_external_models_cache(&oversized).is_none());
let valid = r#"{"schema_version":2,"proxy_node_id":null,"payload":{}}"#;
assert!(parse_admin_external_models_cache(valid).is_some());
}
#[tokio::test]
async fn production_external_models_source_rejects_private_ip_literals() {
for source_url in [
"https://127.0.0.1/api.json",
"https://169.254.169.254/latest/meta-data",
"https://[::1]/api.json",
] {
assert!(
resolve_admin_external_models_source(source_url, false)
.await
.is_err(),
"private source URL should be rejected: {source_url}"
);
}
}
#[test]
fn official_external_models_catalog_allows_benchmarking_fake_ip_addresses() {
let url = url::Url::parse("https://models.dev/api.json")
.expect("official catalog URL should parse");
for address in ["198.18.75.234:443", "198.19.255.254:443"] {
let address = address
.parse()
.expect("fake IP socket address should parse");
assert!(
validate_admin_external_models_source_addresses(&url, &[address], false).is_ok(),
"official catalog should accept local proxy fake IP {address}"
);
}
}
#[test]
fn custom_external_models_sources_reject_benchmarking_fake_ip_addresses() {
let fake_ip = "198.18.75.234:443"
.parse()
.expect("fake IP socket address should parse");
for source_url in [
"https://catalog.example/api.json",
"https://models.dev/other.json",
"https://models.dev:444/api.json",
] {
let url = url::Url::parse(source_url).expect("custom catalog URL should parse");
assert!(
validate_admin_external_models_source_addresses(&url, &[fake_ip], false).is_err(),
"custom source must reject the benchmarking fake IP: {source_url}"
);
}
}
#[test]
fn official_external_models_catalog_still_rejects_private_addresses() {
let url = url::Url::parse("https://models.dev/api.json")
.expect("official catalog URL should parse");
for address in ["127.0.0.1:443", "10.0.0.1:443", "169.254.169.254:443"] {
let address = address
.parse()
.expect("private socket address should parse");
assert!(
validate_admin_external_models_source_addresses(&url, &[address], false).is_err(),
"official catalog must reject private address {address}"
);
}
}
#[tokio::test]
async fn read_external_models_fetches_remote_payload_when_cache_missing() {
let upstream = Router::new().route(
@@ -623,4 +956,38 @@ mod tests {
upstream_handle.abort();
}
#[tokio::test]
async fn direct_external_models_fetch_does_not_follow_redirects() {
let upstream = Router::new()
.route(
"/redirect",
get(|| async { axum::response::Redirect::temporary("/api.json") }),
)
.route(
"/api.json",
get(|| async {
Json(json!({
"openai": {
"name": "redirected payload",
"models": {}
}
}))
}),
);
let (upstream_url, upstream_handle) = start_server(upstream).await;
let _guard =
set_admin_external_models_source_url_for_tests(&format!("{upstream_url}/redirect"));
let state = AppState::new().expect("gateway should build");
let payload = read_admin_external_models_cache(
&AdminAppState::new(&state),
"external-models-redirect",
)
.await
.expect("external models read should not fail");
assert!(payload.is_none(), "redirected payload must not be accepted");
upstream_handle.abort();
}
}
@@ -27,6 +27,36 @@ use axum::{
Json,
};
use serde_json::json;
use std::collections::HashSet;
const MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS: usize = 100;
fn normalize_admin_global_model_batch_ids(
ids: Vec<String>,
field_name: &str,
) -> Result<Vec<String>, String> {
if ids.len() > MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS {
return Err(format!(
"{field_name} 最多 {MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS} 个"
));
}
let mut seen = HashSet::with_capacity(ids.len());
let mut normalized = Vec::with_capacity(ids.len());
for id in ids {
let trimmed = id.trim();
if trimmed.is_empty() {
// Keep the original value so batch-delete retains its existing per-item failure.
normalized.push(id);
continue;
}
let trimmed = trimmed.to_string();
if seen.insert(trimmed.clone()) {
normalized.push(trimmed);
}
}
Ok(normalized)
}
pub(super) async fn maybe_build_local_admin_global_models_write_response(
state: &AdminAppState<'_>,
@@ -212,17 +242,20 @@ async fn build_batch_delete_global_models_response(
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let ids = match normalize_admin_global_model_batch_ids(payload.ids, "ids") {
Ok(ids) => ids,
Err(detail) => return Ok(bad_request_response(detail)),
};
let mut success_count = 0usize;
let mut failed = Vec::new();
for id in payload.ids {
let trimmed = id.trim();
if trimmed.is_empty() {
for id in ids {
if id.trim().is_empty() {
failed.push(json!({"id": id, "error": "not found"}));
continue;
}
let Some(existing) = state.get_admin_global_model_by_id(trimmed).await? else {
failed.push(json!({"id": trimmed, "error": "not found"}));
let Some(existing) = state.get_admin_global_model_by_id(&id).await? else {
failed.push(json!({"id": id, "error": "not found"}));
continue;
};
if state.delete_admin_global_model(&existing.id).await? {
@@ -245,6 +278,46 @@ async fn build_batch_delete_global_models_response(
))
}
#[cfg(test)]
mod batch_boundary_tests {
use super::{normalize_admin_global_model_batch_ids, MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS};
#[test]
fn global_model_batch_ids_are_bounded_and_deduplicated() {
assert_eq!(
normalize_admin_global_model_batch_ids(
vec![
"model-2".to_string(),
"model-1".to_string(),
" model-2 ".to_string(),
" ".to_string(),
],
"ids",
)
.expect("valid ids"),
vec![
"model-2".to_string(),
"model-1".to_string(),
" ".to_string(),
]
);
assert!(normalize_admin_global_model_batch_ids(
(0..=MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS)
.map(|index| format!("model-{index}"))
.collect(),
"ids",
)
.is_err());
assert!(normalize_admin_global_model_batch_ids(
(0..MAX_ADMIN_GLOBAL_MODEL_BATCH_ITEMS)
.map(|index| format!("provider-{index}"))
.collect(),
"provider_ids",
)
.is_ok());
}
}
async fn build_assign_to_providers_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -259,10 +332,15 @@ async fn build_assign_to_providers_response(
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let provider_ids =
match normalize_admin_global_model_batch_ids(payload.provider_ids, "provider_ids") {
Ok(provider_ids) => provider_ids,
Err(detail) => return Ok(bad_request_response(detail)),
};
let payload: serde_json::Value = match build_admin_assign_global_model_to_providers_payload(
state,
&global_model_id,
payload.provider_ids,
provider_ids,
payload.create_models.unwrap_or(false),
)
.await
@@ -67,27 +67,17 @@ pub(crate) async fn build_admin_global_model_routing_payload(
.push(key);
}
let scheduling_mode = state
.read_system_config_json_value("scheduling_mode")
.await
.ok()
.flatten()
.and_then(|value| value.as_str().map(ToOwned::to_owned))
.unwrap_or_else(|| "cache_affinity".to_string());
let priority_mode = state
.read_system_config_json_value("provider_priority_mode")
.await
.ok()
.flatten()
.and_then(|value| value.as_str().map(ToOwned::to_owned))
.unwrap_or_else(|| "provider".to_string());
let keep_priority_on_conversion = state
.read_system_config_json_value("keep_priority_on_conversion")
.await
.ok()
.flatten()
.and_then(|value| value.as_bool())
.unwrap_or(false);
// The admin view reports the system-default routing strategy.
let ordering_config =
match crate::scheduler::config::read_system_default_routing_ordering_config(state.app())
.await
{
Ok(Some(config)) => config,
Ok(None) | Err(_) => crate::scheduler::config::SchedulerOrderingConfig::default(),
};
let scheduling_mode = ordering_config.scheduling_mode_str().to_string();
let priority_mode = ordering_config.priority_mode_str().to_string();
let keep_priority_on_conversion = ordering_config.keep_priority_on_conversion;
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_secs())
@@ -141,9 +141,8 @@ pub(super) async fn build_admin_monitoring_cache_affinities_response(
let key = affinity.key_id.as_ref().and_then(|id| key_by_id.get(id));
let user_api_key_name = user_api_key.and_then(|item| item.name.clone());
let user_api_key_prefix = user_api_key.and_then(|item| {
admin_monitoring_masked_user_api_key_prefix(state, item.key_encrypted.as_deref())
});
let user_api_key_prefix =
user_api_key.and_then(|item| admin_monitoring_masked_user_api_key_prefix(state, item));
let provider_name = provider.map(|item| item.name.clone());
let endpoint_url = endpoint
.map(|item| item.base_url.clone())
@@ -1,28 +1,16 @@
use crate::handlers::admin::request::AdminAppState;
use crate::handlers::shared::{masked_secret_display, open_auth_api_key_secret};
use crate::provider_key_auth::{
provider_key_auth_config_is_agent_identity, provider_key_auth_config_uses_header_authorization,
};
use aether_crypto::decrypt_python_fernet_ciphertext;
#[cfg(test)]
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
pub(super) fn admin_monitoring_masked_user_api_key_prefix(
state: &AdminAppState<'_>,
ciphertext: Option<&str>,
record: &aether_data::repository::auth::StoredAuthApiKeyExportRecord,
) -> Option<String> {
let Some(ciphertext) = ciphertext.map(str::trim).filter(|value| !value.is_empty()) else {
return None;
};
let full_key = admin_monitoring_try_decrypt_secret(state, ciphertext)?;
let prefix_len = full_key.len().min(10);
let prefix = &full_key[..prefix_len];
let suffix = if full_key.len() >= 4 {
&full_key[full_key.len().saturating_sub(4)..]
} else {
""
};
Some(format!("{prefix}...{suffix}"))
let projection = open_auth_api_key_secret(state.app(), record).ok()?;
Some(masked_secret_display(&projection.plaintext, 10, 4, "..."))
}
pub(super) fn admin_monitoring_masked_provider_key_prefix(
@@ -43,59 +31,16 @@ pub(super) fn admin_monitoring_masked_provider_key_prefix(
}
}
_ => {
let full_key = key
.encrypted_api_key
.as_deref()
.and_then(|ciphertext| admin_monitoring_try_decrypt_secret(state, ciphertext))?;
if full_key.len() <= 12 {
Some(format!("{full_key}***"))
} else {
Some(format!(
"{}***{}",
&full_key[..8],
&full_key[full_key.len().saturating_sub(4)..]
))
}
let full_key = state
.app()
.decrypt_provider_catalog_key_api_key(key)
.ok()
.flatten()?;
Some(masked_secret_display(&full_key, 8, 4, "***"))
}
}
}
fn admin_monitoring_try_decrypt_secret(
state: &AdminAppState<'_>,
ciphertext: &str,
) -> Option<String> {
let ciphertext = ciphertext.trim();
if ciphertext.is_empty() {
return None;
}
let encryption_key = state.encryption_key().map(str::trim).unwrap_or("");
if !encryption_key.is_empty() {
if let Ok(value) = decrypt_python_fernet_ciphertext(encryption_key, ciphertext) {
return Some(value);
}
}
for env_key in ["AETHER_GATEWAY_DATA_ENCRYPTION_KEY", "ENCRYPTION_KEY"] {
let Ok(candidate) = std::env::var(env_key) else {
continue;
};
let candidate = candidate.trim();
if candidate.is_empty() || candidate == encryption_key {
continue;
}
if let Ok(value) = decrypt_python_fernet_ciphertext(candidate, ciphertext) {
return Some(value);
}
}
#[cfg(test)]
if encryption_key != DEVELOPMENT_ENCRYPTION_KEY {
if let Ok(value) = decrypt_python_fernet_ciphertext(DEVELOPMENT_ENCRYPTION_KEY, ciphertext)
{
return Some(value);
}
}
None
}
pub(super) fn admin_monitoring_cache_affinity_sort_value(value: Option<&serde_json::Value>) -> f64 {
let Some(value) = value else {
return 0.0;
@@ -122,11 +67,14 @@ pub(super) fn admin_monitoring_cache_affinity_sort_value(value: Option<&serde_js
#[cfg(test)]
mod tests {
use super::admin_monitoring_masked_provider_key_prefix;
use super::{
admin_monitoring_masked_provider_key_prefix, admin_monitoring_masked_user_api_key_prefix,
};
use crate::handlers::admin::request::AdminAppState;
use crate::AppState;
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use sha2::{Digest, Sha256};
#[test]
fn monitoring_labels_agent_identity_instead_of_oauth_token() {
@@ -167,4 +115,41 @@ mod tests {
Some("[Agent Identity]")
);
}
#[test]
fn monitoring_never_exposes_complete_short_credentials() {
let app = AppState::new().expect("gateway should build");
let state = AdminAppState::new(&app);
let plaintext = "short-key";
let ciphertext = encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, plaintext)
.expect("secret should encrypt");
let mut hasher = Sha256::new();
hasher.update(plaintext.as_bytes());
let record = aether_data::repository::auth::StoredAuthApiKeyExportRecord::new(
"owner-1".to_string(),
"key-1".to_string(),
format!("{:x}", hasher.finalize()),
Some(ciphertext),
None,
None,
None,
None,
None,
None,
None,
true,
None,
false,
0,
0,
0.0,
false,
)
.expect("API-key record should build");
let masked = admin_monitoring_masked_user_api_key_prefix(&state, &record)
.expect("secret should decrypt");
assert_ne!(masked, plaintext);
assert!(!masked.contains(plaintext));
}
}
@@ -264,16 +264,12 @@ async fn list_admin_monitoring_cache_affinity_records_matching(
pub(super) async fn build_admin_monitoring_cache_snapshot(
state: &AdminAppState<'_>,
) -> Result<AdminMonitoringCacheSnapshot, GatewayError> {
let scheduling_mode = state
.read_system_config_json_value("scheduling_mode")
.await?
.and_then(|value| value.as_str().map(ToOwned::to_owned))
.unwrap_or_else(|| "cache_affinity".to_string());
let provider_priority_mode = state
.read_system_config_json_value("provider_priority_mode")
.await?
.and_then(|value| value.as_str().map(ToOwned::to_owned))
.unwrap_or_else(|| "provider".to_string());
let ordering_config =
crate::scheduler::config::read_system_default_routing_ordering_config(state.app())
.await?
.unwrap_or_default();
let scheduling_mode = ordering_config.scheduling_mode_str().to_string();
let provider_priority_mode = ordering_config.priority_mode_str().to_string();
let now = chrono::Utc::now();
let usage_summary = if state.has_usage_data_reader() {
@@ -218,7 +218,12 @@ pub(super) async fn build_admin_monitoring_resilience_snapshot(
"model": item.model,
"api_format": item.api_format,
"status_code": item.status_code,
"error_message": item.error_message,
"error_message": item
.error_category
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or("request_failed"),
}
})
})
@@ -3,6 +3,7 @@ use super::test_support::*;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::AppState;
use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
use aether_data_contracts::repository::{
candidates::{RequestCandidateStatus, StoredRequestCandidate},
provider_catalog::{
@@ -119,9 +120,12 @@ async fn admin_monitoring_cache_affinities_and_affinity_return_local_payload_fro
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(
crate::data::GatewayDataState::with_provider_catalog_reader_for_tests(provider_catalog)
.with_user_reader(user_repository)
.with_auth_api_key_reader(auth_repository),
crate::data::GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog,
)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
.with_user_reader(user_repository)
.with_auth_api_key_reader(auth_repository),
)
.with_admin_monitoring_cache_affinity_entry_for_tests(
"cache_affinity:user-key-1:openai:model-alpha",
@@ -228,9 +232,12 @@ async fn admin_monitoring_cache_affinities_and_delete_use_runtime_scheduler_affi
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(
crate::data::GatewayDataState::with_provider_catalog_reader_for_tests(provider_catalog)
.with_user_reader(user_repository)
.with_auth_api_key_reader(auth_repository),
crate::data::GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog,
)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
.with_user_reader(user_repository)
.with_auth_api_key_reader(auth_repository),
);
let affinity_cache_key =
aether_scheduler_core::build_scheduler_affinity_cache_key_for_api_key_id(
@@ -358,9 +365,12 @@ async fn admin_monitoring_cache_affinities_parse_session_scoped_scheduler_affini
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(
crate::data::GatewayDataState::with_provider_catalog_reader_for_tests(provider_catalog)
.with_user_reader(user_repository)
.with_auth_api_key_reader(auth_repository),
crate::data::GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog,
)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
.with_user_reader(user_repository)
.with_auth_api_key_reader(auth_repository),
);
let client_session = aether_scheduler_core::ClientSessionAffinity::new(
Some("Codex".to_string()),
@@ -69,7 +69,7 @@ async fn admin_monitoring_trace_request_returns_local_payload() {
assert_eq!(payload["candidates"][0]["provider_name"], json!("OpenAI"));
assert_eq!(
payload["candidates"][0]["provider_website"],
json!("https://openai.com")
json!("https://openai.com/")
);
assert_eq!(
payload["candidates"][0]["endpoint_name"],
@@ -292,10 +292,10 @@ async fn admin_monitoring_trace_request_falls_back_to_usage_routing_snapshot() {
payload["candidates"][0]["extra_data"]["execution_path"],
json!("local_execution_runtime_miss")
);
assert_eq!(
payload["candidates"][0]["extra_data"]["failure_diagnostic"]["path"],
json!("$.reasoning.summary")
);
assert!(payload["candidates"][0]["error_message"].is_null());
assert!(payload["candidates"][0]["extra_data"]
.get("failure_diagnostic")
.is_none());
}
#[tokio::test]
@@ -340,11 +340,9 @@ async fn admin_monitoring_trace_request_returns_oauth_account_label_from_auth_co
vec![sample_endpoint()],
vec![oauth_key],
));
let data_state = GatewayDataState::with_decision_trace_readers_for_tests(
request_candidates,
provider_catalog,
)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
let data_state = GatewayDataState::with_request_candidate_reader_for_tests(request_candidates)
.attach_provider_catalog_repository_for_tests(provider_catalog)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data_state);
@@ -615,7 +613,7 @@ async fn admin_monitoring_trace_request_exposes_request_path_from_usage_audit()
usage.candidate_id = Some("cand-used".to_string());
usage.request_metadata = Some(json!({
"request_path": "/v1beta/models/gemini-2.5-pro:generateContent",
"request_query_string": "alt=sse"
"request_query_string": "alt=sse&key=gemini-secret&access_token=oauth-secret"
}));
let usage_repository = Arc::new(InMemoryUsageReadRepository::seed(vec![usage]));
let data_state =
@@ -649,14 +647,19 @@ async fn admin_monitoring_trace_request_exposes_request_path_from_usage_audit()
payload["request_path_and_query"],
json!("/v1beta/models/gemini-2.5-pro:generateContent?alt=sse")
);
assert_eq!(
payload["candidates"][0]["extra_data"]["request_path_and_query"],
json!("/v1beta/models/gemini-2.5-pro:generateContent?alt=sse")
);
assert!(payload["candidates"][0]["extra_data"]
.get("request_path")
.is_none());
assert!(payload["candidates"][0]["extra_data"]
.get("request_query_string")
.is_none());
assert!(payload["candidates"][0]["extra_data"]
.get("request_path_and_query")
.is_none());
}
#[tokio::test]
async fn admin_monitoring_trace_request_exposes_failed_candidate_upstream_response_boundary() {
async fn admin_monitoring_trace_request_redacts_failed_candidate_response_payloads() {
let mut candidate = sample_candidate(
"cand-used",
"request-1",
@@ -739,19 +742,18 @@ async fn admin_monitoring_trace_request_exposes_failed_candidate_upstream_respon
let extra = &payload["candidates"][0]["extra_data"];
assert_eq!(extra["upstream_response"]["status_code"], json!(302));
assert_eq!(
extra["upstream_response"]["headers"]["location"],
json!("/")
);
assert_eq!(
extra["upstream_response"]["body"]["error"]["message"],
json!("redirect blocked")
extra["upstream_response"]["source"],
json!("upstream_response")
);
assert!(extra["upstream_response"].get("headers").is_none());
assert!(extra["upstream_response"].get("body").is_none());
assert!(extra["upstream_response"].get("body_ref").is_none());
assert!(extra.get("client_response").is_none());
assert!(extra.get("provider_response").is_none());
}
#[tokio::test]
async fn admin_monitoring_trace_request_prefers_ref_backed_usage_response_body() {
async fn admin_monitoring_trace_request_does_not_hydrate_ref_backed_usage_response_body() {
let mut candidate = sample_candidate(
"cand-used",
"request-ref-body",
@@ -832,30 +834,16 @@ async fn admin_monitoring_trace_request_prefers_ref_backed_usage_response_body()
.expect("body should read");
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
let upstream_response = &payload["candidates"][0]["extra_data"]["upstream_response"];
assert_eq!(
upstream_response["headers"],
json!({
"content-type": "application/json",
"x-request-id": "req_usage-cyber-risk-demo"
})
);
assert_eq!(
upstream_response["body"]["error"],
json!({
"type": "invalid_request",
"message": "This content was flagged for possible cybersecurity risk.",
"code": 400
})
);
assert!(upstream_response["body"].get("input").is_none());
assert_eq!(
upstream_response["body_ref"],
json!("usage://request/request-ref-body/response_body")
);
assert_eq!(upstream_response["status_code"], json!(400));
assert_eq!(upstream_response["source"], json!("upstream_response"));
assert_eq!(upstream_response["body_state"], json!("reference"));
assert!(upstream_response.get("headers").is_none());
assert!(upstream_response.get("body").is_none());
assert!(upstream_response.get("body_ref").is_none());
}
#[tokio::test]
async fn admin_monitoring_trace_request_decodes_connect_json_response_body_refs() {
async fn admin_monitoring_trace_request_does_not_expose_inline_connect_json_response_body() {
let mut candidate = sample_candidate(
"cand-used",
"request-connect",
@@ -922,19 +910,11 @@ async fn admin_monitoring_trace_request_decodes_connect_json_response_body_refs(
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
let upstream_response = &payload["candidates"][0]["extra_data"]["upstream_response"];
assert_eq!(upstream_response["status_code"], json!(429));
assert_eq!(
upstream_response["body"]["error"]["code"],
json!("resource_exhausted")
);
assert_eq!(
upstream_response["body"]["error"]["message"],
json!("quota exhausted")
);
assert_eq!(
upstream_response["body_ref"],
json!("usage://request/request-connect/response_body")
);
assert_eq!(upstream_response["source"], json!("upstream_response"));
assert_eq!(upstream_response["body_state"], json!("inline"));
assert!(upstream_response.get("headers").is_none());
assert!(upstream_response.get("body").is_none());
assert!(upstream_response.get("body_ref").is_none());
}
#[tokio::test]
@@ -11,8 +11,9 @@ use aether_admin::observability::monitoring::{
};
use aether_data_contracts::repository::{
candidates::{
DecisionTrace, DecisionTraceCandidate, RequestCandidateFinalStatus, RequestCandidateStatus,
StoredRequestCandidate,
sanitize_request_candidate_error_type, sanitize_request_candidate_extra_data,
sanitize_request_candidate_skip_reason, DecisionTrace, DecisionTraceCandidate,
RequestCandidateFinalStatus, RequestCandidateStatus, StoredRequestCandidate,
},
provider_catalog::StoredProviderCatalogKey,
usage::StoredRequestUsageAudit,
@@ -30,25 +31,6 @@ struct ResolvedAdminMonitoringTrace {
usage: Option<StoredRequestUsageAudit>,
}
async fn hydrate_admin_monitoring_trace_response_body(
state: &AdminAppState<'_>,
mut usage: StoredRequestUsageAudit,
) -> Result<StoredRequestUsageAudit, GatewayError> {
let is_error_node = !usage.status.eq_ignore_ascii_case("completed")
|| usage
.status_code
.is_some_and(|status| !(200..300).contains(&status));
let response_body_ref = if is_error_node && usage.response_body.is_none() {
usage.response_body_ref.clone()
} else {
None
};
if let Some(body_ref) = response_body_ref.as_deref() {
usage.response_body = state.resolve_request_usage_body_ref(body_ref).await?;
}
Ok(usage)
}
pub(super) async fn build_admin_monitoring_trace_request_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -111,10 +93,6 @@ async fn resolve_admin_monitoring_trace(
.read_request_usage_audit_shallow(request_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let usage = match usage {
Some(usage) => Some(hydrate_admin_monitoring_trace_response_body(state, usage).await?),
None => None,
};
return Ok(Some(ResolvedAdminMonitoringTrace { trace, usage }));
}
@@ -125,10 +103,9 @@ async fn resolve_admin_monitoring_trace(
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
{
usage_candidates.push(hydrate_admin_monitoring_trace_response_body(state, usage).await?);
usage_candidates.push(usage);
}
if let Some(usage) = state.find_request_usage_by_id(request_id).await? {
let usage = hydrate_admin_monitoring_trace_response_body(state, usage).await?;
if !usage_candidates.iter().any(|item| item.id == usage.id) {
usage_candidates.push(usage);
}
@@ -203,17 +180,23 @@ fn build_admin_monitoring_usage_routing_snapshot_trace(
endpoint_id: usage.provider_endpoint_id.clone(),
key_id: usage.provider_api_key_id.clone(),
status,
skip_reason: usage.routing_candidate_skip_reason().map(ToOwned::to_owned),
skip_reason: sanitize_request_candidate_skip_reason(
usage.routing_candidate_skip_reason().map(ToOwned::to_owned),
),
is_cached: false,
status_code: usage.status_code,
error_type: usage
.routing_local_execution_runtime_miss_reason()
.or(usage.error_category.as_deref())
.map(ToOwned::to_owned),
error_message: usage.error_message.clone(),
error_type: sanitize_request_candidate_error_type(
usage
.routing_local_execution_runtime_miss_reason()
.or(usage.error_category.as_deref())
.map(ToOwned::to_owned),
),
error_message: None,
latency_ms: usage.response_time_ms,
concurrent_requests: None,
extra_data: build_admin_monitoring_usage_routing_snapshot_extra_data(usage),
extra_data: sanitize_request_candidate_extra_data(
build_admin_monitoring_usage_routing_snapshot_extra_data(usage),
),
required_capabilities: None,
created_at_unix_ms: usage.created_at_unix_ms,
started_at_unix_ms: Some(usage.created_at_unix_ms),
@@ -485,8 +468,12 @@ fn parse_admin_monitoring_key_auth_config(
state: &AdminAppState<'_>,
key: &StoredProviderCatalogKey,
) -> Option<Map<String, Value>> {
let ciphertext = key.encrypted_auth_config.as_deref()?;
let plaintext = state.decrypt_catalog_secret_with_fallbacks(ciphertext)?;
let _ciphertext = key.encrypted_auth_config.as_deref()?;
let plaintext = state
.app()
.decrypt_provider_catalog_key_auth_config(key)
.ok()
.flatten()?;
serde_json::from_str::<Value>(&plaintext)
.ok()?
.as_object()
@@ -6,7 +6,10 @@ use aether_admin::observability::usage::{
};
use aether_data_contracts::repository::{
provider_catalog::StoredProviderCatalogEndpoint,
usage::{StoredRequestUsageAudit, UsageBodyCaptureState, UsageBodyField},
usage::{
canonical_usage_body_ref_for, StoredRequestUsageAudit, UsageBodyCaptureState,
UsageBodyField,
},
};
use axum::{
body::Body,
@@ -64,7 +67,10 @@ pub(super) async fn admin_usage_resolve_body_value(
}
Some(UsageBodyCaptureState::Reference) | None => {}
}
let resolved_ref_body = match item.body_ref(field) {
let body_ref = item
.body_ref(field)
.and_then(|body_ref| canonical_usage_body_ref_for(body_ref, &item.request_id, field));
let resolved_ref_body = match body_ref.as_deref() {
Some(body_ref) => state.resolve_request_usage_body_ref(body_ref).await?,
None => None,
};
@@ -14,7 +14,9 @@ use aether_admin::observability::usage::{
};
use aether_data::repository::users::StoredUserSummary;
use aether_data_contracts::repository::{
candidates::{RequestCandidateStatus, StoredRequestCandidate},
candidates::{
sanitize_request_candidate_extra_data, RequestCandidateStatus, StoredRequestCandidate,
},
usage::{
StoredRequestUsageAudit, UsageAuditKeywordSearchQuery, UsageAuditListQuery,
UsageAuditSummaryQuery,
@@ -255,10 +257,8 @@ fn latest_admin_usage_image_progress(
candidates
.iter()
.filter_map(|candidate| {
let progress = candidate
.extra_data
.as_ref()
.and_then(|value| value.get("image_progress"))?
let progress = sanitize_request_candidate_extra_data(candidate.extra_data.clone())?
.get("image_progress")?
.clone();
Some((
candidate
@@ -348,9 +348,6 @@ pub(super) fn admin_usage_terminal_candidate_state_override(
if let Some(status_code) = candidate.status_code {
payload["status_code"] = json!(status_code);
}
if let Some(error_message) = candidate.error_message.as_ref() {
payload["error_message"] = json!(error_message);
}
Some(payload)
}
@@ -1038,7 +1035,8 @@ mod tests {
use super::{
admin_usage_terminal_candidate_state_override, build_admin_usage_keyword_search_query,
build_admin_usage_records_query, AdminUsageSearchContext,
build_admin_usage_records_query, latest_admin_usage_image_progress,
AdminUsageSearchContext,
};
fn sample_candidate(
@@ -1140,6 +1138,32 @@ mod tests {
assert!(payload.is_none());
}
#[test]
fn admin_usage_image_progress_sanitizes_untrusted_candidate_data() {
let mut candidate =
sample_candidate(0, RequestCandidateStatus::Streaming, None, None, None);
candidate.extra_data = Some(json!({
"image_progress": {
"phase": "upstream_streaming",
"upstream_sse_frame_count": 3,
"message": "Bearer candidate-secret",
"request_body": {"token": "candidate-secret"}
}
}));
let progress = latest_admin_usage_image_progress(&[candidate])
.expect("safe progress summary should remain");
assert_eq!(
progress,
json!({
"phase": "upstream_streaming",
"upstream_sse_frame_count": 3
})
);
assert!(!progress.to_string().contains("candidate-secret"));
}
#[test]
fn admin_usage_transport_statuses_are_disjoint_in_list_and_keyword_queries() {
for status in ["websocket", "ws", "WS"] {
@@ -3,11 +3,9 @@ use crate::handlers::admin::provider::shared::support::{
ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_KEYS, ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_MODELS,
};
use crate::handlers::admin::request::AdminAppState;
use crate::handlers::admin::shared::{
decrypt_catalog_secret_with_fallbacks, json_string_list, parse_catalog_auth_config_json,
take_secret_prefix, take_secret_suffix,
};
use crate::handlers::admin::shared::{json_string_list, parse_catalog_auth_config_json};
use crate::handlers::public::matches_model_mapping_for_models;
use crate::handlers::shared::masked_secret_display;
use crate::provider_key_auth::provider_key_auth_config_is_agent_identity;
use crate::{GatewayError, LocalProviderDeleteTaskState};
use aether_data_contracts::repository::global_models::{
@@ -182,21 +180,12 @@ pub(crate) fn mapping_preview_masked_catalog_api_key(
return "***".to_string();
}
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)
.map(|value| {
let char_count = value.chars().count();
if char_count > 8 {
format!(
"{}***{}",
take_secret_prefix(&value, 4),
take_secret_suffix(&value, 4)
)
} else if char_count >= 2 {
format!("{}***", take_secret_prefix(&value, 2))
} else {
"***".to_string()
}
})
state
.as_ref()
.decrypt_provider_catalog_key_api_key(key)
.ok()
.flatten()
.map(|value| masked_secret_display(&value, 4, 4, "***"))
.unwrap_or_else(|| "***".to_string())
}
@@ -2,7 +2,9 @@ use crate::handlers::admin::provider::shared::paths::{
admin_export_key_id, admin_provider_id_for_keys, admin_reveal_key_id,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{attach_admin_audit_response, query_param_value};
use crate::handlers::admin::shared::{
attach_admin_audit_response, mark_sensitive_admin_response_no_store, query_param_value,
};
use crate::GatewayError;
use axum::{
body::{Body, Bytes},
@@ -94,13 +96,13 @@ pub(super) async fn maybe_handle(
));
};
return Ok(Some(match state.build_admin_reveal_key_payload(&key) {
Ok(payload) => attach_admin_audit_response(
Ok(payload) => mark_sensitive_admin_response_no_store(attach_admin_audit_response(
Json(payload).into_response(),
"admin_provider_key_revealed",
"reveal_provider_key",
"provider_key",
&key_id,
),
)),
Err(detail) => (
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
@@ -141,13 +143,13 @@ pub(super) async fn maybe_handle(
};
return Ok(Some(
match state.build_admin_export_key_payload(&key).await {
Ok(payload) => attach_admin_audit_response(
Ok(payload) => mark_sensitive_admin_response_no_store(attach_admin_audit_response(
Json(payload).into_response(),
"admin_provider_key_exported",
"export_provider_key",
"provider_key_export",
&key_id,
),
)),
Err(detail) => (
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
@@ -13,6 +13,33 @@ use serde_json::json;
use std::collections::BTreeSet;
use std::time::{SystemTime, UNIX_EPOCH};
const MAX_ADMIN_PROVIDER_MODEL_BATCH_ITEMS: usize = 100;
fn validate_admin_provider_model_batch(
payloads: Vec<AdminProviderModelCreateRequest>,
) -> Result<Vec<(String, AdminProviderModelCreateRequest)>, String> {
if payloads.len() > MAX_ADMIN_PROVIDER_MODEL_BATCH_ITEMS {
return Err(format!(
"批量创建模型最多支持 {MAX_ADMIN_PROVIDER_MODEL_BATCH_ITEMS} 条"
));
}
let mut normalized = Vec::with_capacity(payloads.len());
let mut seen = BTreeSet::new();
for mut payload in payloads {
let normalized_name = payload.provider_model_name.trim().to_string();
if normalized_name.is_empty() {
return Err("provider_model_name 不能为空".to_string());
}
if !seen.insert(normalized_name.clone()) {
return Err(format!("批量请求中包含重复模型 {normalized_name}"));
}
payload.provider_model_name = normalized_name.clone();
normalized.push((normalized_name, payload));
}
Ok(normalized)
}
pub(super) async fn maybe_handle(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
@@ -68,28 +95,23 @@ pub(super) async fn maybe_handle(
));
}
};
let mut created = Vec::new();
let mut seen = BTreeSet::new();
for payload in payloads {
let normalized_name = payload.provider_model_name.trim().to_string();
if normalized_name.is_empty() {
let payloads = match validate_admin_provider_model_batch(payloads) {
Ok(payloads) => payloads,
Err(detail) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "provider_model_name 不能为空" })),
)
.into_response(),
));
}
if !seen.insert(normalized_name.clone()) {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": format!("批量请求中包含重复模型 {normalized_name}") })),
Json(json!({ "detail": detail })),
)
.into_response(),
));
}
};
// Complete request validation before the first write so a bad later item cannot leave
// an earlier subset committed while the endpoint returns a validation error.
let mut staged = Vec::new();
for (normalized_name, payload) in payloads {
if admin_provider_model_name_exists(state, &provider_id, &normalized_name, None).await?
{
continue;
@@ -109,6 +131,11 @@ pub(super) async fn maybe_handle(
));
}
};
staged.push(record);
}
let mut created = Vec::with_capacity(staged.len());
for record in staged {
let Some(model) = state.create_admin_provider_model(&record).await? else {
return Ok(Some(
(
@@ -140,3 +167,36 @@ pub(super) async fn maybe_handle(
Ok(None)
}
#[cfg(test)]
mod tests {
use super::{
validate_admin_provider_model_batch, AdminProviderModelCreateRequest,
MAX_ADMIN_PROVIDER_MODEL_BATCH_ITEMS,
};
fn payload(name: &str) -> AdminProviderModelCreateRequest {
serde_json::from_value(serde_json::json!({
"provider_model_name": name,
"global_model_id": "global-1"
}))
.expect("payload should deserialize")
}
#[test]
fn provider_model_batch_is_bounded_and_prevalidates_duplicates() {
let at_limit = (0..MAX_ADMIN_PROVIDER_MODEL_BATCH_ITEMS)
.map(|index| payload(&format!("model-{index}")))
.collect();
assert!(validate_admin_provider_model_batch(at_limit).is_ok());
let oversized = (0..=MAX_ADMIN_PROVIDER_MODEL_BATCH_ITEMS)
.map(|index| payload(&format!("model-{index}")))
.collect();
assert!(validate_admin_provider_model_batch(oversized).is_err());
let duplicate = vec![payload("model-1"), payload(" model-1 ")];
assert!(validate_admin_provider_model_batch(duplicate).is_err());
assert!(validate_admin_provider_model_batch(vec![payload(" ")]).is_err());
}
}
@@ -394,7 +394,7 @@ async fn resolve_admin_provider_oauth_batch_import_tokens(
});
}
}
return Err(format!("Token 验证失败: {detail}"));
return Err("Token 验证失败".to_string());
}
};
@@ -138,12 +138,12 @@ pub(super) async fn execute_admin_provider_oauth_kiro_batch_import(
.await
{
Ok(config) => config,
Err(err) => {
Err(_) => {
failed += 1;
results.push(json!({
"index": index,
"status": "error",
"error": format!("Token 验证失败: {err}"),
"error": "Token 验证失败",
"replaced": false,
}));
maybe_report_admin_provider_oauth_batch_import_progress(
@@ -16,13 +16,13 @@ use serde_json::json;
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
#[derive(Debug, Clone, Deserialize)]
#[derive(Clone, Deserialize)]
pub(super) struct AdminProviderOAuthBatchImportRequest {
pub credentials: String,
pub proxy_node_id: Option<String>,
}
#[derive(Debug, Clone)]
#[derive(Clone)]
pub(super) struct AdminProviderOAuthBatchImportEntry {
pub parse_error: Option<String>,
pub refresh_token: Option<String>,
@@ -52,7 +52,7 @@ pub(super) struct AdminProviderOAuthBatchImportEntry {
pub rate_limit_tier: Option<String>,
}
#[derive(Debug, Clone)]
#[derive(Clone)]
pub(super) struct AdminProviderOAuthBatchImportOutcome {
pub total: usize,
pub success: usize,
@@ -663,7 +663,7 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
.collect();
}
Ok(_) => {}
Err(error) => return vec![parse_error_entry(format!("JSON 数组解析失败: {error}"))],
Err(_) => return vec![parse_error_entry("JSON 数组解析失败".to_string())],
}
}
@@ -697,8 +697,8 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
"JSON 行必须是账号对象,不能作为 raw token 导入".to_string(),
));
}
Err(error) => {
return Some(parse_error_entry(format!("JSON 行解析失败: {error}")));
Err(_) => {
return Some(parse_error_entry("JSON 行解析失败".to_string()));
}
}
}
@@ -719,7 +719,7 @@ pub(super) fn parse_admin_provider_oauth_agent_identity_import_entries(
return Err("Agent Identity 凭据不能为空".to_string());
}
let value = serde_json::from_str::<serde_json::Value>(raw)
.map_err(|error| format!("Agent Identity JSON 解析失败: {error}"))?;
.map_err(|_| "Agent Identity JSON 解析失败".to_string())?;
let entries = match &value {
serde_json::Value::Array(items) => items
.iter()
@@ -808,6 +808,14 @@ pub(super) fn apply_admin_provider_oauth_batch_import_hints(
return;
}
if provider_type == "antigravity" {
// The Google refresh-token response does not include the account email.
// Preserve the identity supplied by the imported Antigravity credentials so
// account naming and duplicate detection can use it after token exchange.
if let Some(email) = entry.email.as_ref() {
auth_config
.entry("email".to_string())
.or_insert_with(|| json!(email));
}
if let Some(project_id) = entry.project_id.as_ref() {
auth_config
.entry("project_id".to_string())
@@ -932,23 +940,7 @@ pub(super) async fn extract_admin_provider_oauth_batch_error_detail(
response: Response<Body>,
) -> String {
let status = response.status();
let raw_body = to_bytes(response.into_body(), crate::MAX_ERROR_BODY_BYTES)
.await
.ok();
if let Some(raw_body) = raw_body {
if let Ok(value) = serde_json::from_slice::<serde_json::Value>(&raw_body) {
if let Some(detail) = value.get("detail").and_then(serde_json::Value::as_str) {
let normalized = detail.trim();
if !normalized.is_empty() {
return normalized.to_string();
}
}
}
let normalized = String::from_utf8_lossy(&raw_body).trim().to_string();
if !normalized.is_empty() {
return normalized;
}
}
let _ = to_bytes(response.into_body(), crate::MAX_ERROR_BODY_BYTES).await;
format!("HTTP {}", status.as_u16())
}
@@ -1433,7 +1425,7 @@ mod tests {
fn applies_antigravity_project_and_user_agent_hints_to_auth_config() {
let entries = parse_admin_provider_oauth_batch_import_entries(
"antigravity",
r#"{"refreshToken":"rt-1","cloudaicompanionProject":{"id":"project-antigravity-2"},"userAgent":"antigravity"}"#,
r#"{"refreshToken":"rt-1","email":"[email protected]","cloudaicompanionProject":{"id":"project-antigravity-2"},"userAgent":"antigravity"}"#,
);
let mut auth_config = serde_json::Map::new();
@@ -1444,6 +1436,35 @@ mod tests {
Some(&json!("project-antigravity-2"))
);
assert_eq!(auth_config.get("user_agent"), Some(&json!("antigravity")));
assert_eq!(auth_config.get("email"), Some(&json!("[email protected]")));
}
#[test]
fn antigravity_batch_import_keeps_json_email_for_key_naming() {
let entries = parse_admin_provider_oauth_batch_import_entries(
"antigravity",
r#"{"access_token":"at-1","refresh_token":"rt-1","email":"[email protected]","project_id":"project-antigravity-3","type":"antigravity"}"#,
);
// Simulate Google's refresh-token response, which carries no email.
let mut auth_config = json!({
"provider_type": "antigravity",
"refresh_token": "rt-1",
})
.as_object()
.cloned()
.expect("auth config should be an object");
apply_admin_provider_oauth_batch_import_hints("antigravity", &entries[0], &mut auth_config);
assert_eq!(auth_config.get("email"), Some(&json!("[email protected]")));
assert_eq!(
super::super::super::helpers::admin_provider_oauth_key_name_from_auth_config(
"antigravity",
&auth_config,
Some(0),
),
"[email protected]"
);
}
#[test]
@@ -65,13 +65,13 @@ fn codex_agent_identity_import_auth_configs(
.iter()
.enumerate()
.map(|(index, entry)| {
if let Some(error) = entry.parse_error.as_deref() {
return Err(format!("第 {} 个条目无效: {error}", index + 1));
if entry.parse_error.is_some() {
return Err(format!("第 {} 个条目无效", index + 1));
}
match codex_agent_identity_auth_config_from_import(entry) {
Ok(Some(auth_config)) => Ok(auth_config),
Ok(None) => Err(format!("第 {} 个条目不是 Agent Identity", index + 1)),
Err(error) => Err(format!("第 {} 个条目无效: {error}", index + 1)),
Err(_) => Err(format!("第 {} 个条目无效", index + 1)),
}
})
.collect()
@@ -135,11 +135,10 @@ async fn acquire_provider_agent_identity_import_locks(
"其中一个 Agent Identity 正在导入或创建,请稍后重试",
));
}
Err(error) => {
Err(_) => {
tracing::warn!(
provider_id = %provider_id,
lock_key = %lock_key,
error = ?error,
"gateway Agent Identity import lock unavailable"
);
release_provider_agent_identity_import_locks(state, leases).await;
@@ -164,9 +163,8 @@ async fn release_provider_agent_identity_import_locks(
lock_key = %lease.key,
"gateway Agent Identity import lock was not owned during release"
),
Err(error) => tracing::warn!(
Err(_) => tracing::warn!(
lock_key = %lease.key,
error = ?error,
"gateway Agent Identity import lock release failed"
),
}
@@ -329,10 +327,10 @@ async fn handle_admin_provider_oauth_start_import_task(
let agent_identity_auth_configs = if agent_identity_only {
match codex_agent_identity_import_auth_configs(&payload.credentials) {
Ok(auth_configs) => Some(auth_configs),
Err(detail) => {
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("该接口仅接受有效的 Agent Identity JSON: {detail}"),
"该接口仅接受有效的 Agent Identity JSON",
));
}
}
@@ -652,13 +650,12 @@ async fn handle_admin_provider_oauth_start_import_task(
)
.await;
}
Err(err) => {
Err(_) => {
let finished_at = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(started_at);
let error_message = format!("{err:?}");
let failed_state = build_admin_provider_oauth_batch_task_state(
&task_id_for_worker,
&provider_id_for_worker,
@@ -672,7 +669,7 @@ async fn handle_admin_provider_oauth_start_import_task(
0,
0,
Some("导入任务执行失败"),
Some(error_message.as_str()),
Some("provider_oauth_batch_import_failed"),
Vec::new(),
created_at,
Some(started_at),
@@ -688,7 +685,7 @@ async fn handle_admin_provider_oauth_start_import_task(
Some(100),
Some("provider oauth batch import failed".to_string()),
None,
Some(error_message.clone()),
Some("provider_oauth_batch_import_failed".to_string()),
None,
Some(finished_at),
)
@@ -698,13 +695,15 @@ async fn handle_admin_provider_oauth_start_import_task(
&task_id_for_worker,
"failed",
"provider oauth batch import failed",
Some(json!({ "error": error_message.clone() })),
Some(json!({
"error_code": "provider_oauth_batch_import_failed"
})),
)
.await;
tracing::warn!(
task_id = %task_id_for_worker,
provider_id = %provider_id_for_worker,
error = %error_message,
error_category = "provider_oauth_batch_import_failed",
"provider oauth batch import task failed"
);
}
@@ -15,7 +15,8 @@ use super::super::super::state::{
is_fixed_provider_type_for_provider_oauth, json_non_empty_string,
};
use super::shared::{
parse_admin_provider_oauth_complete_callback, parse_admin_provider_oauth_complete_request_body,
admin_provider_oauth_state_matches_principal, parse_admin_provider_oauth_complete_callback,
parse_admin_provider_oauth_complete_request_body,
};
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_complete_key_id;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
@@ -23,7 +24,7 @@ use crate::handlers::shared::sync_provider_key_oauth_status_snapshot;
use crate::provider_key_auth::provider_key_is_oauth_managed;
use crate::GatewayError;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogKeyOAuthCredentialFence, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogUpstreamMetadataNamespaceExpectation,
};
use axum::{
@@ -96,10 +97,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
Err(response) => return Ok(response),
};
let state_data = match state
.consume_provider_oauth_state(&callback.state_nonce)
.await
{
let preview = match state.load_provider_oauth_state(&callback.state_nonce).await {
Ok(Some(state_data)) => state_data,
Ok(None) => {
return Ok(build_internal_control_error_response(
@@ -107,6 +105,15 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
"state 无效或已过期",
));
}
Err(GatewayError::Client {
status: http::StatusCode::BAD_REQUEST,
..
}) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
@@ -114,12 +121,41 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
));
}
};
if state_data.key_id != key_id {
if preview.key_id != key_id
|| !admin_provider_oauth_state_matches_principal(&preview, request_context)
{
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let state_data = match state
.consume_provider_oauth_state(&callback.state_nonce)
.await
{
Ok(Some(state_data)) if state_data == preview => state_data,
Ok(Some(_)) | Ok(None) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
Err(GatewayError::Client {
status: http::StatusCode::BAD_REQUEST,
..
}) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
));
}
};
let key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
@@ -248,7 +284,11 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
}
enrich_admin_provider_oauth_auth_config(&provider_type, &mut auth_config, &token_payload);
let Some(encrypted_api_key) = state.encrypt_catalog_secret_with_fallbacks(&access_token) else {
let Ok(encrypted_api_key) =
state
.app()
.seal_provider_catalog_key_api_key(&provider_id, &key_id, &access_token)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth encryption unavailable",
@@ -256,8 +296,10 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
};
let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let Some(encrypted_auth_config) =
state.encrypt_catalog_secret_with_fallbacks(&auth_config_json)
let Ok(encrypted_auth_config) =
state
.app()
.seal_provider_catalog_key_auth_config(&provider_id, &key_id, &auth_config_json)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
@@ -340,6 +382,12 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
CODEX_CREDENTIAL_GENERATION_KEY: uuid::Uuid::now_v7().to_string()
});
let expected_encrypted_auth_config = state_data.expected_encrypted_auth_config.clone();
let expected_credential = ProviderCatalogKeyOAuthCredentialFence {
encrypted_api_key: key.encrypted_api_key.clone(),
auth_type: key.auth_type.clone(),
provider_id: key.provider_id.clone(),
provider_type: provider.provider_type.clone(),
};
let updated_result: Result<bool, GatewayError> = async {
let max_namespace_retries = if provider_type == "codex" {
CODEX_OAUTH_COMPLETE_NAMESPACE_CAS_MAX_RETRIES
@@ -353,7 +401,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
&ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
key_id: key_id.clone(),
expected_encrypted_auth_config: expected_encrypted_auth_config.clone(),
expected_credential: None,
expected_credential: Some(expected_credential.clone()),
expected_upstream_metadata_namespace: (provider_type == "codex").then(
|| ProviderCatalogUpstreamMetadataNamespaceExpectation {
namespace: "codex".to_string(),
@@ -8,7 +8,8 @@ use super::super::super::state::{
is_fixed_provider_type_for_provider_oauth,
};
use super::shared::{
parse_admin_provider_oauth_complete_callback, parse_admin_provider_oauth_complete_request_body,
admin_provider_oauth_state_matches_principal, parse_admin_provider_oauth_complete_callback,
parse_admin_provider_oauth_complete_request_body,
};
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_complete_provider_id;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
@@ -43,10 +44,7 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
Err(response) => return Ok(response),
};
let state_data = match state
.consume_provider_oauth_state(&callback.state_nonce)
.await
{
let preview = match state.load_provider_oauth_state(&callback.state_nonce).await {
Ok(Some(state_data)) => state_data,
Ok(None) => {
return Ok(build_internal_control_error_response(
@@ -54,6 +52,15 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
"state 无效或已过期",
));
}
Err(GatewayError::Client {
status: http::StatusCode::BAD_REQUEST,
..
}) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
@@ -61,12 +68,42 @@ pub(super) async fn handle_admin_provider_oauth_complete_provider(
));
}
};
if !state_data.key_id.trim().is_empty() || state_data.provider_id != provider_id {
if !preview.key_id.trim().is_empty()
|| preview.provider_id != provider_id
|| !admin_provider_oauth_state_matches_principal(&preview, request_context)
{
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let state_data = match state
.consume_provider_oauth_state(&callback.state_nonce)
.await
{
Ok(Some(state_data)) if state_data == preview => state_data,
Ok(Some(_)) | Ok(None) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
Err(GatewayError::Client {
status: http::StatusCode::BAD_REQUEST,
..
}) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
));
}
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
@@ -1,5 +1,8 @@
use super::super::super::errors::build_internal_control_error_response;
use super::super::super::state::parse_provider_oauth_callback_params;
use crate::control::GatewayAdminPrincipalContext;
use crate::handlers::admin::request::AdminRequestContext;
use aether_data::repository::provider_oauth::StoredAdminProviderOAuthState;
use axum::{
body::{Body, Bytes},
http,
@@ -17,6 +20,31 @@ pub(super) struct AdminProviderOAuthCompleteCallback {
pub(super) state_nonce: String,
}
pub(super) fn admin_provider_oauth_state_matches_principal(
state: &StoredAdminProviderOAuthState,
request_context: &AdminRequestContext<'_>,
) -> bool {
admin_provider_oauth_state_matches_resolved_principal(
state,
request_context
.decision()
.and_then(|decision| decision.admin_principal.as_ref()),
)
}
fn admin_provider_oauth_state_matches_resolved_principal(
state: &StoredAdminProviderOAuthState,
principal: Option<&GatewayAdminPrincipalContext>,
) -> bool {
let Some(principal) = principal else {
return false;
};
state.initiated_by_user_id == principal.user_id
&& state.initiated_by_session_id == principal.session_id
&& state.initiated_by_management_token_id == principal.management_token_id
&& (principal.session_id.is_some() || principal.management_token_id.is_some())
}
pub(super) fn parse_admin_provider_oauth_callback_url(
raw_payload: &serde_json::Map<String, serde_json::Value>,
) -> Result<String, Response<Body>> {
@@ -114,3 +142,59 @@ pub(super) fn parse_admin_provider_oauth_complete_callback(
Ok(AdminProviderOAuthCompleteCallback { code, state_nonce })
}
#[cfg(test)]
mod tests {
use super::admin_provider_oauth_state_matches_resolved_principal;
use crate::control::GatewayAdminPrincipalContext;
use aether_data::repository::provider_oauth::StoredAdminProviderOAuthState;
fn state() -> StoredAdminProviderOAuthState {
StoredAdminProviderOAuthState {
nonce: "a".repeat(64),
key_id: "key-1".to_string(),
provider_id: "provider-1".to_string(),
provider_type: "codex".to_string(),
pkce_verifier: Some("verifier".to_string()),
expected_encrypted_auth_config: None,
initiated_by_user_id: "admin-1".to_string(),
initiated_by_session_id: Some("session-1".to_string()),
initiated_by_management_token_id: None,
created_at: 1,
}
}
fn principal(user_id: &str, session_id: Option<&str>) -> GatewayAdminPrincipalContext {
GatewayAdminPrincipalContext {
user_id: user_id.to_string(),
user_role: "admin".to_string(),
session_id: session_id.map(ToOwned::to_owned),
management_token_id: None,
management_token_permissions: None,
}
}
#[test]
fn provider_oauth_state_is_bound_to_exact_admin_session() {
let state = state();
let matching = principal("admin-1", Some("session-1"));
let wrong_user = principal("admin-2", Some("session-1"));
let wrong_session = principal("admin-1", Some("session-2"));
assert!(admin_provider_oauth_state_matches_resolved_principal(
&state,
Some(&matching)
));
assert!(!admin_provider_oauth_state_matches_resolved_principal(
&state,
Some(&wrong_user)
));
assert!(!admin_provider_oauth_state_matches_resolved_principal(
&state,
Some(&wrong_session)
));
assert!(!admin_provider_oauth_state_matches_resolved_principal(
&state, None
));
}
}
@@ -192,6 +192,18 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
"设备授权仅支持 Kiro / Windsurf provider",
));
}
let Some(principal) = request_context
.decision()
.and_then(|decision| decision.admin_principal.as_ref())
.filter(|principal| {
principal.session_id.is_some() || principal.management_token_id.is_some()
})
else {
return Ok(build_internal_control_error_response(
http::StatusCode::UNAUTHORIZED,
"管理员身份不可用",
));
};
let endpoint_resolution =
resolve_provider_oauth_runtime_endpoints(state, &provider, &provider_type).await?;
let runtime_endpoint = endpoint_resolution.runtime_endpoint;
@@ -240,10 +252,10 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
.build_authorize_url(&ctx, &session_id, None)
{
Ok(authorization) => authorization,
Err(error) => {
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("Windsurf 授权 URL 构建失败: {error}"),
"Windsurf 授权 URL 构建失败",
));
}
};
@@ -251,7 +263,11 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
build_windsurf_authorization_url(&authorization.authorize_url, &login_option);
let now_unix_secs = current_unix_secs();
let session = StoredAdminProviderOAuthDeviceSession {
session_id: session_id.clone(),
provider_id: provider_id.clone(),
initiated_by_user_id: principal.user_id.clone(),
initiated_by_session_id: principal.session_id.clone(),
initiated_by_management_token_id: principal.management_token_id.clone(),
region: String::new(),
client_id: String::new(),
client_secret: String::new(),
@@ -330,7 +346,11 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
);
let now_unix_secs = current_unix_secs();
let session = StoredAdminProviderOAuthDeviceSession {
session_id: session_id.clone(),
provider_id: provider_id.clone(),
initiated_by_user_id: principal.user_id.clone(),
initiated_by_session_id: principal.session_id.clone(),
initiated_by_management_token_id: principal.management_token_id.clone(),
region: "us-east-1".to_string(),
client_id: String::new(),
client_secret: String::new(),
@@ -468,7 +488,11 @@ pub(super) async fn handle_admin_provider_oauth_device_authorize(
let now_unix_secs = current_unix_secs();
let session_id = generate_provider_oauth_nonce();
let session = StoredAdminProviderOAuthDeviceSession {
session_id: session_id.clone(),
provider_id: provider_id.clone(),
initiated_by_user_id: principal.user_id.clone(),
initiated_by_session_id: principal.session_id.clone(),
initiated_by_management_token_id: principal.management_token_id.clone(),
region,
client_id,
client_secret,
@@ -0,0 +1,277 @@
use crate::handlers::admin::request::AdminAppState;
use aether_runtime_state::{RuntimeLockLease, RuntimeState};
use sha2::{Digest, Sha256};
use std::future::Future;
use std::time::Duration;
const DEVICE_POLL_LEASE_TTL: Duration = Duration::from_secs(300);
const DEVICE_POLL_LEASE_RENEW_INTERVAL: Duration = Duration::from_secs(60);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum AdminProviderOAuthDevicePollLeaseFailure {
Lost,
Unavailable,
}
pub(super) enum AdminProviderOAuthDevicePollLeaseAcquire {
Acquired(AdminProviderOAuthDevicePollLease),
Contended,
Unavailable,
}
pub(super) struct AdminProviderOAuthDevicePollLease {
runtime: RuntimeState,
lease: Option<RuntimeLockLease>,
}
impl AdminProviderOAuthDevicePollLease {
pub(super) async fn try_acquire(
state: &AdminAppState<'_>,
session_id: &str,
) -> AdminProviderOAuthDevicePollLeaseAcquire {
let runtime = state.runtime_state().clone();
let lock_key = admin_provider_oauth_device_poll_lock_key(session_id);
let owner = format!(
"aether-gateway-admin-provider-oauth-device-poll:{}",
uuid::Uuid::new_v4()
);
match runtime
.lock_try_acquire(&lock_key, &owner, DEVICE_POLL_LEASE_TTL)
.await
{
Ok(Some(lease)) => AdminProviderOAuthDevicePollLeaseAcquire::Acquired(Self {
runtime,
lease: Some(lease),
}),
Ok(None) => AdminProviderOAuthDevicePollLeaseAcquire::Contended,
Err(error) => {
tracing::warn!(
lock_key = %lock_key,
error = ?error,
"gateway provider OAuth device poll lease acquisition failed"
);
AdminProviderOAuthDevicePollLeaseAcquire::Unavailable
}
}
}
pub(super) async fn run<F, Output>(
&self,
operation: F,
) -> Result<Output, AdminProviderOAuthDevicePollLeaseFailure>
where
F: Future<Output = Output>,
{
let output = prefer_lease_loss(
operation,
wait_for_admin_provider_oauth_device_poll_lease_loss(
self.runtime.clone(),
self.lease
.as_ref()
.expect("an acquired device poll lease must contain its runtime lease")
.clone(),
),
)
.await?;
self.confirm_ownership().await?;
Ok(output)
}
async fn confirm_ownership(&self) -> Result<(), AdminProviderOAuthDevicePollLeaseFailure> {
let lease = self
.lease
.as_ref()
.expect("an acquired device poll lease must contain its runtime lease");
match self.runtime.lock_renew(lease, DEVICE_POLL_LEASE_TTL).await {
Ok(renewed) => ensure_admin_provider_oauth_device_poll_lease_renewed(renewed),
Err(error) => {
tracing::error!(
lock_key = %lease.key,
error = ?error,
"gateway provider OAuth device poll final lease renewal failed"
);
Err(AdminProviderOAuthDevicePollLeaseFailure::Unavailable)
}
}
}
pub(super) async fn release(mut self) {
let Some(lease) = self.lease.as_ref().cloned() else {
return;
};
match self.runtime.lock_release(&lease).await {
Ok(_) => {
self.lease.take();
}
Err(error) => {
tracing::warn!(
lock_key = %lease.key,
error = ?error,
"gateway provider OAuth device poll lease release failed"
);
// Keep the lease in the guard so Drop can make one best-effort retry.
}
}
}
}
impl Drop for AdminProviderOAuthDevicePollLease {
fn drop(&mut self) {
let Some(lease) = self.lease.take() else {
return;
};
let runtime = self.runtime.clone();
let Ok(handle) = tokio::runtime::Handle::try_current() else {
return;
};
handle.spawn(async move {
if let Err(error) = runtime.lock_release(&lease).await {
tracing::warn!(
lock_key = %lease.key,
error = ?error,
"gateway provider OAuth device poll lease Drop release failed"
);
}
});
}
}
fn admin_provider_oauth_device_poll_lock_key(session_id: &str) -> String {
format!(
"admin-provider-oauth-device-poll:sha256:{:x}",
Sha256::digest(session_id.as_bytes())
)
}
fn ensure_admin_provider_oauth_device_poll_lease_renewed(
renewed: bool,
) -> Result<(), AdminProviderOAuthDevicePollLeaseFailure> {
if renewed {
Ok(())
} else {
Err(AdminProviderOAuthDevicePollLeaseFailure::Lost)
}
}
async fn wait_for_admin_provider_oauth_device_poll_lease_loss(
runtime: RuntimeState,
lease: RuntimeLockLease,
) -> AdminProviderOAuthDevicePollLeaseFailure {
let first_renewal = tokio::time::Instant::now() + DEVICE_POLL_LEASE_RENEW_INTERVAL;
let mut renewal_timer =
tokio::time::interval_at(first_renewal, DEVICE_POLL_LEASE_RENEW_INTERVAL);
renewal_timer.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
renewal_timer.tick().await;
match runtime.lock_renew(&lease, DEVICE_POLL_LEASE_TTL).await {
Ok(true) => {}
Ok(false) => {
tracing::error!(
lock_key = %lease.key,
"gateway provider OAuth device poll lease was lost"
);
return AdminProviderOAuthDevicePollLeaseFailure::Lost;
}
Err(error) => {
tracing::error!(
lock_key = %lease.key,
error = ?error,
"gateway provider OAuth device poll lease renewal failed"
);
return AdminProviderOAuthDevicePollLeaseFailure::Unavailable;
}
}
}
}
async fn prefer_lease_loss<Operation, LeaseLoss, Output>(
operation: Operation,
lease_loss: LeaseLoss,
) -> Result<Output, AdminProviderOAuthDevicePollLeaseFailure>
where
Operation: Future<Output = Output>,
LeaseLoss: Future<Output = AdminProviderOAuthDevicePollLeaseFailure>,
{
tokio::pin!(operation);
tokio::pin!(lease_loss);
tokio::select! {
biased;
failure = &mut lease_loss => Err(failure),
output = &mut operation => Ok(output),
}
}
#[cfg(test)]
mod tests {
use super::{
admin_provider_oauth_device_poll_lock_key,
ensure_admin_provider_oauth_device_poll_lease_renewed, prefer_lease_loss,
AdminProviderOAuthDevicePollLeaseFailure,
};
use std::future::Future;
use std::pin::Pin;
use std::sync::{
atomic::{AtomicBool, Ordering},
Arc,
};
use std::task::{Context, Poll};
struct ReadyOperation {
polled: Arc<AtomicBool>,
dropped: Arc<AtomicBool>,
}
impl Future for ReadyOperation {
type Output = &'static str;
fn poll(self: Pin<&mut Self>, _context: &mut Context<'_>) -> Poll<Self::Output> {
self.polled.store(true, Ordering::Release);
Poll::Ready("must not be published")
}
}
impl Drop for ReadyOperation {
fn drop(&mut self) {
self.dropped.store(true, Ordering::Release);
}
}
#[test]
fn device_poll_lease_lock_key_hashes_session_id() {
let session_id = "secret-device-session";
let first = admin_provider_oauth_device_poll_lock_key(session_id);
let second = admin_provider_oauth_device_poll_lock_key(session_id);
assert_eq!(first, second);
assert!(!first.contains(session_id));
}
#[test]
fn device_poll_lease_renew_false_fails_closed() {
assert_eq!(
ensure_admin_provider_oauth_device_poll_lease_renewed(false),
Err(AdminProviderOAuthDevicePollLeaseFailure::Lost)
);
assert!(ensure_admin_provider_oauth_device_poll_lease_renewed(true).is_ok());
}
#[tokio::test]
async fn device_poll_lease_loss_future_has_priority_and_cancels_operation() {
let polled = Arc::new(AtomicBool::new(false));
let dropped = Arc::new(AtomicBool::new(false));
let operation = ReadyOperation {
polled: Arc::clone(&polled),
dropped: Arc::clone(&dropped),
};
let result = prefer_lease_loss(
operation,
std::future::ready(AdminProviderOAuthDevicePollLeaseFailure::Lost),
)
.await;
assert_eq!(result, Err(AdminProviderOAuthDevicePollLeaseFailure::Lost));
assert!(!polled.load(Ordering::Acquire));
assert!(dropped.load(Ordering::Acquire));
}
}
@@ -1,4 +1,5 @@
mod authorize;
mod lease;
mod poll;
mod session;
File diff suppressed because it is too large Load Diff
@@ -17,7 +17,7 @@ pub(super) struct AdminProviderOAuthDeviceAuthorizePayload {
pub(super) proxy_node_id: Option<String>,
}
#[derive(Debug, Deserialize)]
#[derive(Deserialize)]
pub(super) struct AdminProviderOAuthDevicePollPayload {
pub(super) session_id: String,
pub(super) callback_url: Option<String>,
@@ -225,10 +225,9 @@ async fn prepare_codex_agent_identity_enrollment(
"该 ChatGPT 账号正在创建 Agent Identity,请稍后重试",
));
}
Err(error) => {
Err(_) => {
tracing::warn!(
provider_id = %provider_id,
error = ?error,
"gateway Agent Identity enrollment lock unavailable"
);
release_codex_agent_identity_leases(state, leases).await;
@@ -296,11 +295,10 @@ fn spawn_codex_agent_identity_enrollment_heartbeat(
);
return;
}
Err(error) => {
Err(_) => {
lease_lost.store(true, Ordering::Release);
tracing::error!(
lock_key = %lease.key,
error = ?error,
"gateway Agent Identity enrollment lock renewal failed"
);
return;
@@ -316,10 +314,9 @@ async fn release_codex_agent_identity_leases(
leases: Vec<RuntimeLockLease>,
) {
for lease in leases {
if let Err(error) = state.runtime_state().lock_release(&lease).await {
if state.runtime_state().lock_release(&lease).await.is_err() {
tracing::warn!(
lock_key = %lease.key,
error = ?error,
"gateway Agent Identity enrollment lock release failed"
);
}
@@ -460,6 +457,9 @@ fn apply_single_import_hints(
.or_insert_with(|| json!(project_id));
}
for (target, keys) in [
// Antigravity token responses omit the account email, so retain the
// identity carried by the imported credential payload.
("email", &["email", "oauth_email"][..]),
(
"client_version",
&[
@@ -1198,11 +1198,10 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
match state.resolve_local_oauth_request_auth(&transport).await {
Ok(Some(_)) => true,
Ok(None) => false,
Err(error) => {
Err(_) => {
tracing::warn!(
provider_id = %provider_id,
key_id = %persisted_key.id,
error = ?error,
"gateway Agent Identity initial task registration failed"
);
false
@@ -1210,11 +1209,10 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
}
}
Ok(None) => false,
Err(error) => {
Err(_) => {
tracing::warn!(
provider_id = %provider_id,
key_id = %persisted_key.id,
error = ?error,
"gateway Agent Identity pending transport reload failed"
);
false
@@ -1461,7 +1459,8 @@ mod tests {
},
"clientVersion": "1.99.0",
"sessionId": "session-antigravity-1",
"userAgent": "antigravity"
"userAgent": "antigravity",
"email": "[email protected]"
})
.as_object()
.cloned()
@@ -1470,6 +1469,7 @@ mod tests {
apply_single_import_hints("antigravity", &payload, &mut auth_config);
assert_eq!(auth_config.get("email"), Some(&json!("[email protected]")));
assert_eq!(
auth_config.get("project_id"),
Some(&json!("project-antigravity-1"))
@@ -56,21 +56,50 @@ fn admin_provider_oauth_kiro_refresh_error(
"social refresh"
};
match error {
OAuthError::HttpStatus {
status_code,
body_excerpt,
} => {
let detail = body_excerpt.trim();
if detail.is_empty() {
format!("{prefix} 失败: HTTP {status_code}")
} else {
format!("{prefix} 失败: {detail}")
}
OAuthError::HttpStatus { status_code, .. } => {
format!("{prefix} 失败: HTTP {status_code}")
}
OAuthError::Transport(message) => format!("{prefix} 请求失败: {message}"),
OAuthError::InvalidRequest(message) => format!("{prefix} 参数无效: {message}"),
OAuthError::InvalidResponse(message) => format!("{prefix} 返回无效响应: {message}"),
error => format!("{prefix} 失败: {error}"),
OAuthError::InvalidRequest(_) => format!("{prefix} 参数无效"),
OAuthError::InvalidResponse(_) => format!("{prefix} 返回无效响应"),
OAuthError::Transport(_) => format!("{prefix} 请求失败"),
_ => format!("{prefix} 失败"),
}
}
#[cfg(test)]
mod refresh_error_tests {
use super::admin_provider_oauth_kiro_refresh_error;
use crate::handlers::admin::request::AdminKiroAuthConfig;
use aether_oauth::core::OAuthError;
#[test]
fn kiro_refresh_error_does_not_reflect_upstream_body() {
let auth_config = AdminKiroAuthConfig {
auth_method: None,
refresh_token: None,
expires_at: None,
profile_arn: None,
region: None,
auth_region: None,
api_region: None,
client_id: None,
client_secret: None,
machine_id: None,
kiro_version: None,
system_version: None,
node_version: None,
access_token: None,
};
let detail = admin_provider_oauth_kiro_refresh_error(
&auth_config,
OAuthError::HttpStatus {
status_code: 502,
body_excerpt: "authorization=Bearer upstream-secret".to_string(),
},
);
assert_eq!(detail, "social refresh 失败: HTTP 502");
assert!(!detail.contains("upstream-secret"));
}
}
@@ -88,43 +88,44 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
response::oauth_refresh_failed_bad_request_response(&error_reason),
));
}
Err(AdminLocalOAuthRefreshError::Transport { source, .. }) => {
Err(AdminLocalOAuthRefreshError::Transport { .. }) => {
tracing::warn!(
trace_id = %trace_id,
key_id = %key_id,
provider_id = %provider.id,
provider_type = %provider_type,
error = %source,
"gateway manual provider oauth refresh transport failed"
);
return Ok(RefreshDispatch::Respond(
response::oauth_refresh_failed_service_unavailable_response(source.to_string()),
response::oauth_refresh_failed_service_unavailable_response(
"Token 刷新网络请求失败",
),
));
}
Err(AdminLocalOAuthRefreshError::TransportMessage { message, .. }) => {
Err(AdminLocalOAuthRefreshError::TransportMessage { .. }) => {
tracing::warn!(
trace_id = %trace_id,
key_id = %key_id,
provider_id = %provider.id,
provider_type = %provider_type,
error = %message,
"gateway manual provider oauth refresh transport failed"
);
return Ok(RefreshDispatch::Respond(
response::oauth_refresh_failed_service_unavailable_response(message),
response::oauth_refresh_failed_service_unavailable_response(
"Token 刷新网络请求失败",
),
));
}
Err(AdminLocalOAuthRefreshError::InvalidResponse { message, .. }) => {
Err(AdminLocalOAuthRefreshError::InvalidResponse { .. }) => {
tracing::warn!(
trace_id = %trace_id,
key_id = %key_id,
provider_id = %provider.id,
provider_type = %provider_type,
reason = %message,
"gateway manual provider oauth refresh returned invalid response"
);
return Ok(RefreshDispatch::Respond(
response::oauth_refresh_failed_bad_request_response(&message),
response::oauth_refresh_failed_bad_request_response("Token 刷新响应无效"),
));
}
};
@@ -140,12 +141,7 @@ pub(super) async fn execute_admin_provider_oauth_refresh(
.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(),
)
});
.unwrap_or_else(|| helpers::refreshed_auth_config_object(state, &refreshed_key));
let refreshed_expires_at_unix_secs = refreshed_entry
.as_ref()
.and_then(|entry| entry.expires_at_unix_secs)
@@ -1,5 +1,4 @@
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::handlers::admin::shared::decrypt_catalog_secret_with_fallbacks;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
@@ -31,9 +30,13 @@ pub(super) struct RefreshSuccessContext {
pub(super) fn decrypt_auth_config(
state: &AdminAppState<'_>,
encrypted_auth_config: &str,
key: &StoredProviderCatalogKey,
) -> Option<String> {
state.decrypt_catalog_secret_with_fallbacks(encrypted_auth_config)
state
.app()
.decrypt_provider_catalog_key_auth_config(key)
.ok()
.flatten()
}
pub(super) fn parse_auth_config_object(plaintext: &str) -> Map<String, Value> {
@@ -45,10 +48,9 @@ pub(super) fn parse_auth_config_object(plaintext: &str) -> Map<String, Value> {
pub(super) fn refreshed_auth_config_object(
state: &AdminAppState<'_>,
encrypted_auth_config: Option<&str>,
key: &StoredProviderCatalogKey,
) -> Map<String, Value> {
encrypted_auth_config
.and_then(|ciphertext| decrypt_auth_config(state, ciphertext))
decrypt_auth_config(state, key)
.map(|plaintext| parse_auth_config_object(&plaintext))
.unwrap_or_default()
}
@@ -29,14 +29,13 @@ pub(super) async fn parse_admin_provider_oauth_refresh_request(
"Key 不存在",
)));
};
let Some(encrypted_auth_config) = key.encrypted_auth_config.as_deref() else {
let Some(_encrypted_auth_config) = key.encrypted_auth_config.as_deref() else {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::BAD_REQUEST,
"缺少 auth_config,无法 refresh",
)));
};
let Some(decrypted_auth_config) = helpers::decrypt_auth_config(state, encrypted_auth_config)
else {
let Some(decrypted_auth_config) = helpers::decrypt_auth_config(state, &key) else {
return Ok(RefreshDispatch::Respond(response::control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth encryption unavailable",
@@ -8,6 +8,7 @@ use crate::handlers::admin::provider::shared::paths::{
admin_provider_oauth_start_key_id, admin_provider_oauth_start_provider_id,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::mark_sensitive_admin_response_no_store;
use crate::provider_key_auth::provider_key_is_oauth_managed;
use crate::GatewayError;
use axum::{
@@ -75,6 +76,15 @@ pub(super) async fn handle_admin_provider_oauth_start_key(
"该 Provider 不支持 OAuth 授权",
));
};
let Some(principal) = request_context
.decision()
.and_then(|decision| decision.admin_principal.as_ref())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::UNAUTHORIZED,
"管理员身份不可用",
));
};
let pkce_verifier = template
.use_pkce
@@ -87,6 +97,9 @@ pub(super) async fn handle_admin_provider_oauth_start_key(
&provider_type,
pkce_verifier.as_deref(),
key.encrypted_auth_config.as_deref(),
&principal.user_id,
principal.session_id.as_deref(),
principal.management_token_id.as_deref(),
)
.await
{
@@ -99,12 +112,19 @@ pub(super) async fn handle_admin_provider_oauth_start_key(
}
};
Ok(Json(build_provider_oauth_start_response(
template,
&nonce,
code_challenge.as_deref(),
let payload =
match build_provider_oauth_start_response(template, &nonce, code_challenge.as_deref()) {
Ok(payload) => payload,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"OAuth 客户端配置不可用",
));
}
};
Ok(mark_sensitive_admin_response_no_store(
Json(payload).into_response(),
))
.into_response())
}
pub(super) async fn handle_admin_provider_oauth_start_provider(
@@ -153,6 +173,15 @@ pub(super) async fn handle_admin_provider_oauth_start_provider(
"该 Provider 不支持 OAuth 授权",
));
};
let Some(principal) = request_context
.decision()
.and_then(|decision| decision.admin_principal.as_ref())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::UNAUTHORIZED,
"管理员身份不可用",
));
};
let pkce_verifier = template
.use_pkce
@@ -165,6 +194,9 @@ pub(super) async fn handle_admin_provider_oauth_start_provider(
&provider_type,
pkce_verifier.as_deref(),
None,
&principal.user_id,
principal.session_id.as_deref(),
principal.management_token_id.as_deref(),
)
.await
{
@@ -177,10 +209,17 @@ pub(super) async fn handle_admin_provider_oauth_start_provider(
}
};
Ok(Json(build_provider_oauth_start_response(
template,
&nonce,
code_challenge.as_deref(),
let payload =
match build_provider_oauth_start_response(template, &nonce, code_challenge.as_deref()) {
Ok(payload) => payload,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"OAuth 客户端配置不可用",
));
}
};
Ok(mark_sensitive_admin_response_no_store(
Json(payload).into_response(),
))
.into_response())
}
@@ -7,8 +7,17 @@ use base64::{
use serde_json::{json, Map, Value};
use std::collections::BTreeMap;
const MAX_UNVERIFIED_JWT_PART_BYTES: usize = 64 * 1024;
fn decode_base64_url_part(value: &str) -> Option<Vec<u8>> {
URL_SAFE_NO_PAD
if value.len()
> crate::execution_runtime::transport::maximum_base64_len_for_decoded_limit(
MAX_UNVERIFIED_JWT_PART_BYTES,
)
{
return None;
}
let bytes = URL_SAFE_NO_PAD
.decode(value.as_bytes())
.or_else(|_| URL_SAFE.decode(value.as_bytes()))
.or_else(|_| {
@@ -19,7 +28,8 @@ fn decode_base64_url_part(value: &str) -> Option<Vec<u8>> {
}
URL_SAFE.decode(padded.as_bytes())
})
.ok()
.ok()?;
(bytes.len() <= MAX_UNVERIFIED_JWT_PART_BYTES).then_some(bytes)
}
fn decode_unverified_jwt_json_part(part: &str) -> Option<Map<String, Value>> {
@@ -31,15 +41,26 @@ fn decode_unverified_jwt_json_part(part: &str) -> Option<Map<String, Value>> {
}
pub(super) fn looks_like_access_token(token: &str) -> bool {
let parts = token.trim().split('.').collect::<Vec<_>>();
if parts.len() != 3 || parts.iter().any(|part| part.is_empty()) {
let mut parts = token.trim().split('.');
let Some(header_part) = parts.next().filter(|part| !part.is_empty()) else {
return false;
};
let Some(payload_part) = parts.next().filter(|part| !part.is_empty()) else {
return false;
};
let Some(_signature_part) = parts.next().filter(|part| !part.is_empty()) else {
return false;
};
// A JWT has exactly three dot-separated parts. Avoid collecting all
// attacker-controlled segments into a temporary Vec just to reject extras.
if parts.next().is_some() {
return false;
}
let Some(header) = decode_unverified_jwt_json_part(parts[0]) else {
let Some(header) = decode_unverified_jwt_json_part(header_part) else {
return false;
};
let Some(payload) = decode_unverified_jwt_json_part(parts[1]) else {
let Some(payload) = decode_unverified_jwt_json_part(payload_part) else {
return false;
};
@@ -367,6 +388,13 @@ mod tests {
assert_eq!(access_token.as_deref(), Some(token.as_str()));
}
#[test]
fn rejects_jwt_with_many_extra_segments_without_collecting() {
let token = unsigned_jwt(json!({"exp": 2_000_000_000u64}));
let oversized = format!("{token}.{}", "x.".repeat(4096));
assert!(!looks_like_access_token(&oversized));
}
#[test]
fn builds_codex_temporary_auth_config_from_access_token() {
let token = unsigned_jwt(json!({
@@ -211,12 +211,11 @@ pub(crate) async fn acquire_codex_oauth_account_locks(
release_codex_oauth_account_locks(state, leases).await;
return Err(CodexOAuthAccountLockError::Contended);
}
Err(error) => {
Err(_) => {
tracing::warn!(
provider_id = %provider_id,
lock_key = %lock_key,
operation,
error = ?error,
"gateway Codex OAuth account lock unavailable"
);
release_codex_oauth_account_locks(state, leases).await;
@@ -249,9 +248,8 @@ pub(crate) async fn release_provider_oauth_account_locks(
lock_key = %lease.key,
"gateway provider OAuth account lock was not owned during release"
),
Err(error) => tracing::warn!(
Err(_) => tracing::warn!(
lock_key = %lease.key,
error = ?error,
"gateway provider OAuth account lock release failed"
),
}
@@ -308,12 +306,11 @@ pub(crate) async fn acquire_claude_oauth_account_lock(
{
Ok(Some(lease)) => lease,
Ok(None) => return Err(ClaudeOAuthAccountLockError::Contended),
Err(error) => {
Err(_) => {
tracing::warn!(
provider_id = %provider_id,
lock_key = %lock_key,
operation,
error = ?error,
"gateway Claude OAuth account lock unavailable"
);
return Err(ClaudeOAuthAccountLockError::Unavailable);
@@ -41,7 +41,6 @@ pub(crate) fn normalize_provider_oauth_refresh_error_message(
) -> String {
let mut message = None::<String>;
let mut error_code = None::<String>;
let mut error_type = None::<String>;
if let Some(body_excerpt) = body_excerpt {
if let Ok(value) = serde_json::from_str::<serde_json::Value>(body_excerpt) {
@@ -62,12 +61,6 @@ pub(crate) fn normalize_provider_oauth_refresh_error_message(
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
error_type = error_object
.get("type")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
}
if message.is_none() {
message = object
@@ -86,14 +79,6 @@ pub(crate) fn normalize_provider_oauth_refresh_error_message(
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
}
if error_type.is_none() {
error_type = object
.get("type")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
}
}
}
}
@@ -108,7 +93,6 @@ pub(crate) fn normalize_provider_oauth_refresh_error_message(
.unwrap_or_default();
let lowered = message.to_ascii_lowercase();
let error_code = error_code.unwrap_or_default();
let error_type = error_type.unwrap_or_default();
if error_code == "refresh_token_reused"
|| lowered.contains("already been used to generate a new access token")
@@ -126,15 +110,9 @@ pub(crate) fn normalize_provider_oauth_refresh_error_message(
{
return "refresh_token 无效、已过期或已撤销,请重新登录授权".to_string();
}
if error_type == "invalid_request_error" && !message.is_empty() {
return message;
}
if !message.is_empty() {
return message;
}
status_code
.map(|status_code| format!("HTTP {status_code}"))
.unwrap_or_else(|| "未知错误".to_string())
.unwrap_or_else(|| "Token 刷新失败".to_string())
}
pub(crate) fn merge_provider_oauth_refresh_failure_reason(
@@ -182,6 +160,17 @@ mod tests {
);
}
#[test]
fn refresh_error_does_not_reflect_unknown_upstream_text_or_credentials() {
let body = r#"{"error":{"message":"authorization=Bearer upstream-secret https://user:[email protected]?q=secret","type":"invalid_request_error","code":"unexpected"}}"#;
let normalized = normalize_provider_oauth_refresh_error_message(Some(502), Some(body));
assert_eq!(normalized, "HTTP 502");
for secret in ["upstream-secret", "user:pass", "q=secret"] {
assert!(!normalized.contains(secret), "leaked {secret}");
}
}
#[test]
fn refresh_failure_does_not_replace_account_level_block() {
assert_eq!(
@@ -353,13 +353,20 @@ pub(crate) async fn create_provider_oauth_catalog_key(
proxy: Option<serde_json::Value>,
expires_at_unix_secs: Option<u64>,
) -> Result<Option<StoredProviderCatalogKey>, GatewayError> {
let Some(encrypted_api_key) = state.encrypt_catalog_secret_with_fallbacks(access_token) else {
let key_id = Uuid::new_v4().to_string();
let Ok(encrypted_api_key) =
state
.app()
.seal_provider_catalog_key_api_key(provider_id, &key_id, access_token)
else {
return Ok(None);
};
let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let Some(encrypted_auth_config) =
state.encrypt_catalog_secret_with_fallbacks(&auth_config_json)
let Ok(encrypted_auth_config) =
state
.app()
.seal_provider_catalog_key_auth_config(provider_id, &key_id, &auth_config_json)
else {
return Ok(None);
};
@@ -369,7 +376,7 @@ pub(crate) async fn create_provider_oauth_catalog_key(
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut record = StoredProviderCatalogKey::new(
Uuid::new_v4().to_string(),
key_id,
provider_id.to_string(),
name.to_string(),
"oauth".to_string(),
@@ -422,14 +429,20 @@ pub(crate) async fn update_existing_provider_oauth_catalog_key(
proxy: Option<serde_json::Value>,
expires_at_unix_secs: Option<u64>,
) -> Result<Option<StoredProviderCatalogKey>, GatewayError> {
let Some(encrypted_api_key) = state.encrypt_catalog_secret_with_fallbacks(access_token) else {
let Ok(encrypted_api_key) = state.app().seal_provider_catalog_key_api_key(
&existing_key.provider_id,
&existing_key.id,
access_token,
) else {
return Ok(None);
};
let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let Some(encrypted_auth_config) =
state.encrypt_catalog_secret_with_fallbacks(&auth_config_json)
else {
let Ok(encrypted_auth_config) = state.app().seal_provider_catalog_key_auth_config(
&existing_key.provider_id,
&existing_key.id,
&auth_config_json,
) else {
return Ok(None);
};
let now_unix_secs = SystemTime::now()
@@ -492,11 +505,10 @@ pub(super) async fn seed_provider_oauth_pool_score(
.await
{
Ok(mut providers) => providers.pop(),
Err(err) => {
Err(_) => {
tracing::debug!(
provider_id = %provider_id,
key_id = %key.id,
error = ?err,
"gateway provider oauth provisioning: failed to read provider for pool score seed"
);
return;
@@ -524,11 +536,10 @@ pub(super) async fn seed_provider_oauth_pool_score(
.await
{
Ok(mut scores) => scores.pop(),
Err(err) => {
Err(_) => {
tracing::debug!(
provider_id = %provider_id,
key_id = %key.id,
error = ?err,
"gateway provider oauth provisioning: failed to read existing pool score"
);
return;
@@ -541,11 +552,16 @@ pub(super) async fn seed_provider_oauth_pool_score(
now_unix_secs,
pool_config.score_rules,
);
if let Err(err) = state.app().data.upsert_pool_member_score(upsert).await {
if state
.app()
.data
.upsert_pool_member_score(upsert)
.await
.is_err()
{
tracing::debug!(
provider_id = %provider_id,
key_id = %key.id,
error = ?err,
"gateway provider oauth provisioning: failed to refresh pool score row"
);
}
@@ -5,17 +5,81 @@ use super::shared::{
quota_key_auto_removed, quota_refresh_success_invalid_state,
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::provider::shared::payloads::AdminImportProviderModelsRequest;
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_admin::provider::quota::parse_antigravity_usage_response;
use aether_admin::provider::quota::{
parse_antigravity_quota_summary_response, parse_antigravity_usage_response,
};
use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json;
use aether_contracts::ProxySnapshot;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_provider_pool::build_antigravity_pool_quota_request;
use aether_provider_pool::{
build_antigravity_pool_quota_request, build_antigravity_pool_quota_summary_request,
};
use serde_json::json;
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
use tracing::warn;
fn antigravity_discovered_model_ids(metadata_update: Option<&serde_json::Value>) -> Vec<String> {
metadata_update
.and_then(|value| value.pointer("/antigravity/quota_by_model"))
.and_then(serde_json::Value::as_object)
.into_iter()
.flat_map(|models| models.keys())
.map(String::as_str)
.filter(|model_id| aether_model_fetch::antigravity_model_id_is_routable(model_id))
.map(ToOwned::to_owned)
.collect()
}
async fn sync_antigravity_discovered_models(
state: &AdminAppState<'_>,
provider_id: &str,
metadata_update: Option<&serde_json::Value>,
) {
if !state.has_global_model_data_reader() || !state.has_global_model_data_writer() {
return;
}
let model_ids = antigravity_discovered_model_ids(metadata_update);
if model_ids.is_empty() {
return;
}
let result = state
.build_admin_import_provider_models_payload(
provider_id,
AdminImportProviderModelsRequest {
model_ids,
tiered_pricing: None,
price_per_request: None,
},
)
.await;
match result {
Ok(payload) => {
let errors = payload
.get("errors")
.and_then(serde_json::Value::as_array)
.map(Vec::len)
.unwrap_or(0);
if errors > 0 {
warn!(
provider_id,
errors, "Antigravity discovered-model catalog sync completed with item errors"
);
}
}
Err(error) => warn!(
provider_id,
error = %error,
"Antigravity discovered-model catalog sync failed"
),
}
}
async fn execute_antigravity_quota_plan(
state: &AdminAppState<'_>,
@@ -55,6 +119,79 @@ async fn execute_antigravity_quota_plan(
execute_provider_quota_plan(state, transport, plan, "antigravity").await
}
async fn fetch_antigravity_quota_summary_best_effort(
state: &AdminAppState<'_>,
transport: &AdminGatewayProviderTransportSnapshot,
authorization: (String, String),
project_id: &str,
identity_headers: BTreeMap<String, String>,
proxy_override: Option<&ProxySnapshot>,
) -> Option<serde_json::Value> {
let mut request_project_id = Some(project_id);
loop {
let proxy = match proxy_override {
Some(proxy) => Some(proxy.clone()),
None => {
state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
.await
}
};
let timeouts = Some(resolve_provider_quota_execution_timeouts(
state.resolve_transport_execution_timeouts(transport),
proxy.as_ref(),
));
let spec = build_antigravity_pool_quota_summary_request(
&transport.key.id,
&transport.endpoint.base_url,
authorization.clone(),
request_project_id,
identity_headers.clone(),
);
let plan = build_provider_quota_execution_plan(
transport,
spec,
proxy,
state.resolve_transport_profile(transport),
timeouts,
);
let outcome = match execute_provider_quota_plan(state, transport, plan, "antigravity").await
{
Ok(outcome) => outcome,
Err(error) => {
warn!(error = ?error, "Antigravity grouped quota request failed");
return None;
}
};
let result = match outcome {
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
warn!(detail = %detail, "Antigravity grouped quota execution failed");
return None;
}
};
if result.status_code == 200 {
return result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
.and_then(parse_antigravity_quota_summary_response);
}
if result.status_code == 403 && request_project_id.is_some() {
request_project_id = None;
continue;
}
warn!(
status_code = result.status_code,
"Antigravity grouped quota request returned a non-success status"
);
return None;
}
}
pub(crate) async fn refresh_antigravity_provider_quota_locally(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
@@ -130,21 +267,21 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
let result = match execute_antigravity_quota_plan(
state,
&transport,
authorization,
authorization.clone(),
&project_id,
identity_headers,
identity_headers.clone(),
proxy_override.as_ref(),
)
.await?
{
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
ProviderQuotaExecutionOutcome::Failure(_) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": format!("fetchAvailableModels 请求执行失败: {detail}"),
"message": "fetchAvailableModels 请求执行失败",
"status_code": 502,
}));
continue;
@@ -168,9 +305,31 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
.as_ref()
.and_then(|body| body.json_body.as_ref())
{
metadata_update = parse_antigravity_usage_response(body_json, now_unix_secs)
.map(|metadata| json!({ "antigravity": metadata }));
if metadata_update.is_some() {
if let Some(mut metadata) =
parse_antigravity_usage_response(body_json, now_unix_secs)
{
if let Some(metadata) = metadata.as_object_mut() {
metadata.insert("project_id".to_string(), json!(project_id));
}
if let Some(quota_groups) = fetch_antigravity_quota_summary_best_effort(
state,
&transport,
authorization,
&project_id,
identity_headers,
proxy_override.as_ref(),
)
.await
{
if let Some(metadata) = metadata.as_object_mut() {
metadata.insert("quota_groups".to_string(), quota_groups);
metadata.insert(
"quota_groups_updated_at".to_string(),
json!(now_unix_secs),
);
}
}
metadata_update = Some(json!({ "antigravity": metadata }));
status = "success".to_string();
} else {
status = "no_metadata".to_string();
@@ -181,21 +340,12 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
message = Some("响应中未包含配额信息".to_string());
}
} else {
let err_msg = extract_execution_error_message(&result);
message = Some(match err_msg.as_deref() {
Some(detail) if !detail.is_empty() => {
format!(
"fetchAvailableModels 返回状态码 {}: {}",
result.status_code, detail
)
}
_ => format!("fetchAvailableModels 返回状态码 {}", result.status_code),
});
message = Some(format!(
"fetchAvailableModels 返回状态码 {}",
result.status_code
));
if result.status_code == 403 {
let reason = err_msg
.clone()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| "账户访问被禁止".to_string());
let reason = "账户访问被禁止".to_string();
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason = Some(format!("账户访问被禁止: {reason}"));
metadata_update = Some(json!({
@@ -230,6 +380,10 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
continue;
}
if status == "success" {
sync_antigravity_discovered_models(state, &provider.id, metadata_update.as_ref()).await;
}
if status == "success" {
success_count += 1;
} else {
@@ -246,9 +400,11 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
if let Some(metadata) = metadata_update
.as_ref()
.and_then(|value| value.get("antigravity"))
.cloned()
{
payload.insert("metadata".to_string(), metadata);
payload.insert(
"metadata".to_string(),
admin_provider_metadata_bucket_safe_json("antigravity", Some(metadata)),
);
}
if let Some(quota_snapshot) = build_quota_snapshot_payload(
"antigravity",
@@ -1,15 +1,17 @@
use super::shared::{
build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message,
build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message_ref,
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
quota_key_auto_removed, quota_refresh_success_invalid_state,
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
};
use crate::execution_runtime::transport::decode_base64_body_with_limit;
use crate::handlers::admin::provider::shared::payloads::{
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
};
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_admin::provider::quota::parse_chatgpt_web_conversation_init_response;
use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json;
use aether_contracts::{
ExecutionResult, ProxySnapshot, ResolvedTransportProfile, TRANSPORT_BACKEND_BROWSER_WREQ,
TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY,
@@ -21,7 +23,6 @@ use aether_provider_pool::{
build_chatgpt_web_pool_quota_request, enrich_chatgpt_web_quota_metadata,
normalize_chatgpt_web_image_quota_limit,
};
use base64::Engine as _;
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
@@ -125,14 +126,21 @@ fn default_chatgpt_web_quota_transport_profile() -> ResolvedTransportProfile {
}
fn chatgpt_web_quota_error_detail(result: &ExecutionResult) -> Option<String> {
extract_execution_error_message(result).or_else(|| {
let body = result.body.as_ref()?.body_bytes_b64.as_deref()?;
let decoded = base64::engine::general_purpose::STANDARD
.decode(body)
.ok()?;
let text = String::from_utf8_lossy(&decoded).trim().to_string();
(!text.is_empty()).then_some(text)
})
extract_execution_error_message_ref(result)
.map(bound_quota_error_detail)
.or_else(|| {
let body = result.body.as_ref()?.body_bytes_b64.as_deref()?;
let decoded = decode_base64_body_with_limit(body, crate::MAX_ERROR_BODY_BYTES).ok()?;
let text = String::from_utf8_lossy(&decoded);
let text = text.trim();
(!text.is_empty()).then(|| bound_quota_error_detail(text))
})
}
fn bound_quota_error_detail(value: &str) -> String {
let value = value.trim();
let end = value.floor_char_boundary(value.len().min(crate::MAX_ERROR_BODY_BYTES));
value[..end].to_string()
}
fn chatgpt_web_is_structured_account_block(message: &str) -> bool {
@@ -162,11 +170,8 @@ fn chatgpt_web_is_structured_account_block(message: &str) -> bool {
}
fn chatgpt_web_quota_403_refresh_failed_reason(message: Option<&str>) -> String {
let detail = message
.map(str::trim)
.filter(|value| !value.is_empty())
.filter(|value| !value.contains('<'))
.unwrap_or("ChatGPT Web 访问验证失败,请检查浏览器指纹、Cloudflare 验证或代理/地区限制");
let _ = message;
let detail = "ChatGPT Web 访问验证失败,请检查浏览器指纹、Cloudflare 验证或代理/地区限制";
format!("{OAUTH_REFRESH_FAILED_PREFIX}{detail}")
}
@@ -175,14 +180,10 @@ fn chatgpt_web_quota_invalid_reason(status_code: u16, upstream_message: Option<&
if status_code == 403 && !chatgpt_web_is_structured_account_block(message) {
return chatgpt_web_quota_403_refresh_failed_reason(upstream_message);
}
let detail = if message.is_empty() {
match status_code {
401 => "ChatGPT Web Token 无效或已过期",
403 => "ChatGPT Web 账户访问受限",
_ => "ChatGPT Web 请求失败",
}
} else {
message
let detail = match status_code {
401 => "ChatGPT Web Token 无效或已过期",
403 => "ChatGPT Web 账户访问受限",
_ => "ChatGPT Web 请求失败",
};
match status_code {
401 => format!("{OAUTH_EXPIRED_PREFIX}{detail}"),
@@ -263,13 +264,13 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
.await?
{
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
ProviderQuotaExecutionOutcome::Failure(_) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": format!("conversation/init 请求执行失败: {detail}"),
"message": "conversation/init 请求执行失败",
"status_code": 502,
}));
continue;
@@ -304,6 +305,8 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
&mut metadata,
key.upstream_metadata.as_ref(),
);
metadata =
admin_provider_metadata_bucket_safe_json("chatgpt_web", Some(&metadata));
metadata_update = Some(json!({ "chatgpt_web": metadata }));
(oauth_invalid_at_unix_secs, oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
@@ -328,8 +331,7 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
};
let display_detail = invalid_reason
.as_deref()
.map(chatgpt_web_quota_result_message)
.or_else(|| err_msg.clone());
.map(chatgpt_web_quota_result_message);
message = Some(match display_detail.as_deref() {
Some(detail) if !detail.is_empty() => {
format!(
@@ -395,9 +397,11 @@ pub(crate) async fn refresh_chatgpt_web_provider_quota_locally(
if let Some(metadata) = metadata_update
.as_ref()
.and_then(|value| value.get("chatgpt_web"))
.cloned()
{
payload.insert("metadata".to_string(), metadata);
payload.insert(
"metadata".to_string(),
admin_provider_metadata_bucket_safe_json("chatgpt_web", Some(metadata)),
);
}
if let Some(quota_snapshot) = build_quota_snapshot_payload(
"chatgpt_web",
@@ -496,4 +500,65 @@ mod tests {
assert!(reason.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX));
}
#[test]
fn quota_error_detail_is_bounded_before_copying_json_message() {
let message = format!("{}界", "x".repeat(crate::MAX_ERROR_BODY_BYTES));
let result = ExecutionResult {
request_id: "chatgpt-web-quota:oversized-json".to_string(),
candidate_id: None,
status_code: 500,
headers: BTreeMap::new(),
response_observation: None,
body: Some(ResponseBody {
json_body: Some(json!({"error": {"message": message}})),
body_bytes_b64: None,
}),
telemetry: None,
error: None,
};
let detail = chatgpt_web_quota_error_detail(&result).expect("JSON error detail");
assert_eq!(detail.len(), crate::MAX_ERROR_BODY_BYTES);
assert!(detail.bytes().all(|byte| byte == b'x'));
}
#[test]
fn quota_error_detail_rejects_oversized_base64_before_decode() {
let encoded_limit =
crate::execution_runtime::transport::maximum_base64_len_for_decoded_limit(
crate::MAX_ERROR_BODY_BYTES,
);
let result = ExecutionResult {
request_id: "chatgpt-web-quota:oversized-base64".to_string(),
candidate_id: None,
status_code: 500,
headers: BTreeMap::new(),
response_observation: None,
body: Some(ResponseBody {
json_body: None,
body_bytes_b64: Some("A".repeat(encoded_limit + 1)),
}),
telemetry: None,
error: None,
};
assert_eq!(chatgpt_web_quota_error_detail(&result), None);
}
#[test]
fn quota_invalid_reason_does_not_persist_upstream_credentials() {
let reason = chatgpt_web_quota_invalid_reason(
401,
Some("authorization=Bearer upstream-secret https://user:[email protected]?q=secret"),
);
assert_eq!(
reason,
format!("{OAUTH_EXPIRED_PREFIX}ChatGPT Web Token 无效或已过期")
);
assert!(!reason.contains("upstream-secret"));
assert!(!reason.contains("user:pass"));
}
}
@@ -30,6 +30,7 @@ use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTranspo
use crate::provider_key_auth::provider_key_is_oauth_managed;
use crate::state::ProviderTransportCredentialFence;
use crate::GatewayError;
use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json;
use aether_contracts::ProxySnapshot;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyOAuthCredentialCasDelete,
@@ -282,12 +283,58 @@ fn truncate_codex_reset_credit_detail_error(message: impl Into<String>) -> Strin
let message = message.into();
let mut sanitized = message.replace('\n', " ");
if sanitized.len() > 240 {
sanitized.truncate(240);
let mut truncate_at = 240;
while !sanitized.is_char_boundary(truncate_at) {
truncate_at -= 1;
}
sanitized.truncate(truncate_at);
sanitized.push('…');
}
sanitized
}
fn safe_codex_quota_refresh_error(error: &GatewayError) -> &'static str {
if matches!(
error,
GatewayError::LocalExecutionPlanningTimeout { .. } | GatewayError::AdmissionTimeout { .. }
) {
return "Quota refresh timed out";
}
let message = match error {
GatewayError::UpstreamUnavailable { message, .. }
| GatewayError::ControlUnavailable { message, .. }
| GatewayError::Client { message, .. }
| GatewayError::Internal(message) => message.as_str(),
GatewayError::PlanUsageLimited(_)
| GatewayError::LastActiveAdminUpdateDenied
| GatewayError::LastActiveAdminDeleteDenied => return "Quota refresh failed",
GatewayError::LocalExecutionPlanningTimeout { .. }
| GatewayError::AdmissionTimeout { .. } => unreachable!(),
};
let lower = message.to_ascii_lowercase();
if lower.contains("timeout") || lower.contains("timed out") {
"Quota refresh timed out"
} else if [
"connection",
"connect",
"dns",
"network",
"proxy",
"socket",
"tls",
"certificate",
"transport",
]
.iter()
.any(|marker| lower.contains(marker))
{
"Quota refresh connection failed"
} else {
"Quota refresh failed"
}
}
fn merge_codex_reset_credit_detail_metadata(
codex_metadata: &mut Map<String, Value>,
detail_metadata: &Value,
@@ -355,26 +402,21 @@ async fn enrich_codex_reset_credit_details(
.await?
{
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
ProviderQuotaExecutionOutcome::Failure(_) => {
mark_codex_reset_credit_detail_failed(
codex_metadata,
now_unix_secs,
format!("reset credit detail 请求执行失败: {detail}"),
"reset credit detail 请求执行失败".to_string(),
);
return Ok(());
}
};
if result.status_code != 200 {
let detail = extract_execution_error_message(&result)
.unwrap_or_else(|| format!("HTTP {}", result.status_code));
mark_codex_reset_credit_detail_failed(
codex_metadata,
now_unix_secs,
format!(
"reset credit detail 返回状态码 {}: {detail}",
result.status_code
),
format!("reset credit detail 返回状态码 {}", result.status_code),
);
return Ok(());
}
@@ -506,8 +548,7 @@ async fn finish_codex_reset_replay(
}
Err(err) => {
refresh_status = "failed".to_string();
refresh_error =
Some(truncate_codex_reset_credit_detail_error(err.into_message()));
refresh_error = Some(safe_codex_quota_refresh_error(&err).to_string());
}
}
}
@@ -529,7 +570,10 @@ async fn finish_codex_reset_replay(
payload.insert("refresh_error".to_string(), json!(refresh_error));
}
if let Some(metadata) = metadata {
payload.insert("metadata".to_string(), metadata);
payload.insert(
"metadata".to_string(),
admin_provider_metadata_bucket_safe_json("codex", Some(&metadata)),
);
}
if let Some(quota_snapshot) = quota_snapshot {
payload.insert("quota_snapshot".to_string(), quota_snapshot);
@@ -702,14 +746,14 @@ pub(crate) async fn consume_codex_reset_credit_locally(
let result =
match execute_codex_reset_credit_plan(state, &transport, request_spec, None).await? {
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
ProviderQuotaExecutionOutcome::Failure(_) => {
return Ok((
StatusCode::BAD_GATEWAY,
json!({
"key_id": key.id,
"status": "error",
"outcome": "error",
"message": format!("reset credit consume 请求执行失败: {detail}"),
"message": "reset credit consume 请求执行失败",
}),
));
}
@@ -726,8 +770,6 @@ pub(crate) async fn consume_codex_reset_credit_locally(
"reset" | "already_redeemed" | "nothing_to_reset" | "no_credit"
);
if !known_terminal_outcome {
let detail = extract_execution_error_message(&result)
.unwrap_or_else(|| format!("HTTP {}", result.status_code));
return Ok((
StatusCode::BAD_GATEWAY,
json!({
@@ -735,7 +777,7 @@ pub(crate) async fn consume_codex_reset_credit_locally(
"status": "error",
"outcome": "error",
"idempotency_key": idempotency_key,
"message": format!("reset credit consume outcome is ambiguous: {detail}"),
"message": "reset credit consume outcome is ambiguous",
"status_code": result.status_code,
}),
));
@@ -802,7 +844,7 @@ pub(crate) async fn consume_codex_reset_credit_locally(
}
Err(err) => (
"failed".to_string(),
Some(truncate_codex_reset_credit_detail_error(err.into_message())),
Some(safe_codex_quota_refresh_error(&err).to_string()),
None,
None,
),
@@ -825,7 +867,7 @@ pub(crate) async fn consume_codex_reset_credit_locally(
}
Err(err) => (
"failed".to_string(),
Some(truncate_codex_reset_credit_detail_error(err.into_message())),
Some(safe_codex_quota_refresh_error(&err).to_string()),
None,
None,
),
@@ -845,7 +887,10 @@ pub(crate) async fn consume_codex_reset_credit_locally(
payload.insert("refresh_error".to_string(), json!(refresh_error));
}
if let Some(metadata) = metadata {
payload.insert("metadata".to_string(), metadata);
payload.insert(
"metadata".to_string(),
admin_provider_metadata_bucket_safe_json("codex", Some(&metadata)),
);
}
if let Some(quota_snapshot) = quota_snapshot {
payload.insert("quota_snapshot".to_string(), quota_snapshot);
@@ -1003,13 +1048,13 @@ async fn refresh_codex_provider_quota_locally_with_reset_fence(
.await?
{
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
ProviderQuotaExecutionOutcome::Failure(_) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": format!("wham/usage 请求执行失败: {detail}"),
"message": "wham/usage 请求执行失败",
"status_code": 502,
}));
continue;
@@ -1080,15 +1125,7 @@ async fn refresh_codex_provider_quota_locally_with_reset_fence(
}
} else {
let err_msg = extract_execution_error_message(&result);
message = Some(match err_msg.as_deref() {
Some(detail) if !detail.is_empty() => {
format!(
"wham/usage API 返回状态码 {}: {}",
result.status_code, detail
)
}
_ => format!("wham/usage API 返回状态码 {}", result.status_code),
});
message = Some(format!("wham/usage API 返回状态码 {}", result.status_code));
match result.status_code {
401 => {
@@ -1114,12 +1151,7 @@ async fn refresh_codex_provider_quota_locally_with_reset_fence(
codex_meta.insert("updated_at".to_string(), json!(now_unix_secs));
codex_meta.insert("account_disabled".to_string(), json!(true));
codex_meta.insert("reason".to_string(), json!("deactivated_workspace"));
codex_meta.insert(
"message".to_string(),
json!(err_msg
.clone()
.unwrap_or_else(|| "deactivated_workspace".to_string())),
);
codex_meta.insert("message".to_string(), json!("deactivated_workspace"));
let plan_type = transport
.key
.decrypted_auth_config
@@ -1206,7 +1238,7 @@ async fn refresh_codex_provider_quota_locally_with_reset_fence(
account_reset_fence_id,
coverage: quota_window_coverage,
},
Some(&expected_credential.credential),
&expected_credential.credential,
)
.await?
} else {
@@ -1214,9 +1246,6 @@ async fn refresh_codex_provider_quota_locally_with_reset_fence(
state,
&key.id,
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason.clone(),
None,
aether_admin::provider::quota::CodexQuotaMergeContext {
observed_at_unix_secs: now_unix_secs,
request_started_at_unix_ms: Some(quota_request_started_at_unix_ms),
@@ -1366,9 +1395,11 @@ async fn refresh_codex_provider_quota_locally_with_reset_fence(
if let Some(metadata_update) = metadata_update
.as_ref()
.and_then(|value| value.get("codex"))
.cloned()
{
payload.insert("metadata".to_string(), metadata_update);
payload.insert(
"metadata".to_string(),
admin_provider_metadata_bucket_safe_json("codex", Some(metadata_update)),
);
}
if let Some(quota_snapshot) = build_quota_snapshot_payload(
"codex",
@@ -1470,6 +1501,53 @@ mod tests {
);
}
#[test]
fn codex_reset_credit_detail_truncation_preserves_utf8_boundaries() {
let detail = format!("{}密钥", "a".repeat(239));
let truncated = truncate_codex_reset_credit_detail_error(detail);
assert_eq!(truncated, format!("{}…", "a".repeat(239)));
}
#[test]
fn codex_quota_refresh_error_projection_discards_internal_details() {
let secrets = [
"https://user:[email protected]/v1/quota?q=secret",
"Authorization: Bearer upstream-secret",
"user:password",
"upstream-secret",
];
let connection_error = GatewayError::Internal(format!(
"connection failed for {}; {}",
secrets[0], secrets[1]
));
let generic_error = GatewayError::Internal(format!(
"repository failure while processing {}; {}",
secrets[0], secrets[1]
));
let timeout_error = GatewayError::LocalExecutionPlanningTimeout {
trace_id: secrets[1].to_string(),
phase: "quota_refresh",
timeout_ms: 5_000,
};
let projected = [
safe_codex_quota_refresh_error(&connection_error),
safe_codex_quota_refresh_error(&generic_error),
safe_codex_quota_refresh_error(&timeout_error),
];
assert_eq!(projected[0], "Quota refresh connection failed");
assert_eq!(projected[1], "Quota refresh failed");
assert_eq!(projected[2], "Quota refresh timed out");
for safe_error in projected {
for secret in secrets {
assert!(!safe_error.contains(secret));
}
}
}
#[test]
fn codex_quota_coverage_only_replaces_observed_window_families() {
assert_eq!(
@@ -8,6 +8,7 @@ use super::shared::{
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_admin::provider::quota::parse_gemini_cli_retrieve_user_quota_response;
use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json;
use aether_contracts::ProxySnapshot;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
@@ -137,13 +138,13 @@ pub(crate) async fn refresh_gemini_cli_provider_quota_locally(
.await?
{
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
ProviderQuotaExecutionOutcome::Failure(_) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": format!("retrieveUserQuota 请求执行失败: {detail}"),
"message": "retrieveUserQuota 请求执行失败",
"status_code": 502,
}));
continue;
@@ -181,21 +182,12 @@ pub(crate) async fn refresh_gemini_cli_provider_quota_locally(
message = Some("响应中未包含配额信息".to_string());
}
} else {
let err_msg = extract_execution_error_message(&result);
message = Some(match err_msg.as_deref() {
Some(detail) if !detail.is_empty() => {
format!(
"retrieveUserQuota 返回状态码 {}: {}",
result.status_code, detail
)
}
_ => format!("retrieveUserQuota 返回状态码 {}", result.status_code),
});
message = Some(format!(
"retrieveUserQuota 返回状态码 {}",
result.status_code
));
if result.status_code == 403 {
let reason = err_msg
.clone()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| "账户访问被禁止".to_string());
let reason = "账户访问被禁止".to_string();
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason = Some(format!("账户访问被禁止: {reason}"));
metadata_update = Some(json!({
@@ -246,9 +238,11 @@ pub(crate) async fn refresh_gemini_cli_provider_quota_locally(
if let Some(metadata) = metadata_update
.as_ref()
.and_then(|value| value.get("gemini_cli"))
.cloned()
{
payload.insert("metadata".to_string(), metadata);
payload.insert(
"metadata".to_string(),
admin_provider_metadata_bucket_safe_json("gemini_cli", Some(metadata)),
);
}
if let Some(quota_snapshot) = build_quota_snapshot_payload(
"gemini_cli",
@@ -1,13 +1,15 @@
use super::shared::{
build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message,
build_quota_snapshot_payload, execute_provider_quota_plan, extract_execution_error_message_ref,
persist_provider_quota_refresh_state, quota_refresh_success_invalid_state,
resolve_provider_quota_execution_timeouts, ProviderQuotaExecutionOutcome,
};
use crate::execution_runtime::transport::decode_base64_body_with_limit;
use crate::handlers::admin::provider::shared::payloads::{
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
};
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json;
use aether_contracts::{
ExecutionPlan, ExecutionResult, ProxySnapshot, RequestBody, ResolvedTransportProfile,
};
@@ -18,7 +20,6 @@ use aether_provider_pool::{
grok_pool_tier_from_quota_bucket, grok_supported_quota_windows_for_tier,
};
use aether_provider_transport::grok_browser_profile_metadata_from_resolved_transport_profile;
use base64::Engine as _;
use serde_json::json;
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
@@ -288,14 +289,21 @@ async fn execute_grok_quota_plan(
}
fn grok_quota_error_detail(result: &ExecutionResult) -> Option<String> {
extract_execution_error_message(result).or_else(|| {
let body = result.body.as_ref()?.body_bytes_b64.as_deref()?;
let decoded = base64::engine::general_purpose::STANDARD
.decode(body)
.ok()?;
let text = String::from_utf8_lossy(&decoded).trim().to_string();
(!text.is_empty()).then_some(text)
})
extract_execution_error_message_ref(result)
.map(bound_quota_error_detail)
.or_else(|| {
let body = result.body.as_ref()?.body_bytes_b64.as_deref()?;
let decoded = decode_base64_body_with_limit(body, crate::MAX_ERROR_BODY_BYTES).ok()?;
let text = String::from_utf8_lossy(&decoded);
let text = text.trim();
(!text.is_empty()).then(|| bound_quota_error_detail(text))
})
}
fn bound_quota_error_detail(value: &str) -> String {
let value = value.trim();
let end = value.floor_char_boundary(value.len().min(crate::MAX_ERROR_BODY_BYTES));
value[..end].to_string()
}
fn grok_is_cloudflare_challenge(message: &str) -> bool {
@@ -313,14 +321,10 @@ fn grok_quota_invalid_reason(status_code: u16, upstream_message: Option<&str>) -
"{OAUTH_REFRESH_FAILED_PREFIX}Grok Cloudflare 验证失败,请重新从同一浏览器复制最新 Cookie 和 User-Agent,或配置可通过 Cloudflare 的代理运行时"
);
}
let detail = if message.is_empty() {
match status_code {
401 => "Grok Token 无效或已过期",
403 => "Grok 账户访问受限",
_ => "Grok 请求失败",
}
} else {
message
let detail = match status_code {
401 => "Grok Token 无效或已过期",
403 => "Grok 账户访问受限",
_ => "Grok 请求失败",
};
match status_code {
401 => format!("{OAUTH_EXPIRED_PREFIX}{detail}"),
@@ -406,8 +410,8 @@ pub(crate) async fn refresh_grok_provider_quota_locally(
.await?
{
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
last_error_message = Some(format!("rate-limits 请求执行失败: {detail}"));
ProviderQuotaExecutionOutcome::Failure(_) => {
last_error_message = Some("rate-limits 请求执行失败".to_string());
continue;
}
};
@@ -467,10 +471,8 @@ pub(crate) async fn refresh_grok_provider_quota_locally(
));
last_error_message = invalid_reason.as_deref().map(grok_quota_result_message);
} else {
let error_detail =
grok_quota_error_detail(&result).unwrap_or_else(|| "Grok 请求失败".to_string());
last_error_message = Some(format!(
"Grok rate-limits 请求失败({}): {error_detail}",
"Grok rate-limits 请求失败 ({})",
result.status_code
));
}
@@ -544,8 +546,11 @@ pub(crate) async fn refresh_grok_provider_quota_locally(
"status".to_string(),
json!(if refreshed { "success" } else { "error" }),
);
if let Some(metadata) = metadata_update.get("quota_by_model").cloned() {
payload.insert("metadata".to_string(), metadata);
if let Some(metadata) = metadata_update.get("quota_by_model") {
payload.insert(
"metadata".to_string(),
admin_provider_metadata_bucket_safe_json("grok", Some(metadata)),
);
}
if let Some(quota_snapshot) = build_quota_snapshot_payload(
"grok",
@@ -813,6 +818,52 @@ mod tests {
assert!(reason.contains("Cloudflare"));
}
#[test]
fn quota_error_detail_is_bounded_before_copying_json_message() {
let message = format!("{}界", "x".repeat(crate::MAX_ERROR_BODY_BYTES));
let result = ExecutionResult {
request_id: "grok-quota:oversized-json".to_string(),
candidate_id: None,
status_code: 500,
headers: BTreeMap::new(),
response_observation: None,
body: Some(ResponseBody {
json_body: Some(json!({"error": {"message": message}})),
body_bytes_b64: None,
}),
telemetry: None,
error: None,
};
let detail = grok_quota_error_detail(&result).expect("JSON error detail");
assert_eq!(detail.len(), crate::MAX_ERROR_BODY_BYTES);
assert!(detail.bytes().all(|byte| byte == b'x'));
}
#[test]
fn quota_error_detail_rejects_oversized_base64_before_decode() {
let encoded_limit =
crate::execution_runtime::transport::maximum_base64_len_for_decoded_limit(
crate::MAX_ERROR_BODY_BYTES,
);
let result = ExecutionResult {
request_id: "grok-quota:oversized-base64".to_string(),
candidate_id: None,
status_code: 500,
headers: BTreeMap::new(),
response_observation: None,
body: Some(ResponseBody {
json_body: None,
body_bytes_b64: Some("A".repeat(encoded_limit + 1)),
}),
telemetry: None,
error: None,
};
assert_eq!(grok_quota_error_detail(&result), None);
}
#[test]
fn quota_result_message_removes_status_prefix() {
let reason = format!("{OAUTH_REFRESH_FAILED_PREFIX}Grok Cloudflare 验证失败");
@@ -822,4 +873,16 @@ mod tests {
"Grok Cloudflare 验证失败"
);
}
#[test]
fn quota_invalid_reason_does_not_persist_upstream_credentials() {
let reason = grok_quota_invalid_reason(
401,
Some("authorization=Bearer upstream-secret https://user:[email protected]?q=secret"),
);
assert_eq!(reason, "[OAUTH_EXPIRED] Grok Token 无效或已过期");
assert!(!reason.contains("upstream-secret"));
assert!(!reason.contains("user:pass"));
}
}
@@ -5,18 +5,17 @@ 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,
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
oauth_refresh_auto_removed_result, persist_credential_fenced_provider_quota_refresh_state,
persist_quota_oauth_refresh_failure_state, provider_auto_remove_quota_exhausted_keys,
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::request::{AdminAppState, AdminLocalOAuthRefreshError};
use crate::GatewayError;
use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json;
use aether_contracts::ProxySnapshot;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_provider_transport::kiro::build_kiro_request_auth_from_config;
use aether_provider_transport::{CachedOAuthEntry, LocalResolvedOAuthRequestAuth};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
@@ -54,20 +53,6 @@ fn kiro_quota_error_is_account_banned(detail: Option<&str>) -> bool {
.any(|keyword| normalized.contains(keyword))
}
fn kiro_auth_from_refreshed_entry(
entry: &CachedOAuthEntry,
) -> Option<LocalResolvedOAuthRequestAuth> {
if !entry.provider_type.trim().eq_ignore_ascii_case("kiro") {
return None;
}
let auth_config = entry
.metadata
.as_ref()
.and_then(aether_provider_transport::kiro::KiroAuthConfig::from_json_value)?;
let auth = build_kiro_request_auth_from_config(auth_config, None)?;
Some(LocalResolvedOAuthRequestAuth::Kiro(auth))
}
fn kiro_quota_refresh_failure_status(err: &AdminLocalOAuthRefreshError) -> Option<u16> {
match err {
AdminLocalOAuthRefreshError::HttpStatus { status_code, .. } => Some(*status_code),
@@ -77,12 +62,13 @@ fn kiro_quota_refresh_failure_status(err: &AdminLocalOAuthRefreshError) -> Optio
fn kiro_quota_refresh_failure_message(err: &AdminLocalOAuthRefreshError) -> String {
match err {
AdminLocalOAuthRefreshError::HttpStatus {
status_code,
body_excerpt,
..
} => format!("Kiro Token 刷新失败 ({status_code}): {body_excerpt}"),
_ => format!("Kiro Token 刷新失败: {err}"),
AdminLocalOAuthRefreshError::HttpStatus { status_code, .. } => {
format!("Kiro Token 刷新失败: HTTP {status_code}")
}
AdminLocalOAuthRefreshError::InvalidResponse { .. } => {
"Kiro Token 刷新失败: 无效响应".to_string()
}
_ => "Kiro Token 刷新失败: 网络错误".to_string(),
}
}
@@ -118,36 +104,8 @@ pub(crate) async fn refresh_kiro_provider_quota_locally(
}
};
let auth = match state.force_local_oauth_refresh_entry(&transport).await {
Ok(Some(entry)) => match kiro_auth_from_refreshed_entry(&entry) {
Some(LocalResolvedOAuthRequestAuth::Kiro(auth)) => auth,
_ => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Kiro Token 刷新成功但认证信息解析失败",
}));
continue;
}
},
Ok(None) => match state
.resolve_local_oauth_kiro_request_auth(&transport)
.await?
{
Some(auth) => auth,
None => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 Kiro 认证配置 (auth_config)",
}));
continue;
}
},
match state.force_local_oauth_refresh_entry(&transport).await {
Ok(_) => {}
Err(err) => {
if persist_quota_oauth_refresh_failure_state(state, &transport, &err).await?
|| super::shared::quota_key_auto_removed(state, &key.id).await?
@@ -171,19 +129,60 @@ pub(crate) async fn refresh_kiro_provider_quota_locally(
results.push(serde_json::Value::Object(payload));
continue;
}
}
let Some(transport) = state
.read_provider_transport_snapshot_uncached(&provider.id, &endpoint.id, &key.id)
.await?
else {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Kiro Token 刷新后凭据快照不可用",
}));
continue;
};
let Some(auth) = state
.resolve_local_oauth_kiro_request_auth(&transport)
.await?
else {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 Kiro 认证配置 (auth_config)",
}));
continue;
};
let Some(credential_fence) = state
.app()
.capture_provider_transport_credential_fence(&transport)
.await?
else {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Kiro 凭据在请求前已变化",
}));
continue;
};
let result =
match execute_kiro_quota_plan(state, &transport, &auth, proxy_override.as_ref()).await?
{
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
ProviderQuotaExecutionOutcome::Failure(_) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": format!("getUsageLimits 请求执行失败: {detail}"),
"message": "getUsageLimits 请求执行失败",
"status_code": 502,
}));
continue;
@@ -230,11 +229,12 @@ pub(crate) async fn refresh_kiro_provider_quota_locally(
.or_insert_with(|| json!("kiro"));
let auth_config_json =
serde_json::Value::Object(auth_config_object).to_string();
if let Some(auth_config_json) =
state.encrypt_catalog_secret_with_fallbacks(auth_config_json.as_str())
{
encrypted_auth_config = Some(auth_config_json);
}
encrypted_auth_config =
Some(state.app().seal_provider_catalog_key_auth_config(
&transport.provider.id,
&transport.key.id,
auth_config_json.as_str(),
)?);
status = "success".to_string();
} else {
status = "no_metadata".to_string();
@@ -246,25 +246,14 @@ pub(crate) async fn refresh_kiro_provider_quota_locally(
}
} else {
let err_msg = extract_execution_error_message(&result);
message = Some(match err_msg.as_deref() {
Some(detail) if !detail.is_empty() => {
format!(
"getUsageLimits 返回状态码 {}: {}",
result.status_code, detail
)
}
_ => format!("getUsageLimits 返回状态码 {}", result.status_code),
});
message = Some(format!("getUsageLimits 返回状态码 {}", result.status_code));
match result.status_code {
401 => {
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason = Some("Kiro Token 无效或已过期".to_string());
}
403 | 423 => {
let reason = err_msg
.clone()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| format!("HTTP {}", result.status_code));
let reason = format!("HTTP {}", result.status_code);
if kiro_quota_error_is_token_invalid(err_msg.as_deref()) {
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason = Some("Kiro Token 无效或已过期".to_string());
@@ -286,13 +275,14 @@ pub(crate) async fn refresh_kiro_provider_quota_locally(
}
}
if !persist_provider_quota_refresh_state(
if !persist_credential_fenced_provider_quota_refresh_state(
state,
&key.id,
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason,
encrypted_auth_config,
&credential_fence,
)
.await?
{
@@ -336,12 +326,11 @@ pub(crate) async fn refresh_kiro_provider_quota_locally(
if let Some(message) = message {
payload.insert("message".to_string(), json!(message));
}
if let Some(metadata) = metadata_update
.as_ref()
.and_then(|value| value.get("kiro"))
.cloned()
{
payload.insert("metadata".to_string(), metadata);
if let Some(metadata) = metadata_update.as_ref().and_then(|value| value.get("kiro")) {
payload.insert(
"metadata".to_string(),
admin_provider_metadata_bucket_safe_json("kiro", Some(metadata)),
);
}
if let Some(quota_snapshot) = build_quota_snapshot_payload(
"kiro",
@@ -369,7 +358,11 @@ pub(crate) async fn refresh_kiro_provider_quota_locally(
#[cfg(test)]
mod tests {
use super::{kiro_quota_error_is_account_banned, kiro_quota_error_is_token_invalid};
use super::{
kiro_quota_error_is_account_banned, kiro_quota_error_is_token_invalid,
kiro_quota_refresh_failure_message,
};
use crate::handlers::admin::request::AdminLocalOAuthRefreshError;
#[test]
fn bearer_token_invalid_is_not_classified_as_banned() {
@@ -394,4 +387,20 @@ mod tests {
assert!(!kiro_quota_error_is_token_invalid(detail));
assert!(kiro_quota_error_is_account_banned(detail));
}
#[test]
fn quota_refresh_failure_does_not_reflect_upstream_body() {
let message =
kiro_quota_refresh_failure_message(&AdminLocalOAuthRefreshError::HttpStatus {
provider_type: "kiro",
status_code: 502,
body_excerpt:
"authorization=Bearer upstream-secret https://user:[email protected]?q=secret"
.to_string(),
});
assert_eq!(message, "Kiro Token 刷新失败: HTTP 502");
assert!(!message.contains("upstream-secret"));
assert!(!message.contains("user:pass"));
}
}
File diff suppressed because it is too large Load Diff
@@ -6,6 +6,7 @@ use super::shared::{
};
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
use crate::GatewayError;
use aether_admin::provider::redaction::admin_provider_metadata_bucket_safe_json;
use aether_contracts::ProxySnapshot;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
@@ -132,7 +133,7 @@ fn merge_windsurf_probe_metadata(
user_status_metadata
}
fn append_windsurf_probe_warning(metadata: &mut serde_json::Value, probe: &str, message: String) {
fn append_windsurf_probe_warning(metadata: &mut serde_json::Value, probe: &str, code: &str) {
let Some(target) = metadata.as_object_mut() else {
return;
};
@@ -142,7 +143,7 @@ fn append_windsurf_probe_warning(metadata: &mut serde_json::Value, probe: &str,
if let Some(items) = warnings.as_array_mut() {
items.push(json!({
"probe": probe,
"message": message,
"code": code,
}));
}
}
@@ -162,7 +163,13 @@ fn build_windsurf_metadata_update(
for (key, value) in patch_object {
merged_bucket.insert(key.clone(), value.clone());
}
json!({ "windsurf": merged_bucket })
let merged_bucket = serde_json::Value::Object(merged_bucket);
json!({
"windsurf": admin_provider_metadata_bucket_safe_json(
"windsurf",
Some(&merged_bucket),
)
})
}
fn sanitize_windsurf_probe_detail(detail: impl AsRef<str>) -> String {
@@ -170,77 +177,7 @@ fn sanitize_windsurf_probe_detail(detail: impl AsRef<str>) -> String {
if detail.is_empty() {
return "-".to_string();
}
if let Ok(mut value) = serde_json::from_str::<serde_json::Value>(detail) {
redact_windsurf_sensitive_json(&mut value);
return value.to_string().chars().take(500).collect();
}
if contains_windsurf_sensitive_marker(detail) {
"[REDACTED upstream error body]".to_string()
} else {
detail.chars().take(500).collect()
}
}
fn redact_windsurf_sensitive_json(value: &mut serde_json::Value) {
match value {
serde_json::Value::Object(object) => {
for (key, value) in object {
if is_windsurf_sensitive_key(key) {
*value = json!("[REDACTED]");
} else {
redact_windsurf_sensitive_json(value);
}
}
}
serde_json::Value::Array(items) => {
for item in items {
redact_windsurf_sensitive_json(item);
}
}
serde_json::Value::String(text) if looks_like_windsurf_secret(text) => {
*text = "[REDACTED]".to_string();
}
_ => {}
}
}
fn is_windsurf_sensitive_key(key: &str) -> bool {
let normalized = key
.chars()
.filter(|ch| ch.is_ascii_alphanumeric())
.collect::<String>()
.to_ascii_lowercase();
normalized.contains("token")
|| normalized.contains("apikey")
|| normalized.contains("password")
|| normalized.contains("authorization")
|| normalized.contains("secret")
}
fn looks_like_windsurf_secret(value: &str) -> bool {
let value = value.trim();
value.starts_with("devin-session-token$")
|| value.starts_with("sk-")
|| (value.len() > 80 && value.split('.').count() == 3)
}
fn contains_windsurf_sensitive_marker(value: &str) -> bool {
let lowered = value.to_ascii_lowercase();
[
"token",
"api_key",
"apikey",
"sessiontoken",
"firebase_id_token",
"idtoken",
"authorization",
"password",
"secret",
"devin-session-token$",
]
.iter()
.any(|marker| lowered.contains(marker))
|| value.contains("sk-")
"[REDACTED upstream error body]".to_string()
}
pub(crate) async fn refresh_windsurf_provider_quota_locally(
@@ -293,14 +230,13 @@ pub(crate) async fn refresh_windsurf_provider_quota_locally(
.await?
{
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
ProviderQuotaExecutionOutcome::Failure(_) => {
failed_count += 1;
let detail = sanitize_windsurf_probe_detail(detail);
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": format!("GetUserStatus 请求执行失败: {detail}"),
"message": "GetUserStatus 请求执行失败",
"status_code": 502,
}));
continue;
@@ -353,22 +289,18 @@ pub(crate) async fn refresh_windsurf_provider_quota_locally(
})
}
ProviderQuotaExecutionOutcome::Response(model_result) => {
let detail = extract_execution_error_message(&model_result)
.unwrap_or_else(|| format!("HTTP {}", model_result.status_code));
let detail = sanitize_windsurf_probe_detail(detail);
append_windsurf_probe_warning(
&mut metadata,
"model_configs",
format!("GetCascadeModelConfigs 返回: {detail}"),
"response_failed",
);
None
}
ProviderQuotaExecutionOutcome::Failure(detail) => {
let detail = sanitize_windsurf_probe_detail(detail);
ProviderQuotaExecutionOutcome::Failure(_) => {
append_windsurf_probe_warning(
&mut metadata,
"model_configs",
format!("GetCascadeModelConfigs 执行失败: {detail}"),
"execution_failed",
);
None
}
@@ -396,22 +328,18 @@ pub(crate) async fn refresh_windsurf_provider_quota_locally(
})
}
ProviderQuotaExecutionOutcome::Response(rate_limit_result) => {
let detail = extract_execution_error_message(&rate_limit_result)
.unwrap_or_else(|| format!("HTTP {}", rate_limit_result.status_code));
let detail = sanitize_windsurf_probe_detail(detail);
append_windsurf_probe_warning(
&mut metadata,
"rate_limit",
format!("CheckUserMessageRateLimit 返回: {detail}"),
"response_failed",
);
None
}
ProviderQuotaExecutionOutcome::Failure(detail) => {
let detail = sanitize_windsurf_probe_detail(detail);
ProviderQuotaExecutionOutcome::Failure(_) => {
append_windsurf_probe_warning(
&mut metadata,
"rate_limit",
format!("CheckUserMessageRateLimit 执行失败: {detail}"),
"execution_failed",
);
None
}
@@ -457,7 +385,7 @@ pub(crate) async fn refresh_windsurf_provider_quota_locally(
401 | 403 => {
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason =
Some(format!("Windsurf token 无效或已被拒绝: {}", detail));
Some("Windsurf token is invalid or rejected".to_string());
metadata.insert("banned".to_string(), json!(result.status_code == 403));
status = if result.status_code == 401 {
"auth_invalid".to_string()
@@ -470,10 +398,6 @@ pub(crate) async fn refresh_windsurf_provider_quota_locally(
"rate_limit".to_string(),
json!({
"limited": true,
"message": metadata
.get("last_error")
.cloned()
.unwrap_or_else(|| json!("rate limited")),
}),
);
status = "rate_limited".to_string();
@@ -525,9 +449,11 @@ pub(crate) async fn refresh_windsurf_provider_quota_locally(
if let Some(metadata) = metadata_update
.as_ref()
.and_then(|value| value.get("windsurf"))
.cloned()
{
payload.insert("metadata".to_string(), metadata);
payload.insert(
"metadata".to_string(),
admin_provider_metadata_bucket_safe_json("windsurf", Some(metadata)),
);
}
if let Some(quota_snapshot) = build_quota_snapshot_payload(
"windsurf",
@@ -592,7 +518,7 @@ mod tests {
r#"{"error":{"message":"bad"},"apiKey":"sk-secret","sessionToken":"devin-session-token$secret"}"#,
);
assert!(detail.contains("[REDACTED]"));
assert_eq!(detail, "[REDACTED upstream error body]");
assert!(!detail.contains("sk-secret"));
assert!(!detail.contains("devin-session-token$secret"));
}
@@ -4,8 +4,9 @@ use super::super::errors::{
use crate::handlers::admin::request::{AdminAppState, AdminProviderOAuthTemplate};
use aether_contracts::ProxySnapshot;
use aether_oauth::provider::providers::{
ClaudeCodeProviderOAuthAdapter, GenericProviderOAuthAdapter, CLAUDE_CODE_PROVIDER_TYPE,
CLAUDE_CODE_TOKEN_URL, CLAUDE_CODE_WEB_BASE_URL,
AntigravityProviderOAuthAdapter, ClaudeCodeProviderOAuthAdapter, GenericProviderOAuthAdapter,
ANTIGRAVITY_USER_INFO_URL, CLAUDE_CODE_PROVIDER_TYPE, CLAUDE_CODE_TOKEN_URL,
CLAUDE_CODE_WEB_BASE_URL,
};
use aether_oauth::provider::{
ProviderOAuthCookieAuthorizationInput, ProviderOAuthService, ProviderOAuthTransportContext,
@@ -13,12 +14,8 @@ use aether_oauth::provider::{
use axum::{body::Body, http, response::Response};
use std::sync::Arc;
fn provider_oauth_transport_error_detail(prefix: &str, error: &str) -> String {
let error = error.trim();
if error.is_empty() {
return prefix.to_string();
}
format!("{prefix}: {error}")
fn provider_oauth_transport_error_detail(prefix: &str, _error: &str) -> String {
prefix.to_string()
}
fn provider_oauth_exchange_context(
@@ -43,7 +40,19 @@ fn provider_oauth_exchange_context(
fn provider_oauth_service_for_template(
template: AdminProviderOAuthTemplate,
token_url: String,
antigravity_user_info_url: String,
) -> Result<ProviderOAuthService, Response<Body>> {
if template.provider_type.eq_ignore_ascii_case("antigravity") {
let adapter = AntigravityProviderOAuthAdapter::default()
.with_token_url_override(token_url)
.with_user_info_url_override(antigravity_user_info_url);
#[cfg(test)]
let adapter = adapter.with_oauth_credentials_for_tests(
"gateway-test-antigravity-client-id",
"gateway-test-antigravity-client-secret",
);
return Ok(ProviderOAuthService::new().with_adapter(Arc::new(adapter)));
}
GenericProviderOAuthAdapter::for_provider_type(template.provider_type)
.map(|adapter| adapter.with_token_url_override(token_url))
.map(|adapter| ProviderOAuthService::new().with_adapter(Arc::new(adapter)))
@@ -75,7 +84,10 @@ pub(crate) async fn exchange_admin_provider_oauth_code(
proxy: Option<ProxySnapshot>,
) -> Result<serde_json::Value, Response<Body>> {
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
let service = provider_oauth_service_for_template(template, token_url)?;
let antigravity_user_info_url =
state.provider_oauth_token_url("antigravity_user_info", ANTIGRAVITY_USER_INFO_URL);
let service =
provider_oauth_service_for_template(template, token_url, antigravity_user_info_url)?;
let ctx = provider_oauth_exchange_context(template.provider_type, proxy);
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
let result = service
@@ -103,7 +115,10 @@ pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
proxy: Option<ProxySnapshot>,
) -> Result<serde_json::Value, Response<Body>> {
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
let service = provider_oauth_service_for_template(template, token_url)?;
let antigravity_user_info_url =
state.provider_oauth_token_url("antigravity_user_info", ANTIGRAVITY_USER_INFO_URL);
let service =
provider_oauth_service_for_template(template, token_url, antigravity_user_info_url)?;
let ctx = provider_oauth_exchange_context(template.provider_type, proxy);
let executor = crate::oauth::GatewayOAuthHttpExecutor::new(*state);
let input = aether_oauth::provider::ProviderOAuthImportInput {
@@ -182,3 +197,20 @@ pub(crate) async fn authorize_admin_provider_oauth_with_cookie(
)
})
}
#[cfg(test)]
mod tests {
use super::provider_oauth_transport_error_detail;
#[test]
fn provider_oauth_transport_error_does_not_reflect_network_details() {
let detail = provider_oauth_transport_error_detail(
"token exchange 失败",
"request failed for https://user:[email protected]/token?secret=value authorization=Bearer upstream-secret",
);
assert_eq!(detail, "token exchange 失败");
assert!(!detail.contains("upstream-secret"));
assert!(!detail.contains("user:pass"));
}
}
@@ -1,31 +1,29 @@
use crate::handlers::admin::request::AdminProviderOAuthTemplate;
use aether_oauth::core::OAuthError;
use aether_oauth::provider::{ProviderOAuthService, ProviderOAuthTransportContext};
use serde_json::json;
use url::form_urlencoded;
pub(crate) fn build_provider_oauth_start_response(
template: AdminProviderOAuthTemplate,
nonce: &str,
code_challenge: Option<&str>,
) -> serde_json::Value {
let authorization_url = build_provider_oauth_authorization_url(template, nonce, code_challenge)
.unwrap_or_else(|| {
build_provider_oauth_authorization_url_legacy(template, nonce, code_challenge)
});
) -> Result<serde_json::Value, OAuthError> {
let authorization_url =
build_provider_oauth_authorization_url(template, nonce, code_challenge)?;
json!({
Ok(json!({
"authorization_url": authorization_url,
"redirect_uri": template.redirect_uri,
"provider_type": template.provider_type,
"instructions": "1) 打开 authorization_url 完成授权\n2) 复制授权页面显示的授权码或浏览器中的完整回调 URL\n3) 调用 complete 接口粘贴 callback_url",
})
}))
}
fn build_provider_oauth_authorization_url(
template: AdminProviderOAuthTemplate,
nonce: &str,
code_challenge: Option<&str>,
) -> Option<String> {
) -> Result<String, OAuthError> {
let ctx = ProviderOAuthTransportContext {
provider_id: String::new(),
provider_type: template.provider_type.to_string(),
@@ -41,32 +39,5 @@ fn build_provider_oauth_authorization_url(
};
ProviderOAuthService::with_builtin_adapters()
.build_authorize_url(&ctx, nonce, code_challenge)
.ok()
.map(|response| response.authorize_url)
}
fn build_provider_oauth_authorization_url_legacy(
template: AdminProviderOAuthTemplate,
nonce: &str,
code_challenge: Option<&str>,
) -> String {
let mut serializer = form_urlencoded::Serializer::new(String::new());
serializer.append_pair("client_id", template.client_id);
serializer.append_pair("response_type", "code");
serializer.append_pair("redirect_uri", template.redirect_uri);
serializer.append_pair("scope", &template.scopes.join(" "));
serializer.append_pair("state", nonce);
if template.provider_type == "codex" {
serializer.append_pair("prompt", "login");
serializer.append_pair("id_token_add_organizations", "true");
serializer.append_pair("codex_cli_simplified_flow", "true");
}
if template.use_pkce {
if let Some(code_challenge) = code_challenge {
serializer.append_pair("code_challenge", code_challenge);
serializer.append_pair("code_challenge_method", "S256");
}
}
format!("{}?{}", template.authorize_url, serializer.finish())
}
@@ -24,11 +24,13 @@ pub(in super::super) async fn admin_provider_ops_probe_new_api_checkin(
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or("/api/user/checkin");
let url = admin_provider_ops_request_url(
let Ok(url) = admin_provider_ops_request_url(
base_url,
&admin_provider_ops_json_object_map(json!({ "endpoint": endpoint })),
endpoint,
);
) else {
return None;
};
let (status, response_json) = match admin_provider_ops_execute_json_request(
state,
"provider-ops-action:probe_checkin",
@@ -58,8 +60,15 @@ pub(in super::super) async fn admin_provider_ops_probe_new_api_checkin(
cookie_expired: true,
});
}
if status != http::StatusCode::OK {
return Some(AdminProviderOpsCheckinOutcome {
success: Some(false),
message: "签到失败".to_string(),
cookie_expired: false,
});
}
let message = response_json
let upstream_message = response_json
.get("message")
.and_then(serde_json::Value::as_str)
.unwrap_or_default()
@@ -72,43 +81,27 @@ pub(in super::super) async fn admin_provider_ops_probe_new_api_checkin(
{
return Some(AdminProviderOpsCheckinOutcome {
success: Some(true),
message: if message.is_empty() {
"签到成功".to_string()
} else {
message
},
message: "签到成功".to_string(),
cookie_expired: false,
});
}
if admin_provider_ops_checkin_already_done(&message) {
if admin_provider_ops_checkin_already_done(&upstream_message) {
return Some(AdminProviderOpsCheckinOutcome {
success: None,
message: if message.is_empty() {
"今日已签到".to_string()
} else {
message
},
message: "今日已签到".to_string(),
cookie_expired: false,
});
}
if admin_provider_ops_checkin_auth_failure(&message) {
if admin_provider_ops_checkin_auth_failure(&upstream_message) {
return has_cookie.then(|| AdminProviderOpsCheckinOutcome {
success: None,
message: if message.is_empty() {
"Cookie 已失效".to_string()
} else {
message
},
message: "Cookie 已失效".to_string(),
cookie_expired: true,
});
}
Some(AdminProviderOpsCheckinOutcome {
success: Some(false),
message: if message.is_empty() {
"签到失败".to_string()
} else {
message
},
message: "签到失败".to_string(),
cookie_expired: false,
})
}
@@ -32,7 +32,12 @@ pub(in super::super) async fn admin_provider_ops_run_checkin_action(
);
}
let url = admin_provider_ops_request_url(base_url, action_config, "/api/user/checkin");
let url = match admin_provider_ops_request_url(base_url, action_config, "/api/user/checkin") {
Ok(url) => url,
Err(message) => {
return admin_provider_ops_action_error("not_configured", "checkin", message, None)
}
};
let method = admin_provider_ops_request_method(action_config, "POST");
let (status, response_json) = match admin_provider_ops_execute_json_request(
state,
@@ -129,7 +134,7 @@ pub(in super::super) async fn admin_provider_ops_run_checkin_action(
);
}
let message = response_json
let upstream_message = response_json
.get("message")
.and_then(serde_json::Value::as_str)
.unwrap_or_default()
@@ -143,23 +148,23 @@ pub(in super::super) async fn admin_provider_ops_run_checkin_action(
return admin_provider_ops_action_response(
"success",
"checkin",
admin_provider_ops_checkin_payload(&response_json, Some(message)),
admin_provider_ops_checkin_payload(&response_json, Some("签到成功".to_string())),
None,
response_time_ms,
3600,
);
}
if admin_provider_ops_checkin_already_done(&message) {
if admin_provider_ops_checkin_already_done(&upstream_message) {
return admin_provider_ops_action_response(
"already_done",
"checkin",
admin_provider_ops_checkin_payload(&response_json, Some(message)),
admin_provider_ops_checkin_payload(&response_json, Some("今日已签到".to_string())),
None,
response_time_ms,
3600,
);
}
if admin_provider_ops_checkin_auth_failure(&message) {
if admin_provider_ops_checkin_auth_failure(&upstream_message) {
return admin_provider_ops_action_error(
if has_cookie {
"auth_expired"
@@ -167,28 +172,15 @@ pub(in super::super) async fn admin_provider_ops_run_checkin_action(
"auth_failed"
},
"checkin",
if message.is_empty() {
if has_cookie {
"Cookie 已失效"
} else {
"认证失败"
}
if has_cookie {
"Cookie 已失效"
} else {
message.as_str()
"认证失败"
},
response_time_ms,
);
}
admin_provider_ops_action_error(
"unknown_error",
"checkin",
if message.is_empty() {
"签到失败"
} else {
message.as_str()
},
response_time_ms,
)
admin_provider_ops_action_error("unknown_error", "checkin", "签到失败", response_time_ms)
}
fn admin_provider_ops_network_error_message(error: &str) -> String {
@@ -197,5 +189,5 @@ fn admin_provider_ops_network_error_message(error: &str) -> String {
if lower.contains("timeout") || normalized.contains("超时") {
return "请求超时".to_string();
}
format!("网络错误: {normalized}")
"网络错误".to_string()
}
@@ -34,7 +34,7 @@ pub(super) fn admin_provider_ops_checkin_auth_failure(message: &str) -> bool {
pub(super) fn admin_provider_ops_checkin_payload(
response_json: &serde_json::Value,
fallback_message: Option<String>,
message: Option<String>,
) -> serde_json::Value {
let details = response_json
.get("data")
@@ -54,30 +54,45 @@ pub(super) fn admin_provider_ops_checkin_payload(
let next_reward = details.and_then(|value| {
admin_provider_ops_value_as_f64(value.get("next_reward").or_else(|| value.get("next")))
});
let message = fallback_message.or_else(|| {
response_json
.get("message")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned)
});
let mut extra = serde_json::Map::new();
if let Some(details) = details {
for (key, value) in details {
if matches!(
key.as_str(),
"reward"
| "quota"
| "amount"
| "streak_days"
| "streak"
| "next_reward"
| "next"
| "message"
) {
continue;
}
extra.insert(key.clone(), value.clone());
}
admin_provider_ops_checkin_data(
reward,
streak_days,
next_reward,
message,
serde_json::Map::new(),
)
}
#[cfg(test)]
mod tests {
use super::admin_provider_ops_checkin_payload;
use serde_json::json;
#[test]
fn checkin_payload_keeps_metrics_without_copying_upstream_secrets() {
let payload = admin_provider_ops_checkin_payload(
&json!({
"success": true,
"message": "authorization=Bearer upstream-secret",
"data": {
"reward": 1.5,
"streak_days": 3,
"next_reward": 2.0,
"api_key": "secret-api-key",
"profile": {"access_token": "secret-token"}
}
}),
Some("签到成功".to_string()),
);
assert_eq!(payload["reward"], json!(1.5));
assert_eq!(payload["streak_days"], json!(3));
assert_eq!(payload["next_reward"], json!(2.0));
assert_eq!(payload["message"], json!("签到成功"));
assert_eq!(payload["extra"], json!({}));
let serialized = payload.to_string();
assert!(!serialized.contains("upstream-secret"));
assert!(!serialized.contains("secret-api-key"));
assert!(!serialized.contains("secret-token"));
}
admin_provider_ops_checkin_data(reward, streak_days, next_reward, message, extra)
}
@@ -5,7 +5,7 @@ mod support;
use super::config::{
admin_provider_ops_config_object, admin_provider_ops_connector_object,
admin_provider_ops_decrypted_credentials, resolve_admin_provider_ops_base_url,
admin_provider_ops_credential_snapshot,
};
use super::support::ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE;
use super::verify::{
@@ -36,13 +36,23 @@ pub(crate) async fn admin_provider_ops_local_action_response(
state: &AdminAppState<'_>,
provider_id: &str,
provider: Option<&StoredProviderCatalogProvider>,
endpoints: &[StoredProviderCatalogEndpoint],
_endpoints: &[StoredProviderCatalogEndpoint],
action_type: &str,
request_config: Option<&serde_json::Map<String, serde_json::Value>>,
) -> serde_json::Value {
let Some(provider) = provider else {
return responses::admin_provider_ops_action_not_configured(action_type, "未配置操作设置");
};
let credential_snapshot = match admin_provider_ops_credential_snapshot(state, provider).await {
Ok(snapshot) => snapshot,
Err(_) => {
return responses::admin_provider_ops_action_not_configured(
action_type,
"已保存的 Provider Ops 凭据无法解密或迁移",
)
}
};
let provider = &credential_snapshot.provider;
let Some(provider_ops_config) = admin_provider_ops_config_object(provider) else {
return responses::admin_provider_ops_action_not_configured(action_type, "未配置操作设置");
};
@@ -57,14 +67,11 @@ pub(crate) async fn admin_provider_ops_local_action_response(
ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE,
);
};
let Some(base_url) =
resolve_admin_provider_ops_base_url(provider, endpoints, Some(provider_ops_config))
else {
return responses::admin_provider_ops_action_not_configured(
action_type,
"Provider 未配置 base_url",
);
};
let base_url = credential_snapshot
.binding
.destination
.base_url()
.to_string();
let mut connector_config = admin_provider_ops_connector_object(provider_ops_config)
.and_then(|connector| connector.get("config"))
@@ -84,12 +91,7 @@ pub(crate) async fn admin_provider_ops_local_action_response(
let proxy_snapshot =
admin_provider_ops_resolve_proxy_snapshot(state, Some(&connector_config)).await;
let credentials = admin_provider_ops_decrypted_credentials(
state,
admin_provider_ops_config_object(provider)
.and_then(admin_provider_ops_connector_object)
.and_then(|connector| connector.get("credentials")),
);
let credentials = credential_snapshot.credentials;
let headers = match build_headers(
architecture.architecture_id,
&connector_config,
@@ -73,7 +73,17 @@ pub(super) async fn admin_provider_ops_run_query_balance_action(
}
let start = std::time::Instant::now();
let url = admin_provider_ops_request_url(base_url, action_config, "/api/user/balance");
let url = match admin_provider_ops_request_url(base_url, action_config, "/api/user/balance") {
Ok(url) => url,
Err(message) => {
return admin_provider_ops_action_error(
"not_configured",
"query_balance",
message,
None,
)
}
};
let method = admin_provider_ops_request_method(action_config, "GET");
let (status, response_json) = match admin_provider_ops_execute_json_request(
state,
@@ -200,5 +210,5 @@ fn admin_provider_ops_network_error_message(error: &str) -> String {
if lower.contains("timeout") || normalized.contains("超时") {
return "请求超时".to_string();
}
format!("网络错误: {normalized}")
"网络错误".to_string()
}
@@ -63,7 +63,17 @@ pub(super) async fn admin_provider_ops_sub2api_balance_payload(
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or("/api/v1/auth/me?timezone=Asia/Shanghai");
let me_url = admin_provider_ops_sub2api_request_url(base_url, me_endpoint);
let me_url = match admin_provider_ops_sub2api_request_url(base_url, me_endpoint) {
Ok(url) => url,
Err(message) => {
return admin_provider_ops_action_error(
"not_configured",
"query_balance",
message,
None,
)
}
};
let subscription_endpoint = admin_provider_ops_json_object_map(json!({
"endpoint": action_config
.get("subscription_endpoint")
@@ -77,7 +87,17 @@ pub(super) async fn admin_provider_ops_sub2api_balance_payload(
.unwrap_or("/api/v1/subscriptions/summary")
.to_string();
let subscription_url =
admin_provider_ops_sub2api_request_url(base_url, subscription_endpoint.as_str());
match admin_provider_ops_sub2api_request_url(base_url, subscription_endpoint.as_str()) {
Ok(url) => url,
Err(message) => {
return admin_provider_ops_action_error(
"not_configured",
"query_balance",
message,
None,
)
}
};
let auth_value = match reqwest::header::HeaderValue::from_str(&format!("Bearer {access_token}"))
{
@@ -197,8 +217,5 @@ fn network_error_message(error: &str) -> String {
if lower.contains("timeout") || normalized.contains("超时") {
return "请求超时".to_string();
}
if normalized.starts_with("网络错误:") {
return normalized.to_string();
}
format!("网络错误: {normalized}")
"网络错误".to_string()
}
@@ -5,6 +5,9 @@ use super::super::responses::{
admin_provider_ops_action_error, admin_provider_ops_action_response,
};
use crate::handlers::admin::request::AdminAppState;
use crate::handlers::shared::{
canonicalize_provider_ops_base_url, resolve_provider_ops_same_origin_url,
};
use aether_admin::provider::ops::parse_yescode_combined_balance_payload;
use aether_contracts::ProxySnapshot;
use serde_json::json;
@@ -17,8 +20,44 @@ pub(super) async fn admin_provider_ops_yescode_balance_payload(
proxy_snapshot: Option<&ProxySnapshot>,
) -> serde_json::Value {
let start = std::time::Instant::now();
let balance_url = format!("{}/api/v1/user/balance", base_url.trim_end_matches('/'));
let profile_url = format!("{}/api/v1/auth/profile", base_url.trim_end_matches('/'));
// Resolve fixed action paths against a canonical origin. Avoid string
// concatenation so malformed bases (or path-like inputs) cannot redirect
// credentials to another host.
let destination = match canonicalize_provider_ops_base_url(base_url) {
Ok(destination) => destination,
Err(_) => {
return admin_provider_ops_action_error(
"auth_failed",
"query_balance",
"Cookie 已失效,请重新配置",
Some(start.elapsed().as_millis() as u64),
);
}
};
let balance_url =
match resolve_provider_ops_same_origin_url(&destination, "/api/v1/user/balance") {
Ok(url) => url,
Err(_) => {
return admin_provider_ops_action_error(
"auth_failed",
"query_balance",
"Cookie 已失效,请重新配置",
Some(start.elapsed().as_millis() as u64),
);
}
};
let profile_url =
match resolve_provider_ops_same_origin_url(&destination, "/api/v1/auth/profile") {
Ok(url) => url,
Err(_) => {
return admin_provider_ops_action_error(
"auth_failed",
"query_balance",
"Cookie 已失效,请重新配置",
Some(start.elapsed().as_millis() as u64),
);
}
};
let (balance_result, profile_result) = tokio::join!(
admin_provider_ops_execute_json_request(
state,
@@ -1,5 +1,9 @@
use serde_json::json;
use crate::handlers::shared::{
canonicalize_provider_ops_base_url, resolve_provider_ops_same_origin_url,
};
pub(super) fn admin_provider_ops_checkin_data(
reward: Option<f64>,
streak_days: Option<i64>,
@@ -26,18 +30,15 @@ pub(super) fn admin_provider_ops_request_url(
base_url: &str,
action_config: &serde_json::Map<String, serde_json::Value>,
default_endpoint: &str,
) -> String {
) -> Result<String, String> {
let endpoint = action_config
.get("endpoint")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(default_endpoint);
if endpoint.starts_with("http://") || endpoint.starts_with("https://") {
endpoint.to_string()
} else {
format!("{}{}", base_url.trim_end_matches('/'), endpoint)
}
let destination = canonicalize_provider_ops_base_url(base_url).map_err(ToString::to_string)?;
resolve_provider_ops_same_origin_url(&destination, endpoint).map_err(ToString::to_string)
}
pub(super) fn admin_provider_ops_request_method(
@@ -1,7 +1,7 @@
use super::actions::admin_provider_ops_local_action_response;
use crate::handlers::admin::request::AdminAppState;
use crate::task_runtime::{spawn_fire_and_forget, TASK_KEY_PROVIDER_BALANCE_REFRESH};
use serde_json::{json, Value};
use serde_json::{json, Map, Value};
use std::collections::HashSet;
use std::time::Duration;
use tokio::sync::{Mutex, Semaphore};
@@ -75,10 +75,13 @@ pub(crate) async fn store_admin_provider_ops_balance_cache(
provider_id: &str,
payload: &Value,
) {
let Some(ttl_seconds) = balance_cache_ttl_seconds(payload) else {
let Some(projected) = project_admin_provider_ops_balance_cache_payload(payload) else {
return;
};
let serialized = match serde_json::to_string(payload) {
let Some(ttl_seconds) = balance_cache_ttl_seconds(&projected) else {
return;
};
let serialized = match serde_json::to_string(&projected) {
Ok(serialized) => serialized,
Err(err) => {
warn!(
@@ -222,6 +225,284 @@ fn balance_cache_ttl_seconds(payload: &Value) -> Option<u64> {
}
}
const BALANCE_CACHE_EXTRA_NUMERIC_FIELDS: &[&str] = &[
"balance",
"points",
"active_subscriptions",
"total_used_usd",
"normal_balance",
"subscription_balance",
"charity_balance",
"pay_as_you_go_balance",
"daily_limit",
"weekly_limit",
"weekly_spent",
"daily_spent",
"daily_used_quota",
"daily_quota_limit",
"daily_remaining_quota",
];
const BALANCE_CACHE_EXTRA_STRING_FIELDS: &[&str] = &[
"plan_name",
"subscription_status",
"status",
"group_name",
"effective_start_date",
"effective_end_date",
];
const BALANCE_CACHE_EXTRA_BOOL_FIELDS: &[&str] = &["checkin_success", "cookie_expired"];
const BALANCE_CACHE_EXTRA_NESTED_FIELDS: &[&str] = &[
"five_hour_limit",
"weekly_limit",
"month_stats",
"subscriptions",
];
const BALANCE_CACHE_LIMIT_FIELDS: &[&str] = &["limit", "used", "remaining", "resets_at"];
const BALANCE_CACHE_MONTH_STATS_FIELDS: &[&str] = &[
"total_input_tokens",
"total_output_tokens",
"total_quota",
"total_requests",
];
fn project_admin_provider_ops_balance_cache_payload(payload: &Value) -> Option<Value> {
let source = payload.as_object()?;
let status = source.get("status").and_then(Value::as_str)?.trim();
if !matches!(status, "success" | "auth_expired" | "auth_failed") {
return None;
}
if source.get("action_type").and_then(Value::as_str) != Some("query_balance") {
return None;
}
let mut projected = Map::new();
projected.insert("status".to_string(), Value::String(status.to_string()));
projected.insert(
"action_type".to_string(),
Value::String("query_balance".to_string()),
);
let data = match source.get("data") {
Some(Value::Null) | None => Value::Null,
Some(value) => project_admin_provider_ops_balance_data(value)?,
};
projected.insert("data".to_string(), data);
projected.insert(
"message".to_string(),
match status {
"auth_failed" => Value::String("认证失败".to_string()),
"auth_expired" => Value::String("认证已过期".to_string()),
_ => Value::Null,
},
);
if let Some(value) = source
.get("executed_at")
.and_then(project_admin_provider_ops_safe_string)
{
projected.insert("executed_at".to_string(), Value::String(value));
}
if let Some(value) = source
.get("response_time_ms")
.and_then(project_admin_provider_ops_finite_number)
{
projected.insert("response_time_ms".to_string(), value);
}
projected.insert(
"cache_ttl_seconds".to_string(),
Value::from(if status == "auth_failed" {
ADMIN_PROVIDER_OPS_BALANCE_AUTH_FAILED_CACHE_TTL_SECS
} else {
ADMIN_PROVIDER_OPS_BALANCE_CACHE_TTL_SECS
}),
);
Some(Value::Object(projected))
}
fn project_admin_provider_ops_balance_data(value: &Value) -> Option<Value> {
let source = value.as_object()?;
let mut projected = Map::new();
for field in ["total_granted", "total_used", "total_available"] {
if let Some(value) = source.get(field) {
projected.insert(
field.to_string(),
project_admin_provider_ops_finite_number_or_null(value)?,
);
}
}
if let Some(value) = source.get("expires_at") {
projected.insert(
"expires_at".to_string(),
project_admin_provider_ops_finite_number_or_null(value)?,
);
}
if let Some(value) = source.get("currency") {
let currency = project_admin_provider_ops_safe_string(value)?;
if currency.len() > 32
|| !currency
.chars()
.all(|ch| ch.is_ascii_alphanumeric() || matches!(ch, '_' | '-' | '.' | '/'))
{
return None;
}
projected.insert("currency".to_string(), Value::String(currency));
}
if let Some(extra) = source.get("extra") {
projected.insert(
"extra".to_string(),
project_admin_provider_ops_balance_extra(extra)?,
);
}
Some(Value::Object(projected))
}
fn project_admin_provider_ops_balance_extra(value: &Value) -> Option<Value> {
let source = value.as_object()?;
let mut projected = Map::new();
for (field, value) in source {
let projected_value = if BALANCE_CACHE_EXTRA_NUMERIC_FIELDS.contains(&field.as_str())
&& project_admin_provider_ops_finite_number(value).is_some()
{
project_admin_provider_ops_finite_number(value)
} else if BALANCE_CACHE_EXTRA_STRING_FIELDS.contains(&field.as_str()) {
project_admin_provider_ops_safe_string(value).map(Value::String)
} else if BALANCE_CACHE_EXTRA_BOOL_FIELDS.contains(&field.as_str()) {
value.as_bool().map(Value::Bool)
} else if BALANCE_CACHE_EXTRA_NESTED_FIELDS.contains(&field.as_str()) {
project_admin_provider_ops_balance_extra_nested(field, value)
} else if matches!(
field.as_str(),
"weekly_resets_at" | "daily_resets_at" | "resets_at"
) {
project_admin_provider_ops_finite_number_or_safe_string(value)
} else if matches!(field.as_str(), "checkin_message" | "cookie_expired_message") {
project_admin_provider_ops_safe_string(value).map(Value::String)
} else {
None
};
if let Some(projected_value) = projected_value {
projected.insert(field.clone(), projected_value);
}
}
Some(Value::Object(projected))
}
fn project_admin_provider_ops_balance_extra_nested(field: &str, value: &Value) -> Option<Value> {
if field == "subscriptions" {
let items = value.as_array()?;
return Some(Value::Array(
items
.iter()
.take(128)
.filter_map(project_admin_provider_ops_subscription)
.collect(),
));
}
let source = value.as_object()?;
let mut projected = Map::new();
let allowed = if field == "month_stats" {
BALANCE_CACHE_MONTH_STATS_FIELDS
} else {
BALANCE_CACHE_LIMIT_FIELDS
};
for key in allowed {
if let Some(value) = source.get(*key) {
let projected_value = if *key == "resets_at" {
project_admin_provider_ops_finite_number_or_safe_string(value)
} else {
project_admin_provider_ops_finite_number(value)
};
if let Some(projected_value) = projected_value {
projected.insert((*key).to_string(), projected_value);
}
}
}
Some(Value::Object(projected))
}
fn project_admin_provider_ops_subscription(value: &Value) -> Option<Value> {
let source = value.as_object()?;
let mut projected = Map::new();
for field in ["group_name", "status"] {
if let Some(value) = source.get(field) {
projected.insert(
field.to_string(),
Value::String(project_admin_provider_ops_safe_string(value)?),
);
}
}
for field in [
"daily_used_usd",
"daily_limit_usd",
"weekly_used_usd",
"weekly_limit_usd",
"monthly_used_usd",
"monthly_limit_usd",
] {
if let Some(value) = source.get(field) {
if let Some(value) = project_admin_provider_ops_finite_number(value) {
projected.insert(field.to_string(), value);
}
}
}
if let Some(value) = source.get("expires_at") {
if let Some(value) = project_admin_provider_ops_finite_number_or_safe_string(value) {
projected.insert("expires_at".to_string(), value);
}
}
Some(Value::Object(projected))
}
fn project_admin_provider_ops_finite_number(value: &Value) -> Option<Value> {
if let Some(number) = value.as_f64() {
return number.is_finite().then(|| value.clone());
}
let number = value.as_str()?.trim().parse::<f64>().ok()?;
number.is_finite().then(|| Value::from(number))
}
fn project_admin_provider_ops_finite_number_or_null(value: &Value) -> Option<Value> {
if value.is_null() {
Some(Value::Null)
} else {
project_admin_provider_ops_finite_number(value)
}
}
fn project_admin_provider_ops_finite_number_or_safe_string(value: &Value) -> Option<Value> {
project_admin_provider_ops_finite_number(value)
.or_else(|| project_admin_provider_ops_safe_string(value).map(Value::String))
}
fn project_admin_provider_ops_safe_string(value: &Value) -> Option<String> {
let value = value.as_str()?.trim();
if value.is_empty() || value.len() > 256 || value.chars().any(char::is_control) {
return None;
}
let lower = value.to_ascii_lowercase();
if [
"authorization",
"bearer ",
"api_key",
"apikey",
"access_token",
"refresh_token",
"password",
"cookie",
"session",
"secret",
"token=",
]
.iter()
.any(|needle| lower.contains(needle))
{
return None;
}
Some(value.to_string())
}
fn admin_provider_ops_balance_refresh_key(state: &AdminAppState<'_>, provider_id: &str) -> String {
let raw_key = format!("{ADMIN_PROVIDER_OPS_BALANCE_REFRESH_PREFIX}{provider_id}");
format!(
@@ -263,7 +544,10 @@ fn admin_provider_ops_action_response(
#[cfg(test)]
mod tests {
use super::{admin_provider_ops_pending_balance_response, balance_cache_ttl_seconds};
use super::{
admin_provider_ops_pending_balance_response, balance_cache_ttl_seconds,
project_admin_provider_ops_balance_cache_payload,
};
use serde_json::json;
#[test]
@@ -292,4 +576,65 @@ mod tests {
None
);
}
#[test]
fn balance_cache_projection_drops_untrusted_messages_and_fields() {
let payload = json!({
"status": "auth_failed",
"action_type": "query_balance",
"message": "authorization=Bearer upstream-secret",
"data": {
"total_available": 1.25,
"currency": "USD",
"extra": {
"balance": 1.0,
"access_token": "upstream-secret",
"today_stats": {"private_note": "upstream-secret"},
"checkin_message": "签到失败"
}
},
"cache_ttl_seconds": 999999
});
let projected = project_admin_provider_ops_balance_cache_payload(&payload)
.expect("known balance payload should project");
assert_eq!(projected["message"], json!("认证失败"));
assert_eq!(projected["data"]["extra"]["balance"], json!(1.0));
assert!(projected.to_string().find("upstream-secret").is_none());
assert!(projected["data"]["extra"].get("access_token").is_none());
assert!(projected["data"]["extra"].get("today_stats").is_none());
assert_eq!(projected["cache_ttl_seconds"], json!(60));
}
#[test]
fn balance_cache_projection_keeps_sub2api_subscription_allowlist() {
let payload = json!({
"status": "success",
"action_type": "query_balance",
"data": {
"total_available": 8.5,
"currency": "USD",
"extra": {
"subscriptions": [{
"group_name": "default",
"status": "active",
"monthly_used_usd": 1.2,
"private_token": "must-drop"
}]
}
}
});
let projected = project_admin_provider_ops_balance_cache_payload(&payload)
.expect("known balance payload should project");
assert_eq!(
projected["data"]["extra"]["subscriptions"][0]["group_name"],
json!("default")
);
assert_eq!(
projected["data"]["extra"]["subscriptions"][0]["monthly_used_usd"],
json!(1.2)
);
assert!(projected["data"]["extra"]["subscriptions"][0]
.get("private_token")
.is_none());
}
}
@@ -1,19 +1,58 @@
use super::support::{
AdminProviderOpsQuotaAlertConfigRequest, AdminProviderOpsSaveConfigRequest,
ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS,
};
use super::support::{AdminProviderOpsQuotaAlertConfigRequest, AdminProviderOpsSaveConfigRequest};
use crate::handlers::admin::request::AdminAppState;
use crate::handlers::shared::{
canonicalize_provider_ops_base_url, masked_secret_display, open_provider_ops_credential,
provider_ops_credential_binding_from_config, provider_ops_credential_field_is_secret,
seal_provider_ops_credential, ProviderOpsCredentialBinding,
PROVIDER_OPS_PERSISTENT_SECRET_FIELDS, PROVIDER_OPS_TRANSIENT_METADATA_FIELDS,
};
use crate::GatewayError;
use aether_admin::provider::ops as admin_provider_ops_pure;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
ProviderCatalogProviderConfigCasUpdate, StoredProviderCatalogEndpoint,
StoredProviderCatalogProvider,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
const PROVIDER_OPS_QUOTA_ALERT_DEFAULT_FETCH_INTERVAL_SECS: u64 = 30;
const PROVIDER_OPS_QUOTA_ALERT_MIN_FETCH_INTERVAL_SECS: u64 = 30;
const PROVIDER_OPS_QUOTA_ALERT_MAX_FETCH_INTERVAL_SECS: u64 = 86_400;
const PROVIDER_OPS_CREDENTIAL_MIGRATION_RETRIES: usize = 8;
struct AdminProviderOpsDecodedCredentials {
values: serde_json::Map<String, serde_json::Value>,
protected_values: serde_json::Map<String, serde_json::Value>,
migration_required: bool,
}
pub(crate) struct AdminProviderOpsCredentialSnapshot {
pub(crate) provider: StoredProviderCatalogProvider,
pub(crate) credentials: serde_json::Map<String, serde_json::Value>,
pub(crate) binding: ProviderOpsCredentialBinding,
}
impl std::fmt::Debug for AdminProviderOpsCredentialSnapshot {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("AdminProviderOpsCredentialSnapshot")
.field("provider_id", &self.provider.id)
.field("credentials", &"[REDACTED]")
.field("binding", &"[REDACTED]")
.finish_non_exhaustive()
}
}
pub(super) struct AdminProviderOpsMergedCredentialSnapshot {
pub(super) provider: StoredProviderCatalogProvider,
pub(super) credentials: serde_json::Map<String, serde_json::Value>,
pub(super) saved_binding: ProviderOpsCredentialBinding,
pub(super) reused_saved_secret: bool,
}
pub(super) struct AdminProviderOpsSavedConfigSnapshot {
pub(super) provider: StoredProviderCatalogProvider,
pub(super) provider_ops_config: serde_json::Value,
}
pub(super) fn admin_provider_ops_config_object(
provider: &StoredProviderCatalogProvider,
@@ -27,57 +66,77 @@ pub(super) fn admin_provider_ops_connector_object(
admin_provider_ops_pure::admin_provider_ops_connector_object(provider_ops_config)
}
fn admin_provider_ops_masked_secret(
pub(super) fn admin_provider_ops_binding_from_config(
provider_id: &str,
provider_ops_config: &serde_json::Map<String, serde_json::Value>,
effective_base_url: &str,
) -> Result<ProviderOpsCredentialBinding, String> {
provider_ops_credential_binding_from_config(
provider_id,
provider_ops_config,
effective_base_url,
)
.map_err(ToString::to_string)
}
async fn admin_provider_ops_binding_for_provider(
state: &AdminAppState<'_>,
field: &str,
ciphertext: &str,
) -> serde_json::Value {
let plaintext = state
.decrypt_catalog_secret_with_fallbacks(ciphertext)
.unwrap_or_else(|| ciphertext.to_string());
provider: &StoredProviderCatalogProvider,
) -> Result<(ProviderOpsCredentialBinding, bool), GatewayError> {
let provider_ops_config = admin_provider_ops_config_object(provider)
.ok_or_else(|| GatewayError::Internal("Provider Ops 配置格式无效".to_string()))?;
let explicit_base_url = provider_ops_config
.get("base_url")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty());
let endpoints = if explicit_base_url.is_some() {
Vec::new()
} else {
state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?
};
let effective_base_url =
resolve_admin_provider_ops_base_url(provider, &endpoints, Some(provider_ops_config))
.ok_or_else(|| GatewayError::Internal("Provider Ops 未配置 base_url".to_string()))?;
let binding = admin_provider_ops_binding_from_config(
&provider.id,
provider_ops_config,
&effective_base_url,
)
.map_err(GatewayError::Internal)?;
let needs_materialized_base_url = explicit_base_url != Some(binding.destination.base_url());
Ok((binding, needs_materialized_base_url))
}
fn admin_provider_ops_masked_secret(field: &str, plaintext: &str) -> serde_json::Value {
if plaintext.is_empty() {
return serde_json::Value::String(String::new());
}
let masked = if field == "password" {
"********".to_string()
} else if plaintext.len() > 12 {
format!(
"{}****{}",
&plaintext[..4],
&plaintext[plaintext.len().saturating_sub(4)..]
)
} else if plaintext.len() > 8 {
format!(
"{}****{}",
&plaintext[..2],
&plaintext[plaintext.len().saturating_sub(2)..]
)
} else {
"*".repeat(plaintext.len())
masked_secret_display(plaintext, 4, 4, "****")
};
serde_json::Value::String(masked)
}
fn admin_provider_ops_masked_credentials(
state: &AdminAppState<'_>,
raw_credentials: Option<&serde_json::Value>,
credentials: &serde_json::Map<String, serde_json::Value>,
) -> serde_json::Value {
let Some(credentials) = raw_credentials.and_then(serde_json::Value::as_object) else {
return json!({});
};
let mut masked = serde_json::Map::new();
for (key, value) in credentials {
if key.starts_with('_') {
continue;
}
if ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS.contains(&key.as_str()) {
if provider_ops_credential_field_is_secret(key) {
if let Some(ciphertext) = value.as_str().filter(|value| !value.is_empty()) {
masked.insert(
key.clone(),
admin_provider_ops_masked_secret(state, key, ciphertext),
admin_provider_ops_masked_secret(key, ciphertext),
);
continue;
}
@@ -91,53 +150,71 @@ fn admin_provider_ops_is_supported_auth_type(auth_type: &str) -> bool {
admin_provider_ops_pure::admin_provider_ops_is_supported_auth_type(auth_type)
}
pub(super) fn admin_provider_ops_decrypted_credentials(
fn admin_provider_ops_decode_credentials(
state: &AdminAppState<'_>,
binding: &ProviderOpsCredentialBinding,
raw_credentials: Option<&serde_json::Value>,
) -> serde_json::Map<String, serde_json::Value> {
) -> Result<AdminProviderOpsDecodedCredentials, String> {
let Some(credentials) = raw_credentials.and_then(serde_json::Value::as_object) else {
return serde_json::Map::new();
return Ok(AdminProviderOpsDecodedCredentials {
values: serde_json::Map::new(),
protected_values: serde_json::Map::new(),
migration_required: false,
});
};
let mut decrypted = serde_json::Map::new();
let mut values = serde_json::Map::new();
let mut protected_values = credentials.clone();
let mut migration_required = false;
for (key, value) in credentials {
if ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS.contains(&key.as_str()) {
if let Some(ciphertext) = value.as_str() {
let plaintext = state
.decrypt_catalog_secret_with_fallbacks(ciphertext)
.unwrap_or_else(|| ciphertext.to_string());
decrypted.insert(key.clone(), serde_json::Value::String(plaintext));
if provider_ops_credential_field_is_secret(key) {
if let Some(stored_value) = value.as_str() {
if stored_value.is_empty() {
values.insert(key.clone(), value.clone());
continue;
}
let projection =
open_provider_ops_credential(state.app(), binding, key, stored_value).map_err(
|message| format!("已保存的 Provider Ops 凭据无法解密: {message}"),
)?;
migration_required |= projection.migration_required;
protected_values
.insert(key.clone(), serde_json::Value::String(projection.protected));
values.insert(key.clone(), serde_json::Value::String(projection.plaintext));
continue;
}
}
decrypted.insert(key.clone(), value.clone());
values.insert(key.clone(), value.clone());
}
decrypted
Ok(AdminProviderOpsDecodedCredentials {
values,
protected_values,
migration_required,
})
}
fn admin_provider_ops_sensitive_placeholder_or_empty(value: Option<&serde_json::Value>) -> bool {
admin_provider_ops_pure::admin_provider_ops_sensitive_placeholder_or_empty(value)
}
pub(super) fn admin_provider_ops_merge_credentials(
pub(super) async fn admin_provider_ops_merge_credentials(
state: &AdminAppState<'_>,
architecture_id: &str,
provider: &StoredProviderCatalogProvider,
mut request_credentials: serde_json::Map<String, serde_json::Value>,
) -> serde_json::Map<String, serde_json::Value> {
let mut saved_credentials = admin_provider_ops_decrypted_credentials(
state,
admin_provider_ops_config_object(provider)
.and_then(admin_provider_ops_connector_object)
.and_then(|connector| connector.get("credentials")),
);
) -> Result<AdminProviderOpsMergedCredentialSnapshot, String> {
let snapshot = admin_provider_ops_credential_snapshot(state, provider)
.await
.map_err(|_| "已保存的 Provider Ops 凭据无法解密或迁移".to_string())?;
let mut saved_credentials = snapshot.credentials;
let preserve_internal_runtime_fields =
admin_provider_ops_pure::normalize_architecture_id(architecture_id) == "sub2api";
if !preserve_internal_runtime_fields {
saved_credentials.retain(|key, _| !key.starts_with('_'));
}
for field in ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS {
let mut reused_saved_secret = false;
for field in PROVIDER_OPS_PERSISTENT_SECRET_FIELDS {
if field.starts_with('_') {
continue;
}
@@ -146,6 +223,7 @@ pub(super) fn admin_provider_ops_merge_credentials(
{
if let Some(saved_value) = saved_credentials.get(*field) {
request_credentials.insert((*field).to_string(), saved_value.clone());
reused_saved_secret = true;
}
}
}
@@ -158,23 +236,29 @@ pub(super) fn admin_provider_ops_merge_credentials(
}
}
request_credentials
Ok(AdminProviderOpsMergedCredentialSnapshot {
provider: snapshot.provider,
credentials: request_credentials,
saved_binding: snapshot.binding,
reused_saved_secret,
})
}
fn admin_provider_ops_encrypt_credentials(
state: &AdminAppState<'_>,
binding: &ProviderOpsCredentialBinding,
credentials: serde_json::Map<String, serde_json::Value>,
) -> Result<serde_json::Map<String, serde_json::Value>, String> {
let mut encrypted = serde_json::Map::new();
for (key, value) in credentials {
if ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS.contains(&key.as_str()) {
if provider_ops_credential_field_is_secret(&key) {
if let Some(plaintext) = value.as_str() {
if plaintext.is_empty() {
encrypted.insert(key, value);
} else {
let ciphertext = state
.encrypt_catalog_secret_with_fallbacks(plaintext)
.ok_or_else(|| "gateway 未配置 Provider Ops 加密密钥".to_string())?;
let ciphertext =
seal_provider_ops_credential(state.app(), binding, &key, plaintext)
.map_err(ToString::to_string)?;
encrypted.insert(key, serde_json::Value::String(ciphertext));
}
continue;
@@ -185,6 +269,104 @@ fn admin_provider_ops_encrypt_credentials(
Ok(encrypted)
}
fn admin_provider_ops_config_with_credentials(
provider: &StoredProviderCatalogProvider,
credentials: serde_json::Map<String, serde_json::Value>,
binding: &ProviderOpsCredentialBinding,
) -> Result<Option<serde_json::Value>, String> {
let mut provider_config = provider
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.cloned()
.ok_or_else(|| "Provider Ops 配置格式无效".to_string())?;
let mut provider_ops_config = provider_config
.get("provider_ops")
.and_then(serde_json::Value::as_object)
.cloned()
.ok_or_else(|| "Provider Ops 配置格式无效".to_string())?;
let mut connector_config = provider_ops_config
.get("connector")
.and_then(serde_json::Value::as_object)
.cloned()
.ok_or_else(|| "Provider Ops connector 配置格式无效".to_string())?;
connector_config.insert(
"credentials".to_string(),
serde_json::Value::Object(credentials),
);
provider_ops_config.insert(
"connector".to_string(),
serde_json::Value::Object(connector_config),
);
provider_ops_config.insert(
"base_url".to_string(),
serde_json::Value::String(binding.destination.base_url().to_string()),
);
provider_config.insert(
"provider_ops".to_string(),
serde_json::Value::Object(provider_ops_config),
);
Ok(Some(serde_json::Value::Object(provider_config)))
}
pub(crate) async fn admin_provider_ops_credential_snapshot(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
) -> Result<AdminProviderOpsCredentialSnapshot, GatewayError> {
let mut current = provider.clone();
for _ in 0..PROVIDER_OPS_CREDENTIAL_MIGRATION_RETRIES {
let (binding, needs_materialized_base_url) =
admin_provider_ops_binding_for_provider(state, &current).await?;
let raw_credentials = admin_provider_ops_config_object(&current)
.and_then(admin_provider_ops_connector_object)
.and_then(|connector| connector.get("credentials"));
let decoded = admin_provider_ops_decode_credentials(state, &binding, raw_credentials)
.map_err(GatewayError::Internal)?;
if !decoded.migration_required && !needs_materialized_base_url {
return Ok(AdminProviderOpsCredentialSnapshot {
provider: current,
credentials: decoded.values,
binding,
});
}
let migrated_config = admin_provider_ops_config_with_credentials(
&current,
decoded.protected_values,
&binding,
)
.map_err(GatewayError::Internal)?;
let update = ProviderCatalogProviderConfigCasUpdate {
provider_id: current.id.clone(),
expected_config: current.config.clone(),
config: migrated_config.clone(),
};
if state
.compare_and_swap_provider_catalog_provider_config(&update)
.await?
{
current.config = migrated_config;
return Ok(AdminProviderOpsCredentialSnapshot {
provider: current,
credentials: decoded.values,
binding,
});
}
current = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&current.id))
.await?
.into_iter()
.next()
.ok_or_else(|| GatewayError::Internal("Provider Ops Provider 不存在".to_string()))?;
}
Err(GatewayError::Internal(
"Provider Ops 凭据迁移未能稳定完成".to_string(),
))
}
pub(super) async fn persist_admin_provider_ops_runtime_credentials(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
@@ -193,82 +375,91 @@ pub(super) async fn persist_admin_provider_ops_runtime_credentials(
if updated_credentials.is_empty() || !state.has_provider_catalog_data_writer() {
return Ok(None);
}
let mut updated_provider = provider.clone();
let mut provider_config = updated_provider
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
let Some(provider_ops_config) = provider_config
.get("provider_ops")
.and_then(serde_json::Value::as_object)
.cloned()
else {
return Ok(None);
};
let Some(connector_config) = provider_ops_config
.get("connector")
.and_then(serde_json::Value::as_object)
.cloned()
else {
return Ok(None);
};
let mut decrypted_credentials =
admin_provider_ops_decrypted_credentials(state, connector_config.get("credentials"));
for (key, value) in updated_credentials {
decrypted_credentials.insert(key.clone(), value.clone());
for key in updated_credentials.keys() {
if key != "refresh_token"
&& key != "_cached_access_token"
&& !PROVIDER_OPS_TRANSIENT_METADATA_FIELDS.contains(&key.as_str())
{
return Err(GatewayError::Internal(format!(
"不允许持久化未知的 Provider Ops runtime credential 字段 '{key}'"
)));
}
}
let encrypted_credentials =
admin_provider_ops_encrypt_credentials(state, decrypted_credentials)
.map_err(GatewayError::Internal)?;
let mut updated_connector = connector_config.clone();
updated_connector.insert(
"credentials".to_string(),
serde_json::Value::Object(encrypted_credentials),
);
let mut updated_provider_ops = provider_ops_config.clone();
updated_provider_ops.insert(
"connector".to_string(),
serde_json::Value::Object(updated_connector),
);
provider_config.insert(
"provider_ops".to_string(),
serde_json::Value::Object(updated_provider_ops),
);
updated_provider.config = Some(serde_json::Value::Object(provider_config));
updated_provider.updated_at_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs());
state
.update_provider_catalog_provider(&updated_provider)
.await
let mut current = provider.clone();
for _ in 0..PROVIDER_OPS_CREDENTIAL_MIGRATION_RETRIES {
let snapshot = admin_provider_ops_credential_snapshot(state, &current).await?;
let mut decrypted_credentials = snapshot.credentials;
for (key, value) in updated_credentials {
decrypted_credentials.insert(key.clone(), value.clone());
}
let encrypted_credentials =
admin_provider_ops_encrypt_credentials(state, &snapshot.binding, decrypted_credentials)
.map_err(GatewayError::Internal)?;
let config = admin_provider_ops_config_with_credentials(
&snapshot.provider,
encrypted_credentials,
&snapshot.binding,
)
.map_err(GatewayError::Internal)?;
let update = ProviderCatalogProviderConfigCasUpdate {
provider_id: snapshot.provider.id.clone(),
expected_config: snapshot.provider.config.clone(),
config,
};
if state
.compare_and_swap_provider_catalog_provider_config(&update)
.await?
{
return Ok(state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider.id))
.await?
.into_iter()
.next());
}
current = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider.id))
.await?
.into_iter()
.next()
.ok_or_else(|| GatewayError::Internal("Provider Ops Provider 不存在".to_string()))?;
}
Err(GatewayError::Internal(
"Provider Ops runtime credential 并发更新未能稳定完成".to_string(),
))
}
pub(super) fn build_admin_provider_ops_saved_config_value(
pub(super) async fn build_admin_provider_ops_saved_config_value(
state: &AdminAppState<'_>,
provider: &StoredProviderCatalogProvider,
payload: AdminProviderOpsSaveConfigRequest,
) -> Result<serde_json::Value, String> {
) -> Result<AdminProviderOpsSavedConfigSnapshot, String> {
let architecture_id = payload.architecture_id.trim();
let normalized_architecture_id =
admin_provider_ops_pure::normalize_architecture_id(architecture_id);
if architecture_id.is_empty() || architecture_id != normalized_architecture_id {
return Err("architecture_id 必须是合法的 Provider Ops 架构".to_string());
}
let auth_type = payload.connector.auth_type.trim().to_string();
if auth_type.is_empty() || !admin_provider_ops_is_supported_auth_type(auth_type.as_str()) {
return Err("connector.auth_type 必须是合法的认证类型".to_string());
}
let merged_credentials = admin_provider_ops_merge_credentials(
let merged = admin_provider_ops_merge_credentials(
state,
payload.architecture_id.as_str(),
normalized_architecture_id,
provider,
payload.connector.credentials,
);
let encrypted_credentials = admin_provider_ops_encrypt_credentials(state, merged_credentials)?;
)
.await?;
let canonical_base_url = payload
.base_url
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or_else(|| merged.saved_binding.destination.base_url());
let canonical_destination =
canonicalize_provider_ops_base_url(canonical_base_url).map_err(ToString::to_string)?;
let actions = payload
.actions
@@ -284,19 +475,48 @@ pub(super) fn build_admin_provider_ops_saved_config_value(
})
.collect::<serde_json::Map<String, serde_json::Value>>();
let quota_alert = normalize_admin_provider_ops_quota_alert(payload.quota_alert)?;
Ok(json!({
"architecture_id": payload.architecture_id,
"base_url": payload.base_url,
let mut provider_ops_config = json!({
"architecture_id": normalized_architecture_id,
"base_url": canonical_destination.base_url(),
"connector": {
"auth_type": auth_type,
"config": payload.connector.config,
"credentials": encrypted_credentials,
"credentials": {},
},
"actions": actions,
"schedule": payload.schedule,
"quota_alert": quota_alert,
}))
});
let new_binding = admin_provider_ops_binding_from_config(
&merged.provider.id,
provider_ops_config
.as_object()
.ok_or_else(|| "Provider Ops 配置格式无效".to_string())?,
canonical_destination.base_url(),
)?;
let same_secret_destination = merged.saved_binding.provider_id == new_binding.provider_id
&& merged.saved_binding.architecture_id == new_binding.architecture_id
&& merged.saved_binding.auth_type == new_binding.auth_type
&& merged.saved_binding.destination == new_binding.destination;
if merged.reused_saved_secret && !same_secret_destination {
return Err("修改 Provider Ops 架构、认证类型或目标地址时必须重新填写凭据".to_string());
}
let mut merged_credentials = merged.credentials;
if merged.saved_binding != new_binding {
for field in PROVIDER_OPS_TRANSIENT_METADATA_FIELDS {
merged_credentials.remove(*field);
}
merged_credentials.retain(|field, _| !field.starts_with("_cached_"));
}
let encrypted_credentials =
admin_provider_ops_encrypt_credentials(state, &new_binding, merged_credentials)?;
provider_ops_config["connector"]["credentials"] =
serde_json::Value::Object(encrypted_credentials);
Ok(AdminProviderOpsSavedConfigSnapshot {
provider: merged.provider,
provider_ops_config,
})
}
fn normalize_admin_provider_ops_quota_alert(
@@ -356,27 +576,35 @@ pub(super) fn build_admin_provider_ops_status_payload(
admin_provider_ops_pure::build_admin_provider_ops_status_payload(provider_id, provider)
}
pub(super) fn build_admin_provider_ops_config_payload(
pub(super) async fn build_admin_provider_ops_config_payload(
state: &AdminAppState<'_>,
provider_id: &str,
provider: Option<&StoredProviderCatalogProvider>,
endpoints: &[StoredProviderCatalogEndpoint],
) -> serde_json::Value {
) -> Result<serde_json::Value, GatewayError> {
let Some(provider) = provider else {
return json!({
return Ok(json!({
"provider_id": provider_id,
"is_configured": false,
});
}));
};
let Some(provider_ops_config) = admin_provider_ops_config_object(provider) else {
return json!({
if admin_provider_ops_config_object(provider).is_none() {
return Ok(json!({
"provider_id": provider_id,
"is_configured": false,
});
}));
}
let snapshot = admin_provider_ops_credential_snapshot(state, provider).await?;
let provider = &snapshot.provider;
let Some(provider_ops_config) = admin_provider_ops_config_object(provider) else {
return Ok(json!({
"provider_id": provider_id,
"is_configured": false,
}));
};
let connector = admin_provider_ops_connector_object(provider_ops_config);
json!({
Ok(json!({
"provider_id": provider_id,
"is_configured": true,
"architecture_id": provider_ops_config
@@ -398,15 +626,159 @@ pub(super) fn build_admin_provider_ops_config_payload(
.filter(|value| value.is_object())
.cloned()
.unwrap_or_else(|| json!({})),
"credentials": admin_provider_ops_masked_credentials(
state,
connector.and_then(|connector| connector.get("credentials")),
),
"credentials": admin_provider_ops_masked_credentials(&snapshot.credentials),
},
"quota_alert": provider_ops_config
.get("quota_alert")
.filter(|value| value.is_object())
.cloned()
.unwrap_or_else(default_admin_provider_ops_quota_alert),
})
}))
}
#[cfg(test)]
mod tests {
use super::{admin_provider_ops_credential_snapshot, open_provider_ops_credential};
use crate::data::GatewayDataState;
use crate::handlers::admin::request::AdminAppState;
use crate::AppState;
use aether_crypto::{
decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext,
looks_like_python_fernet_ciphertext, DEVELOPMENT_ENCRYPTION_KEY,
};
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogReadRepository, StoredProviderCatalogProvider,
};
use serde_json::json;
use std::sync::Arc;
const TEST_PROVIDER_ID: &str = "provider-ops-secret-test";
const TEST_API_KEY: &str = "legacy-provider-ops-api-key";
fn provider_with_api_key(api_key: &str) -> StoredProviderCatalogProvider {
StoredProviderCatalogProvider::new(
TEST_PROVIDER_ID.to_string(),
"Provider Ops Secret Test".to_string(),
None,
"openai".to_string(),
)
.expect("provider should build")
.with_transport_fields(
true,
false,
false,
None,
None,
None,
None,
None,
Some(json!({
"provider_ops": {
"architecture_id": "generic_api",
"base_url": "https://provider.example.com",
"connector": {
"auth_type": "api_key",
"config": {},
"credentials": {
"api_key": api_key,
"account_id": "account-1"
}
},
"actions": {},
"schedule": {}
}
})),
)
}
fn state_with_provider(
provider: StoredProviderCatalogProvider,
) -> (AppState, Arc<InMemoryProviderCatalogReadRepository>) {
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
Vec::new(),
Vec::new(),
));
let state = AppState::new()
.expect("gateway state should build")
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(repository.clone())
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
);
(state, repository)
}
async fn stored_provider(
repository: &InMemoryProviderCatalogReadRepository,
) -> StoredProviderCatalogProvider {
repository
.list_providers_by_ids(&[TEST_PROVIDER_ID.to_string()])
.await
.expect("provider should read")
.into_iter()
.next()
.expect("provider should exist")
}
#[tokio::test]
async fn legacy_provider_ops_credentials_are_lazily_migrated() {
let provider = provider_with_api_key(TEST_API_KEY);
let (state, repository) = state_with_provider(provider.clone());
let admin_state = AdminAppState::new(&state);
let snapshot = admin_provider_ops_credential_snapshot(&admin_state, &provider)
.await
.expect("legacy Provider Ops credential should migrate");
assert_eq!(snapshot.credentials["api_key"], TEST_API_KEY);
assert_eq!(snapshot.credentials["account_id"], "account-1");
let stored = stored_provider(repository.as_ref()).await;
let ciphertext = stored
.config
.as_ref()
.and_then(|config| config.pointer("/provider_ops/connector/credentials/api_key"))
.and_then(serde_json::Value::as_str)
.expect("stored API key should exist");
assert_ne!(ciphertext, TEST_API_KEY);
// New migrations use a binding-aware runtime-secret envelope. Keep
// the legacy Fernet assertion below only for the tamper fixture; a
// migrated value must no longer be treated as an unbound Fernet blob.
assert!(ciphertext.starts_with("aether-provider-ops-credential-v2:"));
assert_eq!(
open_provider_ops_credential(&state, &snapshot.binding, "api_key", ciphertext)
.expect("migrated Provider Ops API key should decrypt")
.plaintext,
TEST_API_KEY
);
}
#[tokio::test]
async fn tampered_provider_ops_ciphertext_fails_closed() {
let mut tampered =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, TEST_API_KEY)
.expect("Provider Ops API key should encrypt");
tampered.replace_range(tampered.len() - 2.., "AA");
assert!(looks_like_python_fernet_ciphertext(&tampered));
let provider = provider_with_api_key(&tampered);
let (state, repository) = state_with_provider(provider.clone());
let admin_state = AdminAppState::new(&state);
let error = admin_provider_ops_credential_snapshot(&admin_state, &provider)
.await
.expect_err("tampered Provider Ops ciphertext must not be used as plaintext");
assert!(format!("{error:?}").contains("无法解密"));
let stored = stored_provider(repository.as_ref()).await;
assert_eq!(
stored
.config
.as_ref()
.and_then(|config| {
config.pointer("/provider_ops/connector/credentials/api_key")
})
.and_then(serde_json::Value::as_str),
Some(tampered.as_str())
);
}
}
@@ -5,4 +5,5 @@ mod routes;
mod support;
mod verify;
pub(crate) use self::balance_cache::store_admin_provider_ops_balance_cache;
pub(crate) use self::config::admin_provider_ops_credential_snapshot;
pub(super) use self::routes::maybe_build_local_admin_provider_ops_providers_response;
@@ -15,7 +15,7 @@ use axum::{
};
use futures_util::stream::{self, StreamExt};
use serde_json::json;
use std::collections::HashMap;
use std::collections::{HashMap, HashSet};
pub(super) async fn handle_admin_provider_ops_batch_balance(
state: &AdminAppState<'_>,
@@ -166,12 +166,48 @@ fn parse_provider_ids(body: &Bytes) -> Result<Vec<String>, Response<Body>> {
)
.into_response()
})?;
let mut seen = HashSet::with_capacity(items.len());
let mut provider_ids = Vec::with_capacity(items.len());
for item in items {
let Some(provider_id) = item.as_str().map(str::trim) else {
return Err((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "provider_ids 必须是非空字符串数组" })),
)
.into_response());
};
if provider_id.is_empty() {
return Err((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "provider_ids 必须是非空字符串数组" })),
)
.into_response());
}
if seen.insert(provider_id) {
provider_ids.push(provider_id.to_string());
}
}
Ok(provider_ids)
}
Ok(items
.iter()
.filter_map(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.collect())
#[cfg(test)]
mod tests {
use super::parse_provider_ids;
use axum::body::Bytes;
#[test]
fn provider_ops_batch_ids_are_deduplicated_without_reordering() {
let body =
Bytes::from_static(br#"{"provider_ids":["provider-1"," provider-1 ","provider-2"]}"#);
assert_eq!(
parse_provider_ids(&body).expect("valid ids"),
vec!["provider-1".to_string(), "provider-2".to_string()]
);
}
#[test]
fn provider_ops_batch_ids_reject_non_string_entries() {
let body = Bytes::from_static(br#"{"provider_ids":["provider-1",42]}"#);
assert!(parse_provider_ids(&body).is_err());
}
}
@@ -3,6 +3,7 @@ use super::super::config::build_admin_provider_ops_saved_config_value;
use super::super::support::AdminProviderOpsSaveConfigRequest;
use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError;
use aether_data_contracts::repository::provider_catalog::ProviderCatalogProviderConfigCasUpdate;
use axum::{
body::{Body, Bytes},
http,
@@ -10,7 +11,8 @@ use axum::{
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
const ADMIN_PROVIDER_OPS_CONFIG_SAVE_RETRIES: usize = 8;
pub(super) async fn handle_admin_provider_ops_save_config(
state: &AdminAppState<'_>,
@@ -23,7 +25,7 @@ pub(super) async fn handle_admin_provider_ops_save_config(
Err(response) => return Ok(Some(response)),
};
let provider_ids = [provider_id.to_string()];
let Some(existing_provider) = state
let Some(mut existing_provider) = state
.read_provider_catalog_providers_by_ids(&provider_ids)
.await?
.into_iter()
@@ -32,31 +34,57 @@ pub(super) async fn handle_admin_provider_ops_save_config(
return Ok(Some(provider_not_found_response()));
};
let provider_ops_config =
match build_admin_provider_ops_saved_config_value(state, &existing_provider, payload) {
Ok(config) => config,
let mut saved = false;
for _ in 0..ADMIN_PROVIDER_OPS_CONFIG_SAVE_RETRIES {
let snapshot = match build_admin_provider_ops_saved_config_value(
state,
&existing_provider,
payload.clone(),
)
.await
{
Ok(snapshot) => snapshot,
Err(detail) => return Ok(Some(bad_request_detail_response(&detail))),
};
let mut updated_provider = existing_provider.clone();
let mut provider_config = updated_provider
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
provider_config.insert("provider_ops".to_string(), provider_ops_config);
updated_provider.config = Some(serde_json::Value::Object(provider_config));
updated_provider.updated_at_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs());
let Some(_updated) = state
.update_provider_catalog_provider(&updated_provider)
.await?
else {
return Ok(None);
};
let mut provider_config = snapshot
.provider
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
provider_config.insert("provider_ops".to_string(), snapshot.provider_ops_config);
let update = ProviderCatalogProviderConfigCasUpdate {
provider_id: snapshot.provider.id.clone(),
expected_config: snapshot.provider.config.clone(),
config: Some(serde_json::Value::Object(provider_config)),
};
if state
.compare_and_swap_provider_catalog_provider_config(&update)
.await?
{
saved = true;
break;
}
let Some(current) = state
.read_provider_catalog_providers_by_ids(&provider_ids)
.await?
.into_iter()
.next()
else {
return Ok(Some(provider_not_found_response()));
};
existing_provider = current;
}
if !saved {
return Ok(Some(
(
http::StatusCode::CONFLICT,
Json(json!({ "detail": "Provider Ops 配置并发更新冲突,请重试" })),
)
.into_response(),
));
}
clear_admin_provider_ops_balance_cache(state, provider_id).await;
Ok(Some(
@@ -73,7 +101,7 @@ pub(super) async fn handle_admin_provider_ops_delete_config(
provider_id: &str,
) -> Result<Option<Response<Body>>, GatewayError> {
let provider_ids = [provider_id.to_string()];
let Some(existing_provider) = state
let Some(mut existing_provider) = state
.read_provider_catalog_providers_by_ids(&provider_ids)
.await?
.into_iter()
@@ -82,25 +110,49 @@ pub(super) async fn handle_admin_provider_ops_delete_config(
return Ok(Some(provider_not_found_response()));
};
let mut updated_provider = existing_provider.clone();
let mut provider_config = updated_provider
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
if provider_config.remove("provider_ops").is_some() {
updated_provider.config = Some(serde_json::Value::Object(provider_config));
updated_provider.updated_at_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs());
let Some(_updated) = state
.update_provider_catalog_provider(&updated_provider)
.await?
else {
return Ok(None);
let mut removed = false;
for _ in 0..ADMIN_PROVIDER_OPS_CONFIG_SAVE_RETRIES {
let mut provider_config = existing_provider
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
if provider_config.remove("provider_ops").is_none() {
break;
}
let update = ProviderCatalogProviderConfigCasUpdate {
provider_id: existing_provider.id.clone(),
expected_config: existing_provider.config.clone(),
config: Some(serde_json::Value::Object(provider_config)),
};
if state
.compare_and_swap_provider_catalog_provider_config(&update)
.await?
{
removed = true;
break;
}
let Some(current) = state
.read_provider_catalog_providers_by_ids(&provider_ids)
.await?
.into_iter()
.next()
else {
return Ok(Some(provider_not_found_response()));
};
existing_provider = current;
}
if admin_provider_ops_config_still_exists(&existing_provider) && !removed {
return Ok(Some(
(
http::StatusCode::CONFLICT,
Json(json!({ "detail": "Provider Ops 配置并发更新冲突,请重试" })),
)
.into_response(),
));
}
if removed {
clear_admin_provider_ops_balance_cache(state, provider_id).await;
}
@@ -113,6 +165,16 @@ pub(super) async fn handle_admin_provider_ops_delete_config(
))
}
fn admin_provider_ops_config_still_exists(
provider: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider,
) -> bool {
provider
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.is_some_and(|config| config.contains_key("provider_ops"))
}
fn parse_json_object_payload<T>(request_body: Option<&Bytes>) -> Result<T, Response<Body>>
where
T: serde::de::DeserializeOwned,
@@ -1,6 +1,6 @@
use super::super::config::{
admin_provider_ops_config_object, admin_provider_ops_connector_object,
admin_provider_ops_decrypted_credentials, resolve_admin_provider_ops_base_url,
admin_provider_ops_config_object, admin_provider_ops_credential_snapshot,
resolve_admin_provider_ops_base_url,
};
use super::super::support::{
AdminProviderOpsConnectRequest, ADMIN_PROVIDER_OPS_CONNECT_RUST_ONLY_MESSAGE,
@@ -33,6 +33,9 @@ pub(super) async fn handle_admin_provider_ops_connect(
else {
return Ok(bad_request_detail_response("Provider 不存在"));
};
let credential_snapshot =
admin_provider_ops_credential_snapshot(state, &existing_provider).await?;
let existing_provider = credential_snapshot.provider;
let Some(provider_ops_config) = admin_provider_ops_config_object(&existing_provider) else {
return Ok(bad_request_detail_response("未配置操作设置"));
};
@@ -52,14 +55,7 @@ pub(super) async fn handle_admin_provider_ops_connect(
let actual_credentials = payload
.credentials
.filter(|value| !value.is_empty())
.unwrap_or_else(|| {
admin_provider_ops_decrypted_credentials(
state,
admin_provider_ops_config_object(&existing_provider)
.and_then(admin_provider_ops_connector_object)
.and_then(|connector| connector.get("credentials")),
)
});
.unwrap_or(credential_snapshot.credentials);
if actual_credentials.is_empty() {
return Ok(bad_request_detail_response("未提供凭据"));
}

Some files were not shown because too many files have changed in this diff Show More