Preserve usage data in system imports

This commit is contained in:
fawney19
2026-05-24 21:40:48 +08:00
parent 18d9004f22
commit e2d5fc9dfb
21 changed files with 2683 additions and 309 deletions
@@ -1684,6 +1684,28 @@ impl GatewayDataState {
}
}
pub(crate) async fn set_api_key_usage_totals(
&self,
api_key_id: &str,
total_requests: u64,
total_tokens: u64,
total_cost_usd: f64,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
match &self.auth_api_key_writer {
Some(repository) => {
repository
.set_api_key_usage_totals(
api_key_id,
total_requests,
total_tokens,
total_cost_usd,
)
.await
}
None => Ok(None),
}
}
pub(crate) async fn set_standalone_api_key_feature_settings(
&self,
api_key_id: &str,
@@ -511,6 +511,45 @@ impl GatewayDataState {
}
}
pub(crate) async fn export_admin_system_usage_aggregates(
&self,
) -> Result<aether_data::repository::system::AdminSystemUsageAggregateSnapshot, DataLayerError>
{
match self.backends.as_ref() {
Some(backends) => backends.export_admin_system_usage_aggregates().await,
None => {
Ok(aether_data::repository::system::AdminSystemUsageAggregateSnapshot::default())
}
}
}
pub(crate) async fn import_admin_system_usage_aggregates(
&self,
snapshot: &aether_data::repository::system::AdminSystemUsageAggregateSnapshot,
user_id_map: &std::collections::BTreeMap<String, String>,
api_key_id_map: &std::collections::BTreeMap<String, String>,
mode: aether_data::repository::system::AdminSystemUsageAggregateImportMode,
) -> Result<
aether_data::repository::system::AdminSystemUsageAggregateImportSummary,
DataLayerError,
> {
match self.backends.as_ref() {
Some(backends) => {
backends
.import_admin_system_usage_aggregates(
snapshot,
user_id_map,
api_key_id_map,
mode,
)
.await
}
None => Ok(
aether_data::repository::system::AdminSystemUsageAggregateImportSummary::default(),
),
}
}
pub(crate) async fn purge_admin_request_bodies_batch(
&self,
batch_size: usize,
@@ -185,6 +185,7 @@ impl<'a> AdminAppState<'a> {
let standalone_wallets = self
.list_wallet_snapshots_by_api_key_ids(&standalone_api_key_ids)
.await?;
let usage_aggregates = self.export_admin_system_usage_aggregates().await?;
let wallets_by_user_id = user_wallets
.into_iter()
@@ -262,6 +263,7 @@ impl<'a> AdminAppState<'a> {
.collect::<Vec<_>>();
json!({
"id": user.id.clone(),
"email": user.email.clone(),
"email_verified": user.email_verified,
"username": user.username.clone(),
@@ -306,6 +308,7 @@ impl<'a> AdminAppState<'a> {
"user_groups": user_groups_data,
"users": users_data,
"standalone_keys": standalone_keys_data,
"usage_aggregates": usage_aggregates,
}))
}
@@ -330,6 +333,7 @@ impl<'a> AdminAppState<'a> {
include_is_standalone: bool,
) -> serde_json::Value {
let mut payload = serde_json::Map::from_iter([
("api_key_id".to_string(), json!(key.api_key_id.clone())),
("key_hash".to_string(), json!(key.key_hash.clone())),
("name".to_string(), json!(key.name.clone())),
(
@@ -38,6 +38,10 @@ use aether_data::repository::auth_modules::StoredLdapModuleConfig;
use aether_data::repository::oauth_providers::{
EncryptedSecretUpdate, UpsertOAuthProviderConfigRecord,
};
use aether_data::repository::system::{
AdminSystemUsageAggregateImportMode, AdminSystemUsageAggregateImportSummary,
AdminSystemUsageAggregateSnapshot,
};
use aether_data::repository::wallet::WalletLookupKey;
use aether_data_contracts::repository::global_models::{
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
@@ -344,54 +348,6 @@ fn imported_oauth_expires_at_unix_secs(normalized_auth_config: Option<&Value>) -
None
}
fn imported_oauth_has_refresh_token(normalized_auth_config: Option<&Value>) -> bool {
normalized_auth_config
.and_then(Value::as_object)
.and_then(|object| object.get("refresh_token"))
.and_then(Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty())
}
async fn refresh_imported_oauth_key_after_persist(
state: &AdminAppState<'_>,
provider: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider,
key_id: &str,
) -> Result<(), GatewayError> {
let Some(endpoint) =
crate::handlers::admin::provider::oauth::runtime::resolve_provider_oauth_runtime_endpoints(
state,
provider,
provider.provider_type.as_str(),
)
.await?
.runtime_endpoint
else {
return Ok(());
};
let Some(transport) = state
.read_provider_transport_snapshot(&provider.id, &endpoint.id, key_id)
.await?
else {
return Ok(());
};
if !crate::provider_transport::supports_local_oauth_request_auth_resolution(&transport) {
return Ok(());
}
if let Err(error) = state.force_local_oauth_refresh_entry(&transport).await {
tracing::warn!(
provider_id = %provider.id,
provider_type = %provider.provider_type,
key_id = %key_id,
error = ?error,
"admin system import oauth refresh after credential import failed"
);
}
Ok(())
}
fn build_import_provider_model_record(
provider_id: &str,
existing_id: Option<&str>,
@@ -435,6 +391,8 @@ struct AdminSystemUsersImportStats {
users: AdminSystemConfigImportCounter,
api_keys: AdminSystemConfigImportCounter,
standalone_keys: AdminSystemConfigImportCounter,
#[serde(skip_serializing_if = "Option::is_none")]
usage_aggregates: Option<AdminSystemUsageAggregateImportSummary>,
errors: Vec<String>,
}
@@ -488,6 +446,16 @@ fn validate_imported_system_users_export_version(version: Option<&Value>) -> Res
Ok(())
}
fn usage_aggregate_import_mode(
merge_mode: AdminImportMergeMode,
) -> AdminSystemUsageAggregateImportMode {
match merge_mode {
AdminImportMergeMode::Skip => AdminSystemUsageAggregateImportMode::Skip,
AdminImportMergeMode::Overwrite => AdminSystemUsageAggregateImportMode::Overwrite,
AdminImportMergeMode::Error => AdminSystemUsageAggregateImportMode::Error,
}
}
fn imported_object_field<'a>(
value: &'a Value,
field_name: &str,
@@ -1560,16 +1528,6 @@ impl<'a> AdminAppState<'a> {
"更新 Provider '{provider_name}' 的 Key 失败"
))));
};
if auth_type == "oauth"
&& imported_oauth_has_refresh_token(normalized_auth_config.as_ref())
{
refresh_imported_oauth_key_after_persist(
self,
&provider,
&persisted.id,
)
.await?;
}
existing_keys[existing_index] = persisted;
stats.keys.updated += 1;
}
@@ -1610,11 +1568,6 @@ impl<'a> AdminAppState<'a> {
"创建 Provider '{provider_name}' 的 Key 失败"
))));
};
if auth_type == "oauth"
&& imported_oauth_has_refresh_token(normalized_auth_config.as_ref())
{
refresh_imported_oauth_key_after_persist(self, &provider, &created.id).await?;
}
existing_keys.push(created);
stats.keys.created += 1;
}
@@ -2052,6 +2005,8 @@ impl<'a> AdminAppState<'a> {
));
let mut stats = AdminSystemUsersImportStats::default();
let mut imported_user_id_map = BTreeMap::<String, String>::new();
let mut imported_api_key_id_map = BTreeMap::<String, String>::new();
let default_group_id = self.effective_default_user_group_id().await?;
let existing_groups = self.list_user_groups().await?;
let mut groups_by_name = existing_groups
@@ -2138,6 +2093,7 @@ impl<'a> AdminAppState<'a> {
Ok(value) => value,
Err(detail) => return Ok(Err(invalid_request(detail))),
};
let source_user_id = invalid_value!(imported_optional_string(user.get("id")));
let role = invalid_value!(imported_optional_string(user.get("role")))
.unwrap_or_else(|| "user".to_string())
.to_ascii_lowercase();
@@ -2471,6 +2427,9 @@ impl<'a> AdminAppState<'a> {
stats.users.created += 1;
created.id
};
if let Some(source_user_id) = source_user_id {
imported_user_id_map.insert(source_user_id, user_id.clone());
}
let existing_api_keys = self
.list_auth_api_key_export_records_by_user_ids(std::slice::from_ref(&user_id))
@@ -2506,6 +2465,8 @@ impl<'a> AdminAppState<'a> {
));
continue;
};
let source_api_key_id =
invalid_value!(imported_optional_string(key.get("api_key_id")));
let name = invalid_value!(imported_optional_string(key.get("name")));
let allowed_providers = invalid_value!(normalize_imported_user_string_list(
key,
@@ -2538,21 +2499,21 @@ impl<'a> AdminAppState<'a> {
let auto_delete_on_expiry =
invalid_value!(imported_optional_bool(key.get("auto_delete_on_expiry")))
.unwrap_or(false);
let total_requests = invalid_value!(imported_optional_u64(
let imported_total_requests = invalid_value!(imported_optional_u64(
key.get("total_requests"),
"total_requests"
))
.unwrap_or(0);
let total_tokens = invalid_value!(imported_optional_u64(
));
let total_requests = imported_total_requests.unwrap_or(0);
let imported_total_tokens = invalid_value!(imported_optional_u64(
key.get("total_tokens"),
"total_tokens"
))
.unwrap_or(0);
let total_cost_usd = invalid_value!(imported_optional_f64(
));
let total_tokens = imported_total_tokens.unwrap_or(0);
let imported_total_cost_usd = invalid_value!(imported_optional_f64(
key.get("total_cost_usd"),
"total_cost_usd"
))
.unwrap_or(0.0);
));
let total_cost_usd = imported_total_cost_usd.unwrap_or(0.0);
let feature_settings = invalid_value!(imported_optional_json_object(
key.get("feature_settings"),
"feature_settings"
@@ -2624,13 +2585,31 @@ impl<'a> AdminAppState<'a> {
is_active,
)
.await?;
if imported_total_requests.is_some()
|| imported_total_tokens.is_some()
|| imported_total_cost_usd.is_some()
{
let updated_usage = self
.set_api_key_usage_totals(
&existing_key.api_key_id,
imported_total_requests
.unwrap_or(existing_key.total_requests),
imported_total_tokens.unwrap_or(existing_key.total_tokens),
imported_total_cost_usd
.unwrap_or(existing_key.total_cost_usd),
)
.await?;
if updated_usage.is_none() {
return Ok(Err((
http::StatusCode::SERVICE_UNAVAILABLE,
json!({ "detail": "Admin system data unavailable" }),
)));
}
}
if key.contains_key("allowed_api_formats")
|| key.contains_key("allowed_models")
|| key.contains_key("expires_at")
|| key.contains_key("auto_delete_on_expiry")
|| key.contains_key("total_requests")
|| key.contains_key("total_tokens")
|| key.contains_key("total_cost_usd")
{
stats.errors.push(format!(
"用户 '{}' 的现有 API Key 仅覆盖基础字段;高级导入字段保持原值",
@@ -2638,6 +2617,10 @@ impl<'a> AdminAppState<'a> {
));
}
stats.api_keys.updated += 1;
if let Some(source_api_key_id) = source_api_key_id.clone() {
imported_api_key_id_map
.insert(source_api_key_id, existing_key.api_key_id.clone());
}
}
}
continue;
@@ -2680,125 +2663,137 @@ impl<'a> AdminAppState<'a> {
)
.await?;
}
let created_api_key_id = created.api_key_id.clone();
existing_api_keys_by_hash.insert(key_hash, created);
if let Some(source_api_key_id) = source_api_key_id {
imported_api_key_id_map.insert(source_api_key_id, created_api_key_id);
}
stats.api_keys.created += 1;
}
}
if standalone_keys.is_empty() {
return Ok(Ok(json!({
"message": "用户数据导入成功",
"stats": stats,
})));
}
let Some(standalone_owner_id) = standalone_owner_id else {
stats.standalone_keys.skipped += standalone_keys.len() as u64;
stats
.errors
.push("无法导入独立余额 Key: 当前管理员用户记录不存在".to_string());
return Ok(Ok(json!({
"message": "用户数据导入成功",
"stats": stats,
})));
};
let existing_standalone_keys = self
.list_auth_api_key_export_standalone_records()
.await?
.into_iter()
.collect::<Vec<_>>();
let mut existing_standalone_by_hash = existing_standalone_keys
.into_iter()
.map(|record| (record.key_hash.clone(), record))
.collect::<BTreeMap<_, _>>();
for (index, raw_key) in standalone_keys.iter().enumerate() {
let key = match imported_object_field(raw_key, &format!("standalone_keys[{index}]")) {
Ok(value) => value,
Err(detail) => return Ok(Err(invalid_request(detail))),
};
let Some((key_hash, key_encrypted)) =
invalid_value!(self.resolve_imported_system_user_api_key_material(key))
else {
stats.standalone_keys.skipped += 1;
if !standalone_keys.is_empty() {
let Some(standalone_owner_id) = standalone_owner_id else {
stats.standalone_keys.skipped += standalone_keys.len() as u64;
stats
.errors
.push(format!("跳过无效独立余额 Key: standalone_keys[{index}]"));
continue;
.push("无法导入独立余额 Key: 当前管理员用户记录不存在".to_string());
if let Some(summary) = self
.import_admin_system_user_usage_aggregates(
root.get("usage_aggregates"),
&imported_user_id_map,
&imported_api_key_id_map,
merge_mode,
)
.await?
{
stats.usage_aggregates = Some(summary);
}
return Ok(Ok(json!({
"message": "用户数据导入成功",
"stats": stats,
})));
};
let name = invalid_value!(imported_optional_string(key.get("name")));
let allowed_providers = invalid_value!(normalize_imported_user_string_list(
key,
"allowed_providers"
));
let allowed_api_formats = invalid_value!(normalize_imported_user_api_formats(
key,
"allowed_api_formats"
));
let allowed_models =
invalid_value!(normalize_imported_user_string_list(key, "allowed_models"));
let ip_rules = invalid_value!(normalize_imported_user_ip_rules(key));
let rate_limit =
invalid_value!(imported_optional_i32(key.get("rate_limit"), "rate_limit"))
.unwrap_or(0);
let concurrent_limit = invalid_value!(imported_optional_i32(
key.get("concurrent_limit"),
"concurrent_limit"
));
if concurrent_limit.is_some_and(|value| value < 0) {
return Ok(Err(invalid_request("concurrent_limit 必须是非负整数")));
}
let force_capabilities = imported_optional_value(key.get("force_capabilities"));
let is_active =
invalid_value!(imported_optional_bool(key.get("is_active"))).unwrap_or(true);
let expires_at_unix_secs = invalid_value!(imported_rfc3339_to_unix_secs(
key.get("expires_at"),
"expires_at"
));
let auto_delete_on_expiry =
invalid_value!(imported_optional_bool(key.get("auto_delete_on_expiry")))
.unwrap_or(false);
let total_requests = invalid_value!(imported_optional_u64(
key.get("total_requests"),
"total_requests"
))
.unwrap_or(0);
let total_tokens = invalid_value!(imported_optional_u64(
key.get("total_tokens"),
"total_tokens"
))
.unwrap_or(0);
let total_cost_usd = invalid_value!(imported_optional_f64(
key.get("total_cost_usd"),
"total_cost_usd"
))
.unwrap_or(0.0);
let feature_settings = invalid_value!(imported_optional_json_object(
key.get("feature_settings"),
"feature_settings"
)
.and_then(normalize_admin_feature_settings));
let wallet_payload = match key.get("wallet") {
Some(Value::Object(map)) => Some(map),
Some(Value::Null) | None => None,
Some(_) => return Ok(Err(invalid_request("wallet 必须是对象"))),
};
let unlimited =
invalid_value!(imported_optional_bool(key.get("unlimited"))).unwrap_or(false);
let wallet_target =
invalid_value!(normalize_imported_wallet_target(wallet_payload, unlimited));
if let Some(existing_key) = existing_standalone_by_hash.get(&key_hash).cloned() {
match merge_mode {
AdminImportMergeMode::Skip => {
stats.standalone_keys.skipped += 1;
}
AdminImportMergeMode::Error => {
return Ok(Err(invalid_request("独立余额 Key 已存在")));
}
AdminImportMergeMode::Overwrite => {
let updated = self
let existing_standalone_keys = self
.list_auth_api_key_export_standalone_records()
.await?
.into_iter()
.collect::<Vec<_>>();
let mut existing_standalone_by_hash = existing_standalone_keys
.into_iter()
.map(|record| (record.key_hash.clone(), record))
.collect::<BTreeMap<_, _>>();
for (index, raw_key) in standalone_keys.iter().enumerate() {
let key = match imported_object_field(raw_key, &format!("standalone_keys[{index}]"))
{
Ok(value) => value,
Err(detail) => return Ok(Err(invalid_request(detail))),
};
let Some((key_hash, key_encrypted)) =
invalid_value!(self.resolve_imported_system_user_api_key_material(key))
else {
stats.standalone_keys.skipped += 1;
stats
.errors
.push(format!("跳过无效独立余额 Key: standalone_keys[{index}]"));
continue;
};
let source_api_key_id =
invalid_value!(imported_optional_string(key.get("api_key_id")));
let name = invalid_value!(imported_optional_string(key.get("name")));
let allowed_providers = invalid_value!(normalize_imported_user_string_list(
key,
"allowed_providers"
));
let allowed_api_formats = invalid_value!(normalize_imported_user_api_formats(
key,
"allowed_api_formats"
));
let allowed_models =
invalid_value!(normalize_imported_user_string_list(key, "allowed_models"));
let ip_rules = invalid_value!(normalize_imported_user_ip_rules(key));
let rate_limit =
invalid_value!(imported_optional_i32(key.get("rate_limit"), "rate_limit"))
.unwrap_or(0);
let concurrent_limit = invalid_value!(imported_optional_i32(
key.get("concurrent_limit"),
"concurrent_limit"
));
if concurrent_limit.is_some_and(|value| value < 0) {
return Ok(Err(invalid_request("concurrent_limit 必须是非负整数")));
}
let force_capabilities = imported_optional_value(key.get("force_capabilities"));
let is_active =
invalid_value!(imported_optional_bool(key.get("is_active"))).unwrap_or(true);
let expires_at_unix_secs = invalid_value!(imported_rfc3339_to_unix_secs(
key.get("expires_at"),
"expires_at"
));
let auto_delete_on_expiry =
invalid_value!(imported_optional_bool(key.get("auto_delete_on_expiry")))
.unwrap_or(false);
let imported_total_requests = invalid_value!(imported_optional_u64(
key.get("total_requests"),
"total_requests"
));
let total_requests = imported_total_requests.unwrap_or(0);
let imported_total_tokens = invalid_value!(imported_optional_u64(
key.get("total_tokens"),
"total_tokens"
));
let total_tokens = imported_total_tokens.unwrap_or(0);
let imported_total_cost_usd = invalid_value!(imported_optional_f64(
key.get("total_cost_usd"),
"total_cost_usd"
));
let total_cost_usd = imported_total_cost_usd.unwrap_or(0.0);
let feature_settings = invalid_value!(imported_optional_json_object(
key.get("feature_settings"),
"feature_settings"
)
.and_then(normalize_admin_feature_settings));
let wallet_payload = match key.get("wallet") {
Some(Value::Object(map)) => Some(map),
Some(Value::Null) | None => None,
Some(_) => return Ok(Err(invalid_request("wallet 必须是对象"))),
};
let unlimited =
invalid_value!(imported_optional_bool(key.get("unlimited"))).unwrap_or(false);
let wallet_target =
invalid_value!(normalize_imported_wallet_target(wallet_payload, unlimited));
if let Some(existing_key) = existing_standalone_by_hash.get(&key_hash).cloned() {
match merge_mode {
AdminImportMergeMode::Skip => {
stats.standalone_keys.skipped += 1;
}
AdminImportMergeMode::Error => {
return Ok(Err(invalid_request("独立余额 Key 已存在")));
}
AdminImportMergeMode::Overwrite => {
let updated = self
.update_standalone_api_key_basic(
aether_data::repository::auth::UpdateStandaloneApiKeyBasicRecord {
api_key_id: existing_key.api_key_id.clone(),
@@ -2819,94 +2814,134 @@ impl<'a> AdminAppState<'a> {
},
)
.await?;
if updated.is_none() {
return Ok(Err((
http::StatusCode::SERVICE_UNAVAILABLE,
json!({ "detail": "Admin system data unavailable" }),
)));
}
let _ = self
.set_standalone_api_key_active(&existing_key.api_key_id, is_active)
.await?;
if key.contains_key("feature_settings") {
if updated.is_none() {
return Ok(Err((
http::StatusCode::SERVICE_UNAVAILABLE,
json!({ "detail": "Admin system data unavailable" }),
)));
}
let _ = self
.set_standalone_api_key_feature_settings(
&existing_key.api_key_id,
feature_settings.clone(),
)
.set_standalone_api_key_active(&existing_key.api_key_id, is_active)
.await?;
if key.contains_key("feature_settings") {
let _ = self
.set_standalone_api_key_feature_settings(
&existing_key.api_key_id,
feature_settings.clone(),
)
.await?;
}
if imported_total_requests.is_some()
|| imported_total_tokens.is_some()
|| imported_total_cost_usd.is_some()
{
let updated_usage = self
.set_api_key_usage_totals(
&existing_key.api_key_id,
imported_total_requests
.unwrap_or(existing_key.total_requests),
imported_total_tokens.unwrap_or(existing_key.total_tokens),
imported_total_cost_usd
.unwrap_or(existing_key.total_cost_usd),
)
.await?;
if updated_usage.is_none() {
return Ok(Err((
http::StatusCode::SERVICE_UNAVAILABLE,
json!({ "detail": "Admin system data unavailable" }),
)));
}
}
if key.contains_key("expires_at")
|| key.contains_key("auto_delete_on_expiry")
|| key.contains_key("force_capabilities")
{
stats.errors.push(
"现有独立余额 Key 仅覆盖基础字段;高级导入字段保持原值"
.to_string(),
);
}
self.sync_imported_api_key_wallet(
&existing_key.api_key_id,
&wallet_target,
key.get("name")
.and_then(Value::as_str)
.unwrap_or("独立余额 Key"),
)
.await?;
stats.standalone_keys.updated += 1;
if let Some(source_api_key_id) = source_api_key_id.clone() {
imported_api_key_id_map
.insert(source_api_key_id, existing_key.api_key_id.clone());
}
}
if key.contains_key("expires_at")
|| key.contains_key("auto_delete_on_expiry")
|| key.contains_key("force_capabilities")
|| key.contains_key("total_requests")
|| key.contains_key("total_tokens")
|| key.contains_key("total_cost_usd")
{
stats.errors.push(
"现有独立余额 Key 仅覆盖基础字段;高级导入字段保持原值".to_string(),
);
}
self.sync_imported_api_key_wallet(
&existing_key.api_key_id,
&wallet_target,
key.get("name")
.and_then(Value::as_str)
.unwrap_or("独立余额 Key"),
)
.await?;
stats.standalone_keys.updated += 1;
}
continue;
}
continue;
}
let created = self
.create_standalone_api_key(
aether_data::repository::auth::CreateStandaloneApiKeyRecord {
user_id: standalone_owner_id.clone(),
api_key_id: Uuid::new_v4().to_string(),
key_hash: key_hash.clone(),
key_encrypted,
name,
allowed_providers,
allowed_api_formats,
allowed_models,
ip_rules,
rate_limit: Some(rate_limit),
concurrent_limit,
force_capabilities,
is_active,
expires_at_unix_secs,
auto_delete_on_expiry,
total_requests,
total_tokens,
total_cost_usd,
},
)
.await?;
let Some(created) = created else {
return Ok(Err((
http::StatusCode::SERVICE_UNAVAILABLE,
json!({ "detail": "Admin system data unavailable" }),
)));
};
if key.contains_key("feature_settings") {
let _ = self
.set_standalone_api_key_feature_settings(
&created.api_key_id,
feature_settings.clone(),
let created = self
.create_standalone_api_key(
aether_data::repository::auth::CreateStandaloneApiKeyRecord {
user_id: standalone_owner_id.clone(),
api_key_id: Uuid::new_v4().to_string(),
key_hash: key_hash.clone(),
key_encrypted,
name,
allowed_providers,
allowed_api_formats,
allowed_models,
ip_rules,
rate_limit: Some(rate_limit),
concurrent_limit,
force_capabilities,
is_active,
expires_at_unix_secs,
auto_delete_on_expiry,
total_requests,
total_tokens,
total_cost_usd,
},
)
.await?;
let Some(created) = created else {
return Ok(Err((
http::StatusCode::SERVICE_UNAVAILABLE,
json!({ "detail": "Admin system data unavailable" }),
)));
};
if key.contains_key("feature_settings") {
let _ = self
.set_standalone_api_key_feature_settings(
&created.api_key_id,
feature_settings.clone(),
)
.await?;
}
self.sync_imported_api_key_wallet(
&created.api_key_id,
&wallet_target,
created.name.as_deref().unwrap_or("独立余额 Key"),
)
.await?;
let created_api_key_id = created.api_key_id.clone();
existing_standalone_by_hash.insert(key_hash, created);
if let Some(source_api_key_id) = source_api_key_id {
imported_api_key_id_map.insert(source_api_key_id, created_api_key_id);
}
stats.standalone_keys.created += 1;
}
self.sync_imported_api_key_wallet(
&created.api_key_id,
&wallet_target,
created.name.as_deref().unwrap_or("独立余额 Key"),
}
if let Some(summary) = self
.import_admin_system_user_usage_aggregates(
root.get("usage_aggregates"),
&imported_user_id_map,
&imported_api_key_id_map,
merge_mode,
)
.await?;
existing_standalone_by_hash.insert(key_hash, created);
stats.standalone_keys.created += 1;
.await?
{
stats.usage_aggregates = Some(summary);
}
Ok(Ok(json!({
@@ -2915,6 +2950,40 @@ impl<'a> AdminAppState<'a> {
})))
}
async fn import_admin_system_user_usage_aggregates(
&self,
value: Option<&Value>,
user_id_map: &BTreeMap<String, String>,
api_key_id_map: &BTreeMap<String, String>,
merge_mode: AdminImportMergeMode,
) -> Result<Option<AdminSystemUsageAggregateImportSummary>, GatewayError> {
let Some(value) = value else {
return Ok(None);
};
if value.is_null() {
return Ok(None);
}
let snapshot = serde_json::from_value::<AdminSystemUsageAggregateSnapshot>(value.clone())
.map_err(|err| GatewayError::Client {
status: http::StatusCode::BAD_REQUEST,
message: format!("usage_aggregates 格式无效: {err}"),
})?;
if snapshot.stats_daily.is_empty()
&& snapshot.stats_user_daily.is_empty()
&& snapshot.stats_daily_api_key.is_empty()
{
return Ok(None);
}
self.import_admin_system_usage_aggregates(
&snapshot,
user_id_map,
api_key_id_map,
usage_aggregate_import_mode(merge_mode),
)
.await
.map(Some)
}
async fn sync_imported_user_wallet(
&self,
user_id: &str,
@@ -70,6 +70,26 @@ impl<'a> AdminAppState<'a> {
self.app.purge_admin_system_data(target).await
}
pub(crate) async fn export_admin_system_usage_aggregates(
&self,
) -> Result<aether_data::repository::system::AdminSystemUsageAggregateSnapshot, GatewayError>
{
self.app.export_admin_system_usage_aggregates().await
}
pub(crate) async fn import_admin_system_usage_aggregates(
&self,
snapshot: &aether_data::repository::system::AdminSystemUsageAggregateSnapshot,
user_id_map: &std::collections::BTreeMap<String, String>,
api_key_id_map: &std::collections::BTreeMap<String, String>,
mode: aether_data::repository::system::AdminSystemUsageAggregateImportMode,
) -> Result<aether_data::repository::system::AdminSystemUsageAggregateImportSummary, GatewayError>
{
self.app
.import_admin_system_usage_aggregates(snapshot, user_id_map, api_key_id_map, mode)
.await
}
pub(crate) async fn run_admin_system_cleanup_once(
&self,
) -> Result<crate::maintenance::AdminSystemCleanupSummary, GatewayError> {
@@ -706,6 +706,19 @@ impl<'a> AdminAppState<'a> {
.await
}
pub(crate) async fn set_api_key_usage_totals(
&self,
api_key_id: &str,
total_requests: u64,
total_tokens: u64,
total_cost_usd: f64,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.app
.set_api_key_usage_totals(api_key_id, total_requests, total_tokens, total_cost_usd)
.await
}
pub(crate) async fn delete_user_api_key(
&self,
user_id: &str,
+30
View File
@@ -629,6 +629,36 @@ impl AppState {
Ok(summary)
}
pub(crate) async fn export_admin_system_usage_aggregates(
&self,
) -> Result<aether_data::repository::system::AdminSystemUsageAggregateSnapshot, GatewayError>
{
self.data
.export_admin_system_usage_aggregates()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn import_admin_system_usage_aggregates(
&self,
snapshot: &aether_data::repository::system::AdminSystemUsageAggregateSnapshot,
user_id_map: &std::collections::BTreeMap<String, String>,
api_key_id_map: &std::collections::BTreeMap<String, String>,
mode: aether_data::repository::system::AdminSystemUsageAggregateImportMode,
) -> Result<aether_data::repository::system::AdminSystemUsageAggregateImportSummary, GatewayError>
{
self.data
.import_admin_system_usage_aggregates(snapshot, user_id_map, api_key_id_map, mode)
.await
.map_err(|err| match err {
aether_data::DataLayerError::InvalidInput(detail) => GatewayError::Client {
status: http::StatusCode::BAD_REQUEST,
message: detail,
},
other => GatewayError::Internal(other.to_string()),
})
}
pub(crate) async fn run_admin_system_cleanup_once(
&self,
) -> Result<crate::maintenance::AdminSystemCleanupSummary, GatewayError> {
@@ -369,6 +369,25 @@ impl AppState {
Ok(api_key)
}
pub(crate) async fn set_api_key_usage_totals(
&self,
api_key_id: &str,
total_requests: u64,
total_tokens: u64,
total_cost_usd: f64,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
let api_key = self
.data
.set_api_key_usage_totals(api_key_id, total_requests, total_tokens, total_cost_usd)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn set_standalone_api_key_feature_settings(
&self,
api_key_id: &str,
@@ -1043,6 +1043,7 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr
payload["users"][0]["group_names"],
json!(["Restricted GPT"])
);
assert_eq!(payload["users"][0]["id"], json!("user-1"));
assert_eq!(payload["users"][0]["wallet"]["balance"], json!(12.5));
assert_eq!(
payload["users"][0]["wallet"]["recharge_balance"],
@@ -1063,6 +1064,10 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr
payload["users"][0]["api_keys"][0]["is_standalone"],
json!(false)
);
assert_eq!(
payload["users"][0]["api_keys"][0]["api_key_id"],
json!("key-user-1")
);
assert_eq!(
payload["users"][0]["api_keys"][0]["total_tokens"],
json!(420)
@@ -1071,12 +1076,22 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr
payload["standalone_keys"][0]["key"],
json!("ak-standalone-live-1")
);
assert_eq!(
payload["standalone_keys"][0]["api_key_id"],
json!("key-standalone-1")
);
assert_eq!(payload["standalone_keys"][0]["total_tokens"], json!(84));
assert_eq!(
payload["standalone_keys"][0]["wallet"]["unlimited"],
json!(true)
);
assert_eq!(payload["standalone_keys"][0].get("is_standalone"), None,);
assert_eq!(payload["usage_aggregates"]["stats_daily"], json!([]));
assert_eq!(payload["usage_aggregates"]["stats_user_daily"], json!([]));
assert_eq!(
payload["usage_aggregates"]["stats_daily_api_key"],
json!([])
);
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort();
@@ -4,7 +4,9 @@ use aether_contracts::ExecutionPlan;
use aether_crypto::{
decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY,
};
use aether_data::repository::auth::InMemoryAuthApiKeySnapshotRepository;
use aether_data::repository::auth::{
InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot,
};
use aether_data::repository::auth_modules::{
AuthModuleReadRepository, InMemoryAuthModuleReadRepository, StoredOAuthProviderModuleConfig,
};
@@ -27,7 +29,7 @@ use axum::{extract::Request, Json, Router};
use http::StatusCode;
use serde_json::{json, Value};
use super::super::helpers::{sample_endpoint, sample_key, sample_provider};
use super::super::helpers::{hash_api_key, sample_endpoint, sample_key, sample_provider};
use super::super::{
build_router_with_state, build_state_with_execution_runtime_override, start_server, AppState,
};
@@ -1052,6 +1054,7 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data() {
.expect("user api keys should load");
assert_eq!(user_api_keys.len(), 1);
assert_eq!(user_api_keys[0].name.as_deref(), Some("Alice CLI"));
assert_eq!(user_api_keys[0].total_requests, 12);
assert_eq!(user_api_keys[0].total_tokens, 3456);
assert_eq!(user_api_keys[0].total_cost_usd, 1.25);
assert_eq!(
@@ -1079,6 +1082,7 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data() {
standalone_keys[0].name.as_deref(),
Some("Imported Standalone")
);
assert_eq!(standalone_keys[0].total_requests, 3);
assert_eq!(standalone_keys[0].total_tokens, 789);
assert_eq!(standalone_keys[0].total_cost_usd, 0.75);
assert_eq!(
@@ -1119,6 +1123,211 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data() {
let _ = upstream_url;
}
#[tokio::test]
async fn gateway_overwrites_existing_admin_system_user_key_usage_totals() {
let user_key_hash = hash_api_key("sk-existing-user-key");
let standalone_key_hash = hash_api_key("sk-existing-standalone-key");
let existing_user = StoredUserAuthRecord::new(
"user-existing".to_string(),
Some("existing@example.com".to_string()),
true,
"existing".to_string(),
Some("existing-hash".to_string()),
"user".to_string(),
"local".to_string(),
None,
None,
None,
true,
false,
Some(chrono::Utc::now()),
Some(chrono::Utc::now()),
)
.expect("existing user should build");
let user_key_snapshot = StoredAuthApiKeySnapshot::new(
"user-existing".to_string(),
"existing".to_string(),
Some("existing@example.com".to_string()),
"user".to_string(),
"local".to_string(),
true,
false,
None,
None,
None,
"key-user-existing".to_string(),
Some("Existing User Key".to_string()),
true,
false,
false,
Some(10),
None,
None,
None,
None,
None,
)
.expect("user key snapshot should build");
let standalone_key_snapshot = StoredAuthApiKeySnapshot::new(
"admin-user-123".to_string(),
"admin".to_string(),
Some("admin@example.com".to_string()),
"admin".to_string(),
"local".to_string(),
true,
false,
None,
None,
None,
"key-standalone-existing".to_string(),
Some("Existing Standalone Key".to_string()),
true,
false,
true,
Some(20),
None,
None,
None,
None,
None,
)
.expect("standalone key snapshot should build");
let auth_repository = Arc::new(
InMemoryAuthApiKeySnapshotRepository::seed(vec![
(Some(user_key_hash.clone()), user_key_snapshot),
(Some(standalone_key_hash.clone()), standalone_key_snapshot),
])
.with_export_records(vec![
StoredAuthApiKeyExportRecord::new(
"user-existing".to_string(),
"key-user-existing".to_string(),
user_key_hash.clone(),
None,
Some("Existing User Key".to_string()),
None,
None,
None,
Some(10),
None,
None,
true,
None,
false,
1,
2,
0.03,
false,
)
.expect("existing user key export should build"),
StoredAuthApiKeyExportRecord::new(
"admin-user-123".to_string(),
"key-standalone-existing".to_string(),
standalone_key_hash.clone(),
None,
Some("Existing Standalone Key".to_string()),
None,
None,
None,
Some(20),
None,
None,
true,
None,
false,
4,
5,
0.06,
true,
)
.expect("existing standalone key export should build"),
]),
);
let user_repository =
Arc::new(aether_data::repository::users::InMemoryUserReadRepository::default());
let state = AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(
GatewayDataState::with_auth_api_key_repository_for_tests(Arc::clone(&auth_repository))
.with_user_reader(user_repository)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
)
.with_auth_users_for_tests([sample_import_admin_user("admin-user-123"), existing_user])
.with_auth_wallets_for_tests(Vec::<StoredWalletSnapshot>::new());
let gateway = build_router_with_state(state.clone());
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!("{gateway_url}/api/admin/system/users/import"))
.header(GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"version": "1.4",
"merge_mode": "overwrite",
"users": [{
"id": "source-user-existing",
"email": "existing@example.com",
"username": "existing",
"password_hash": "existing-hash",
"role": "user",
"is_active": true,
"api_keys": [{
"api_key_id": "source-user-key",
"key_hash": user_key_hash,
"name": "Imported User Key",
"is_active": true,
"total_requests": 222,
"total_tokens": 3333,
"total_cost_usd": 4.56
}]
}],
"standalone_keys": [{
"api_key_id": "source-standalone-key",
"key_hash": standalone_key_hash,
"name": "Imported Standalone Key",
"is_active": true,
"total_requests": 444,
"total_tokens": 5555,
"total_cost_usd": 6.78
}]
}))
.send()
.await
.expect("request should succeed");
let status = response.status();
let payload: Value = response.json().await.expect("json body should parse");
assert_eq!(status, StatusCode::OK, "payload={payload}");
assert_eq!(payload["stats"]["users"]["updated"], json!(1));
assert_eq!(payload["stats"]["api_keys"]["updated"], json!(1));
assert_eq!(payload["stats"]["standalone_keys"]["updated"], json!(1));
let updated_records = state
.list_auth_api_key_export_records_by_ids(&[
"key-user-existing".to_string(),
"key-standalone-existing".to_string(),
])
.await
.expect("api key export records should load");
let user_key = updated_records
.iter()
.find(|record| record.api_key_id == "key-user-existing")
.expect("updated user key should exist");
assert_eq!(user_key.total_requests, 222);
assert_eq!(user_key.total_tokens, 3333);
assert_eq!(user_key.total_cost_usd, 4.56);
let standalone_key = updated_records
.iter()
.find(|record| record.api_key_id == "key-standalone-existing")
.expect("updated standalone key should exist");
assert_eq!(standalone_key.total_requests, 444);
assert_eq!(standalone_key.total_tokens, 5555);
assert_eq!(standalone_key.total_cost_usd, 6.78);
gateway_handle.abort();
}
#[tokio::test]
async fn gateway_imports_admin_system_config_fixture_v22() {
let gateway = build_router_with_state(
@@ -1671,33 +1880,20 @@ async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_imp
}
#[tokio::test]
async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_import_and_forces_refresh(
async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_import_without_refresh(
) {
#[derive(Debug, Clone)]
struct SeenRefreshRequest {
content_type: String,
body: String,
}
let seen_refresh = Arc::new(Mutex::new(None::<SeenRefreshRequest>));
let seen_refresh = Arc::new(Mutex::new(false));
let seen_refresh_clone = Arc::clone(&seen_refresh);
let refresh_hits = Arc::new(Mutex::new(0usize));
let refresh_hits_clone = Arc::clone(&refresh_hits);
let refresh_server = Router::new().route(
"/oauth/token",
post(move |headers: HeaderMap, body: Bytes| {
post(move |_headers: HeaderMap, _body: Bytes| {
let seen_refresh_inner = Arc::clone(&seen_refresh_clone);
let refresh_hits_inner = Arc::clone(&refresh_hits_clone);
async move {
*refresh_hits_inner.lock().expect("mutex should lock") += 1;
*seen_refresh_inner.lock().expect("mutex should lock") = Some(SeenRefreshRequest {
content_type: headers
.get(http::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.unwrap_or_default()
.to_string(),
body: String::from_utf8(body.to_vec()).unwrap_or_default(),
});
*seen_refresh_inner.lock().expect("mutex should lock") = true;
axum::Json(json!({
"access_token": "oauth-access-token-refreshed",
"refresh_token": "oauth-refresh-token-refreshed",
@@ -1792,21 +1988,8 @@ async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_imp
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(*refresh_hits.lock().expect("mutex should lock"), 1);
let seen_refresh = seen_refresh
.lock()
.expect("mutex should lock")
.clone()
.expect("refresh request should be captured");
assert_eq!(
seen_refresh.content_type,
"application/x-www-form-urlencoded"
);
assert!(seen_refresh.body.contains("grant_type=refresh_token"));
assert!(seen_refresh
.body
.contains("refresh_token=oauth-refresh-token-new"));
assert_eq!(*refresh_hits.lock().expect("mutex should lock"), 0);
assert!(!*seen_refresh.lock().expect("mutex should lock"));
let providers = provider_catalog_repository
.list_providers(false)
@@ -1823,7 +2006,7 @@ async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_imp
assert_eq!(key.name, "oauth-primary");
assert_eq!(key.oauth_invalid_at_unix_secs, None);
assert_eq!(key.oauth_invalid_reason, None);
assert!(key.expires_at_unix_secs.is_some());
assert_eq!(key.expires_at_unix_secs, None);
assert_eq!(
decrypt_python_fernet_ciphertext(
DEVELOPMENT_ENCRYPTION_KEY,
@@ -1832,7 +2015,7 @@ async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_imp
.expect("api key should be present"),
)
.expect("oauth access token should decrypt"),
"oauth-access-token-refreshed"
"oauth-access-token-new"
);
let auth_config = decrypt_python_fernet_ciphertext(
@@ -1845,15 +2028,12 @@ async fn gateway_overwrites_oauth_provider_key_credentials_from_admin_system_imp
let auth_config: Value =
serde_json::from_str(&auth_config).expect("oauth auth config json should parse");
assert_eq!(auth_config["provider_type"], "codex");
assert_eq!(
auth_config["refresh_token"],
"oauth-refresh-token-refreshed"
);
assert_eq!(auth_config["refresh_token"], "oauth-refresh-token-new");
assert_eq!(auth_config["email"], "alice@example.com");
assert_eq!(auth_config["account_id"], "acct-codex-123");
assert_eq!(auth_config["plan_type"], "plus");
assert_eq!(auth_config["token_type"], "Bearer");
assert_eq!(auth_config["expires_at"].as_u64(), key.expires_at_unix_secs);
assert!(auth_config.get("token_type").is_none());
assert!(auth_config.get("expires_at").is_none());
gateway_handle.abort();
refresh_handle.abort();