mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 16:37:46 +08:00
perf(gateway): scale request hot paths for 20k streams
Shard and singleflight hot-path caches, batch and prioritize candidate and usage lifecycle persistence, and extend database and pressure-test instrumentation for 20k concurrent streams.
This commit is contained in:
@@ -191,8 +191,7 @@ pub(super) async fn build_admin_monitoring_system_status_response(
|
||||
.unwrap_or(usize::MAX);
|
||||
let tunnel = state.tunnel.stats();
|
||||
let usage_counter_snapshot = state
|
||||
.data
|
||||
.read_usage_counter_health()
|
||||
.read_cached_usage_counter_health()
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let usage_counter =
|
||||
|
||||
@@ -29,8 +29,7 @@ async fn build_usage_counter_health_payload(
|
||||
let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
|
||||
let snapshot = state
|
||||
.as_ref()
|
||||
.data
|
||||
.read_usage_counter_health()
|
||||
.read_cached_usage_counter_health()
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
Ok(build_admin_usage_counter_health_payload(
|
||||
|
||||
+12
-5
@@ -2,6 +2,7 @@ use crate::handlers::admin::provider::shared::paths::admin_reset_cycle_stats_key
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::provider_key_status_snapshot_payload;
|
||||
use crate::GatewayError;
|
||||
use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyStatusSnapshotUpdate;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -33,7 +34,7 @@ pub(super) async fn maybe_handle(
|
||||
let Some(key_id) = admin_reset_cycle_stats_key_id(request_context.path()) else {
|
||||
return Ok(Some(not_found_response("Key 不存在")));
|
||||
};
|
||||
let Some(mut key) = state
|
||||
let Some(key) = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
@@ -65,11 +66,17 @@ pub(super) async fn maybe_handle(
|
||||
return Ok(Some(bad_request_response("当前账号没有可重置的周期窗口")));
|
||||
}
|
||||
|
||||
key.status_snapshot = Some(status_snapshot);
|
||||
key.updated_at_unix_secs = Some(now_unix_secs);
|
||||
let Some(_) = state.update_provider_catalog_key(&key).await? else {
|
||||
let quota = status_snapshot.get("quota").cloned().unwrap_or(Value::Null);
|
||||
if !state
|
||||
.update_provider_catalog_key_status_snapshot(&ProviderCatalogKeyStatusSnapshotUpdate {
|
||||
key_id: key.id.clone(),
|
||||
status_snapshot_patch: json!({"quota":quota}),
|
||||
updated_at_unix_secs: Some(now_unix_secs),
|
||||
})
|
||||
.await?
|
||||
{
|
||||
return Ok(None);
|
||||
};
|
||||
}
|
||||
|
||||
Ok(Some(
|
||||
Json(json!({
|
||||
|
||||
@@ -82,9 +82,22 @@ pub(super) async fn maybe_handle(
|
||||
Ok(record) => record,
|
||||
Err(detail) => return Ok(Some(bad_request_response(detail))),
|
||||
};
|
||||
let Some(updated) = state.update_provider_catalog_key(&updated_record).await? else {
|
||||
let Some(mut updated) = state.update_provider_catalog_key(&updated_record).await? else {
|
||||
return Ok(None);
|
||||
};
|
||||
if updated_record.learned_rpm_limit != existing_key.learned_rpm_limit {
|
||||
let Some(reloaded) = state
|
||||
.set_provider_catalog_key_learned_rpm_limit(
|
||||
&key_id,
|
||||
updated_record.learned_rpm_limit,
|
||||
updated_record.updated_at_unix_secs,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
updated = reloaded;
|
||||
}
|
||||
let should_overwrite_allowed_models_immediately =
|
||||
admin_provider_key_update_requires_immediate_model_fetch(&existing_key, &updated);
|
||||
let updated = if should_overwrite_allowed_models_immediately {
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
use super::super::super::errors::build_internal_control_error_response;
|
||||
use super::super::super::provisioning::provider_oauth_token_payload_expires_at_unix_secs;
|
||||
use super::super::super::provisioning::{
|
||||
provider_oauth_token_payload_expires_at_unix_secs, seed_provider_oauth_pool_score,
|
||||
};
|
||||
use super::super::super::runtime::{
|
||||
resolve_provider_oauth_runtime_endpoints,
|
||||
spawn_provider_oauth_account_state_refresh_after_update,
|
||||
@@ -220,6 +222,25 @@ pub(super) async fn handle_admin_provider_oauth_complete_key(
|
||||
"Key 不存在",
|
||||
));
|
||||
}
|
||||
if !state
|
||||
.clear_provider_catalog_key_oauth_invalid_marker(&key_id)
|
||||
.await?
|
||||
{
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Key 不存在",
|
||||
));
|
||||
}
|
||||
let Some(recovered_key) = state
|
||||
.reset_provider_catalog_key_recovery_state(&key_id)
|
||||
.await?
|
||||
else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Key 不存在",
|
||||
));
|
||||
};
|
||||
seed_provider_oauth_pool_score(state, &provider.id, &recovered_key, now_unix_secs).await;
|
||||
|
||||
spawn_provider_oauth_account_state_refresh_after_update(
|
||||
state.cloned_app(),
|
||||
|
||||
@@ -2,11 +2,16 @@ use super::state::{
|
||||
decode_jwt_claims, enrich_admin_provider_oauth_auth_config, json_non_empty_string,
|
||||
json_u64_value,
|
||||
};
|
||||
use crate::ai_serving::{
|
||||
build_provider_key_pool_score_upsert, provider_key_pool_score_id, provider_key_pool_score_scope,
|
||||
};
|
||||
use crate::handlers::admin::admin_provider_pool_config;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::maintenance::ensure_provider_key_pool_scores_for_keys;
|
||||
use crate::provider_key_auth::provider_active_api_formats;
|
||||
use crate::GatewayError;
|
||||
use aether_data_contracts::repository::pool_scores::{
|
||||
GetPoolMemberScoresByIdsQuery, PoolMemberIdentity,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
};
|
||||
@@ -211,14 +216,22 @@ pub(crate) async fn update_existing_provider_oauth_catalog_key(
|
||||
if updated.fingerprint.is_none() {
|
||||
updated.fingerprint = grok_oauth_catalog_key_fingerprint(provider_type, auth_config);
|
||||
}
|
||||
updated.health_by_format = Some(json!({}));
|
||||
updated.circuit_breaker_by_format = Some(json!({}));
|
||||
updated.error_count = Some(0);
|
||||
if let Some(proxy) = proxy {
|
||||
updated.proxy = Some(proxy);
|
||||
}
|
||||
updated.updated_at_unix_secs = Some(now_unix_secs);
|
||||
let persisted = state.update_provider_catalog_key(&updated).await?;
|
||||
if state.update_provider_catalog_key(&updated).await?.is_none() {
|
||||
return Ok(None);
|
||||
}
|
||||
if !state
|
||||
.clear_provider_catalog_key_oauth_invalid_marker(&updated.id)
|
||||
.await?
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
let persisted = state
|
||||
.reset_provider_catalog_key_recovery_state(&updated.id)
|
||||
.await?;
|
||||
if let Some(key) = persisted.as_ref() {
|
||||
let _ = state
|
||||
.app()
|
||||
@@ -229,7 +242,7 @@ pub(crate) async fn update_existing_provider_oauth_catalog_key(
|
||||
Ok(persisted)
|
||||
}
|
||||
|
||||
async fn seed_provider_oauth_pool_score(
|
||||
pub(super) async fn seed_provider_oauth_pool_score(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
key: &StoredProviderCatalogKey,
|
||||
@@ -257,38 +270,45 @@ async fn seed_provider_oauth_pool_score(
|
||||
let Some(pool_config) = admin_provider_pool_config(&provider) else {
|
||||
return;
|
||||
};
|
||||
let endpoints = match state
|
||||
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
|
||||
if !key.is_active || key.provider_id != provider.id {
|
||||
return;
|
||||
}
|
||||
|
||||
let identity = PoolMemberIdentity::provider_api_key(provider.id.clone(), key.id.clone());
|
||||
let scope = provider_key_pool_score_scope();
|
||||
let score_id = provider_key_pool_score_id(&identity, &scope);
|
||||
let existing = match state
|
||||
.app()
|
||||
.data
|
||||
.get_pool_member_scores_by_ids(&GetPoolMemberScoresByIdsQuery {
|
||||
ids: vec![score_id],
|
||||
})
|
||||
.await
|
||||
{
|
||||
Ok(endpoints) => endpoints,
|
||||
Ok(mut scores) => scores.pop(),
|
||||
Err(err) => {
|
||||
tracing::debug!(
|
||||
provider_id = %provider_id,
|
||||
key_id = %key.id,
|
||||
error = ?err,
|
||||
"gateway provider oauth provisioning: failed to read endpoints for pool score seed"
|
||||
"gateway provider oauth provisioning: failed to read existing pool score"
|
||||
);
|
||||
return;
|
||||
}
|
||||
};
|
||||
let score_ensure_budget = (pool_config.score_fallback_scan_limit as usize).clamp(1, 50_000);
|
||||
if let Err(err) = ensure_provider_key_pool_scores_for_keys(
|
||||
state.as_ref(),
|
||||
&provider,
|
||||
&pool_config,
|
||||
&endpoints,
|
||||
std::slice::from_ref(key),
|
||||
let upsert = build_provider_key_pool_score_upsert(
|
||||
key,
|
||||
provider.provider_type.as_str(),
|
||||
existing.as_ref(),
|
||||
now_unix_secs,
|
||||
score_ensure_budget,
|
||||
)
|
||||
.await
|
||||
{
|
||||
pool_config.score_rules,
|
||||
);
|
||||
if let Err(err) = state.app().data.upsert_pool_member_score(upsert).await {
|
||||
tracing::debug!(
|
||||
provider_id = %provider_id,
|
||||
key_id = %key.id,
|
||||
error = ?err,
|
||||
"gateway provider oauth provisioning: failed to seed pool score row"
|
||||
"gateway provider oauth provisioning: failed to refresh pool score row"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -14,6 +14,7 @@ use aether_contracts::{
|
||||
ResolvedTransportProfile, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
};
|
||||
use aether_provider_pool::{ProviderPoolQuotaRequestSpec, ProviderPoolService};
|
||||
@@ -288,6 +289,30 @@ pub(crate) async fn persist_provider_quota_refresh_state(
|
||||
oauth_invalid_reason: Option<String>,
|
||||
encrypted_auth_config: Option<String>,
|
||||
) -> Result<bool, GatewayError> {
|
||||
persist_provider_quota_refresh_state_after_read(
|
||||
state,
|
||||
key_id,
|
||||
metadata_update,
|
||||
oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason,
|
||||
encrypted_auth_config,
|
||||
std::future::ready(()),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn persist_provider_quota_refresh_state_after_read<F>(
|
||||
state: &AdminAppState<'_>,
|
||||
key_id: &str,
|
||||
metadata_update: Option<&serde_json::Value>,
|
||||
oauth_invalid_at_unix_secs: Option<u64>,
|
||||
oauth_invalid_reason: Option<String>,
|
||||
encrypted_auth_config: Option<String>,
|
||||
after_read: F,
|
||||
) -> Result<bool, GatewayError>
|
||||
where
|
||||
F: std::future::Future<Output = ()>,
|
||||
{
|
||||
let Some(mut latest_key) = state
|
||||
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
|
||||
.await?
|
||||
@@ -296,7 +321,11 @@ pub(crate) async fn persist_provider_quota_refresh_state(
|
||||
else {
|
||||
return Ok(false);
|
||||
};
|
||||
after_read.await;
|
||||
|
||||
// Keep the namespace values observed before applying the refresh response;
|
||||
// each runtime metadata write uses them as its CAS expectation.
|
||||
let observed_upstream_metadata = latest_key.upstream_metadata.clone();
|
||||
let mut quota_snapshot_provider_type = None::<String>;
|
||||
if let Some(metadata_update) = metadata_update {
|
||||
latest_key.upstream_metadata = Some(merge_upstream_metadata(
|
||||
@@ -306,8 +335,8 @@ pub(crate) async fn persist_provider_quota_refresh_state(
|
||||
quota_snapshot_provider_type =
|
||||
aether_provider_pool::provider_pool_quota_metadata_provider_type(metadata_update);
|
||||
}
|
||||
if let Some(encrypted_auth_config) = encrypted_auth_config {
|
||||
latest_key.encrypted_auth_config = Some(encrypted_auth_config);
|
||||
if let Some(encrypted_auth_config) = encrypted_auth_config.as_ref() {
|
||||
latest_key.encrypted_auth_config = Some(encrypted_auth_config.clone());
|
||||
}
|
||||
latest_key.oauth_invalid_at_unix_secs = oauth_invalid_at_unix_secs;
|
||||
latest_key.oauth_invalid_reason = oauth_invalid_reason;
|
||||
@@ -325,10 +354,91 @@ pub(crate) async fn persist_provider_quota_refresh_state(
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs());
|
||||
Ok(state
|
||||
.update_provider_catalog_key(&latest_key)
|
||||
.await?
|
||||
.is_some())
|
||||
let status_patch = provider_quota_refresh_status_patch(latest_key.status_snapshot.as_ref());
|
||||
let metadata_updates = metadata_update
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.map(|updates| {
|
||||
updates
|
||||
.iter()
|
||||
.map(|(namespace, value)| (namespace.clone(), value.clone()))
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.unwrap_or_default();
|
||||
if metadata_updates.is_empty() {
|
||||
if !state
|
||||
.update_provider_catalog_key_oauth_runtime_state(
|
||||
key_id,
|
||||
latest_key.oauth_invalid_at_unix_secs,
|
||||
latest_key.oauth_invalid_reason.as_deref(),
|
||||
encrypted_auth_config.as_deref(),
|
||||
latest_key.updated_at_unix_secs,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
return Ok(false);
|
||||
}
|
||||
return state
|
||||
.update_provider_catalog_key_status_snapshot(&ProviderCatalogKeyStatusSnapshotUpdate {
|
||||
key_id: key_id.to_string(),
|
||||
status_snapshot_patch: status_patch,
|
||||
updated_at_unix_secs: latest_key.updated_at_unix_secs,
|
||||
})
|
||||
.await;
|
||||
}
|
||||
|
||||
for (index, (namespace, value)) in metadata_updates.iter().enumerate() {
|
||||
let patch = if index + 1 == metadata_updates.len() {
|
||||
status_patch.clone()
|
||||
} else {
|
||||
serde_json::json!({})
|
||||
};
|
||||
let mut expected = observed_upstream_metadata
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|metadata| metadata.get(namespace))
|
||||
.cloned();
|
||||
let persisted = state
|
||||
.app()
|
||||
.update_provider_catalog_key_runtime_metadata(
|
||||
&ProviderCatalogKeyRuntimeMetadataUpdate {
|
||||
key_id: key_id.to_string(),
|
||||
namespace: namespace.clone(),
|
||||
expected_upstream_metadata_value: expected.clone(),
|
||||
upstream_metadata_value: value.clone(),
|
||||
status_snapshot_patch: patch.clone(),
|
||||
updated_at_unix_secs: latest_key.updated_at_unix_secs,
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
if !persisted {
|
||||
// The refresh response is an authoritative snapshot. Do not
|
||||
// replay it over a newer namespace after a CAS conflict.
|
||||
return Ok(false);
|
||||
}
|
||||
}
|
||||
state
|
||||
.update_provider_catalog_key_oauth_runtime_state(
|
||||
key_id,
|
||||
latest_key.oauth_invalid_at_unix_secs,
|
||||
latest_key.oauth_invalid_reason.as_deref(),
|
||||
encrypted_auth_config.as_deref(),
|
||||
latest_key.updated_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn provider_quota_refresh_status_patch(
|
||||
status_snapshot: Option<&serde_json::Value>,
|
||||
) -> serde_json::Value {
|
||||
let mut patch = serde_json::Map::new();
|
||||
if let Some(snapshot) = status_snapshot.and_then(serde_json::Value::as_object) {
|
||||
for field in ["quota", "oauth"] {
|
||||
if let Some(value) = snapshot.get(field) {
|
||||
patch.insert(field.to_string(), value.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
serde_json::Value::Object(patch)
|
||||
}
|
||||
|
||||
pub(super) async fn execute_provider_quota_plan(
|
||||
@@ -372,3 +482,95 @@ pub(super) async fn execute_provider_quota_plan(
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::AppState;
|
||||
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogReadRepository, ProviderCatalogWriteRepository, StoredProviderCatalogKey,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::sync::Arc;
|
||||
|
||||
#[tokio::test]
|
||||
async fn metadata_cas_conflict_does_not_persist_stale_oauth_runtime_state() {
|
||||
let mut key = StoredProviderCatalogKey::new(
|
||||
"key-codex-cas".to_string(),
|
||||
"provider-codex-cas".to_string(),
|
||||
"Codex CAS".to_string(),
|
||||
"oauth".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build");
|
||||
key.encrypted_auth_config = Some("old-auth-config".to_string());
|
||||
key.oauth_invalid_at_unix_secs = Some(100);
|
||||
key.oauth_invalid_reason = Some("old-invalid-reason".to_string());
|
||||
key.upstream_metadata = Some(json!({"codex":{"remaining":5}}));
|
||||
key.status_snapshot = Some(json!({"oauth":{"invalid":true}}));
|
||||
|
||||
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![],
|
||||
vec![],
|
||||
vec![key],
|
||||
));
|
||||
let app = AppState::new()
|
||||
.expect("app should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_repository_for_tests(Arc::clone(
|
||||
&repository,
|
||||
)),
|
||||
);
|
||||
let admin_state = AdminAppState::new(&app);
|
||||
let concurrent_repository = Arc::clone(&repository);
|
||||
let metadata_update = json!({"codex":{"remaining":3}});
|
||||
|
||||
let persisted = persist_provider_quota_refresh_state_after_read(
|
||||
&admin_state,
|
||||
"key-codex-cas",
|
||||
Some(&metadata_update),
|
||||
Some(200),
|
||||
Some("new-invalid-reason".to_string()),
|
||||
Some("new-auth-config".to_string()),
|
||||
async move {
|
||||
assert!(concurrent_repository
|
||||
.update_key_runtime_metadata(&ProviderCatalogKeyRuntimeMetadataUpdate {
|
||||
key_id: "key-codex-cas".to_string(),
|
||||
namespace: "codex".to_string(),
|
||||
expected_upstream_metadata_value: Some(json!({"remaining":5})),
|
||||
upstream_metadata_value: json!({"remaining":4}),
|
||||
status_snapshot_patch: json!({}),
|
||||
updated_at_unix_secs: Some(150),
|
||||
})
|
||||
.await
|
||||
.expect("concurrent metadata update should execute"));
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("quota refresh persistence should not error");
|
||||
|
||||
assert!(!persisted, "stale namespace should report a CAS conflict");
|
||||
let stored = repository
|
||||
.list_keys_by_ids(&["key-codex-cas".to_string()])
|
||||
.await
|
||||
.expect("key should reload")
|
||||
.pop()
|
||||
.expect("key should remain");
|
||||
assert_eq!(
|
||||
stored.encrypted_auth_config.as_deref(),
|
||||
Some("old-auth-config")
|
||||
);
|
||||
assert_eq!(stored.oauth_invalid_at_unix_secs, Some(100));
|
||||
assert_eq!(
|
||||
stored.oauth_invalid_reason.as_deref(),
|
||||
Some("old-invalid-reason")
|
||||
);
|
||||
assert_eq!(
|
||||
stored.upstream_metadata.as_ref().unwrap()["codex"],
|
||||
json!({"remaining":4})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -140,6 +140,15 @@ impl<'a> AdminAppState<'a> {
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn reset_provider_catalog_key_error_count(
|
||||
&self,
|
||||
key_id: &str,
|
||||
) -> Result<bool, GatewayError> {
|
||||
self.app
|
||||
.reset_provider_catalog_key_error_count(key_id)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn create_provider_catalog_endpoint(
|
||||
&self,
|
||||
endpoint: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogEndpoint,
|
||||
@@ -177,6 +186,139 @@ impl<'a> AdminAppState<'a> {
|
||||
self.app.update_provider_catalog_key(key).await
|
||||
}
|
||||
|
||||
pub(crate) async fn compare_and_update_provider_catalog_key_adaptive_state(
|
||||
&self,
|
||||
update: &aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyAdaptiveStateUpdate,
|
||||
) -> Result<bool, GatewayError> {
|
||||
self.app
|
||||
.compare_and_update_provider_catalog_key_adaptive_state(update)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn set_provider_catalog_key_learned_rpm_limit(
|
||||
&self,
|
||||
key_id: &str,
|
||||
learned_rpm_limit: Option<u32>,
|
||||
updated_at_unix_secs: Option<u64>,
|
||||
) -> Result<
|
||||
Option<aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey>,
|
||||
GatewayError,
|
||||
> {
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
|
||||
};
|
||||
|
||||
for _ in 0..4 {
|
||||
let Some(current) = self
|
||||
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let expected = ProviderCatalogKeyAdaptiveState::from(¤t);
|
||||
if expected.learned_rpm_limit == learned_rpm_limit {
|
||||
return Ok(Some(current));
|
||||
}
|
||||
let mut next = expected.clone();
|
||||
next.learned_rpm_limit = learned_rpm_limit;
|
||||
if self
|
||||
.compare_and_update_provider_catalog_key_adaptive_state(
|
||||
&ProviderCatalogKeyAdaptiveStateUpdate {
|
||||
key_id: key_id.to_string(),
|
||||
expected,
|
||||
next,
|
||||
status_snapshot_patch: serde_json::json!({
|
||||
"learning_confidence": 0.0,
|
||||
"enforcement_active": false
|
||||
}),
|
||||
updated_at_unix_secs,
|
||||
},
|
||||
)
|
||||
.await?
|
||||
{
|
||||
return Ok(self
|
||||
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
|
||||
.await?
|
||||
.into_iter()
|
||||
.next());
|
||||
}
|
||||
}
|
||||
|
||||
Err(GatewayError::Internal(format!(
|
||||
"provider key {key_id} adaptive state changed repeatedly while updating"
|
||||
)))
|
||||
}
|
||||
|
||||
pub(crate) async fn reset_provider_catalog_key_recovery_state(
|
||||
&self,
|
||||
key_id: &str,
|
||||
) -> Result<
|
||||
Option<aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey>,
|
||||
GatewayError,
|
||||
> {
|
||||
use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyHealthStateUpdate;
|
||||
|
||||
let empty = serde_json::json!({});
|
||||
let mut health_reset = false;
|
||||
for _ in 0..4 {
|
||||
let Some(current) = self
|
||||
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
if current.health_by_format.as_ref() == Some(&empty)
|
||||
&& current.circuit_breaker_by_format.as_ref() == Some(&empty)
|
||||
{
|
||||
health_reset = true;
|
||||
break;
|
||||
}
|
||||
if self
|
||||
.app
|
||||
.compare_and_update_provider_catalog_key_health_state(
|
||||
&ProviderCatalogKeyHealthStateUpdate {
|
||||
key_id: key_id.to_string(),
|
||||
expected_health_by_format: current.health_by_format,
|
||||
expected_circuit_breaker_by_format: current.circuit_breaker_by_format,
|
||||
health_by_format: Some(empty.clone()),
|
||||
circuit_breaker_by_format: Some(empty.clone()),
|
||||
},
|
||||
)
|
||||
.await?
|
||||
{
|
||||
health_reset = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if !health_reset {
|
||||
return Err(GatewayError::Internal(format!(
|
||||
"provider key {key_id} health state changed repeatedly while resetting OAuth recovery state"
|
||||
)));
|
||||
}
|
||||
if !self.reset_provider_catalog_key_error_count(key_id).await? {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(self
|
||||
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
|
||||
.await?
|
||||
.into_iter()
|
||||
.next())
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_key_status_snapshot(
|
||||
&self,
|
||||
update: &aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyStatusSnapshotUpdate,
|
||||
) -> Result<bool, GatewayError> {
|
||||
self.app
|
||||
.update_provider_catalog_key_status_snapshot(update)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_keys(
|
||||
&self,
|
||||
keys: &[aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey],
|
||||
|
||||
@@ -46,6 +46,25 @@ impl<'a> AdminAppState<'a> {
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_key_oauth_runtime_state(
|
||||
&self,
|
||||
key_id: &str,
|
||||
oauth_invalid_at_unix_secs: Option<u64>,
|
||||
oauth_invalid_reason: Option<&str>,
|
||||
encrypted_auth_config_update: Option<&str>,
|
||||
updated_at_unix_secs: Option<u64>,
|
||||
) -> Result<bool, GatewayError> {
|
||||
self.app
|
||||
.update_provider_catalog_key_oauth_runtime_state(
|
||||
key_id,
|
||||
oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason,
|
||||
encrypted_auth_config_update,
|
||||
updated_at_unix_secs,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn clear_provider_catalog_key_oauth_invalid_marker(
|
||||
&self,
|
||||
key_id: &str,
|
||||
|
||||
@@ -749,10 +749,11 @@ impl<'a> AdminAppState<'a> {
|
||||
.collect::<BTreeSet<_>>();
|
||||
|
||||
let staged_records = staged_updates
|
||||
.into_iter()
|
||||
.map(|(_, updated)| updated)
|
||||
.iter()
|
||||
.map(|(_, updated)| updated.clone())
|
||||
.collect::<Vec<_>>();
|
||||
let Some(updated_keys) = self.update_provider_catalog_keys(&staged_records).await? else {
|
||||
let Some(mut updated_keys) = self.update_provider_catalog_keys(&staged_records).await?
|
||||
else {
|
||||
return Ok((
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
Json(json!({ "detail": "Provider 密钥写入能力不可用" })),
|
||||
@@ -760,6 +761,29 @@ impl<'a> AdminAppState<'a> {
|
||||
.into_response());
|
||||
};
|
||||
|
||||
for (existing, requested) in &staged_updates {
|
||||
if requested.learned_rpm_limit == existing.learned_rpm_limit {
|
||||
continue;
|
||||
}
|
||||
let Some(reloaded) = self
|
||||
.set_provider_catalog_key_learned_rpm_limit(
|
||||
&requested.id,
|
||||
requested.learned_rpm_limit,
|
||||
requested.updated_at_unix_secs,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok((
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
Json(json!({ "detail": format!("Provider 密钥 {} 已不存在", requested.id) })),
|
||||
)
|
||||
.into_response());
|
||||
};
|
||||
if let Some(updated) = updated_keys.iter_mut().find(|key| key.id == requested.id) {
|
||||
*updated = reloaded;
|
||||
}
|
||||
}
|
||||
|
||||
let endpoints = self
|
||||
.list_provider_catalog_endpoints_by_provider_ids(&provider_ids)
|
||||
.await?;
|
||||
|
||||
@@ -7,6 +7,9 @@ use aether_admin::system::{
|
||||
build_admin_adaptive_set_limit_payload, build_admin_adaptive_stats_payload,
|
||||
build_admin_adaptive_summary_payload, build_admin_adaptive_toggle_mode_payload,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
|
||||
};
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -130,22 +133,49 @@ impl<'a> AdminAppState<'a> {
|
||||
&self,
|
||||
key_id: &str,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let Some(mut key) = self.find_admin_adaptive_key(key_id).await? else {
|
||||
return Ok(admin_adaptive_key_not_found_response(key_id));
|
||||
};
|
||||
key.learned_rpm_limit = None;
|
||||
key.concurrent_429_count = Some(0);
|
||||
key.rpm_429_count = Some(0);
|
||||
key.last_429_at_unix_secs = None;
|
||||
key.last_429_type = None;
|
||||
key.adjustment_history = None;
|
||||
key.utilization_samples = None;
|
||||
key.last_probe_increase_at_unix_secs = None;
|
||||
key.last_rpm_peak = None;
|
||||
let Some(updated) = self.update_provider_catalog_key(&key).await? else {
|
||||
return Ok(admin_adaptive_key_not_found_response(key_id));
|
||||
};
|
||||
Ok(Json(build_admin_adaptive_reset_learning_payload(&updated.id)).into_response())
|
||||
for _ in 0..4 {
|
||||
let Some(key) = self.find_admin_adaptive_key(key_id).await? else {
|
||||
return Ok(admin_adaptive_key_not_found_response(key_id));
|
||||
};
|
||||
let expected = ProviderCatalogKeyAdaptiveState::from(&key);
|
||||
let next = ProviderCatalogKeyAdaptiveState {
|
||||
learned_rpm_limit: None,
|
||||
concurrent_429_count: Some(0),
|
||||
rpm_429_count: Some(0),
|
||||
last_429_at_unix_secs: None,
|
||||
last_429_type: None,
|
||||
adjustment_history: None,
|
||||
utilization_samples: None,
|
||||
last_probe_increase_at_unix_secs: None,
|
||||
last_rpm_peak: None,
|
||||
};
|
||||
if self
|
||||
.compare_and_update_provider_catalog_key_adaptive_state(
|
||||
&ProviderCatalogKeyAdaptiveStateUpdate {
|
||||
key_id: key.id.clone(),
|
||||
expected,
|
||||
next,
|
||||
status_snapshot_patch: json!({
|
||||
"observation_count": 0,
|
||||
"header_observation_count": 0,
|
||||
"latest_upstream_limit": null,
|
||||
"learning_confidence": 0.0,
|
||||
"enforcement_active": false,
|
||||
"known_boundary": null
|
||||
}),
|
||||
updated_at_unix_secs: None,
|
||||
},
|
||||
)
|
||||
.await?
|
||||
{
|
||||
return Ok(
|
||||
Json(build_admin_adaptive_reset_learning_payload(&key.id)).into_response()
|
||||
);
|
||||
}
|
||||
}
|
||||
Err(GatewayError::Internal(format!(
|
||||
"provider key {key_id} adaptive state changed repeatedly while resetting learning"
|
||||
)))
|
||||
}
|
||||
|
||||
pub(crate) fn admin_adaptive_dispatcher_not_found_response(&self) -> Response<Body> {
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
use super::{
|
||||
AdminAppState, ADMIN_SYSTEM_DATA_EXPORT_VERSION, ADMIN_SYSTEM_DATA_IMPORT_MAX_SIZE_BYTES,
|
||||
};
|
||||
use crate::ai_serving::build_provider_key_pool_score_upsert;
|
||||
use crate::api::ai::admin_endpoint_signature_parts;
|
||||
use crate::handlers::admin::admin_provider_pool_config;
|
||||
use crate::handlers::admin::provider::endpoints_admin::payloads::AdminProviderEndpointUpdatePatch;
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
AdminProviderCreateRequest, AdminProviderKeyCreateRequest, AdminProviderKeyUpdatePatch,
|
||||
@@ -47,6 +49,7 @@ use aether_data_contracts::repository::global_models::{
|
||||
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
|
||||
UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
|
||||
};
|
||||
use aether_data_contracts::repository::pool_scores::PoolMemberScoreUpsertMode;
|
||||
use axum::{body::Bytes, http};
|
||||
use serde_json::{json, Map, Value};
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
@@ -362,18 +365,24 @@ fn apply_imported_oauth_key_credentials(
|
||||
raw_key: &Map<String, Value>,
|
||||
normalized_auth_config: Option<&Value>,
|
||||
record: &mut aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey,
|
||||
) -> Result<(), String> {
|
||||
) -> Result<bool, String> {
|
||||
let mut credentials_supplied = false;
|
||||
let mut api_key_supplied = false;
|
||||
if let Some(api_key_value) = raw_key.get("api_key") {
|
||||
let plaintext = api_key_value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
record.encrypted_api_key = match plaintext {
|
||||
Some(plaintext) => Some(
|
||||
state
|
||||
.encrypt_catalog_secret_with_fallbacks(plaintext)
|
||||
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?,
|
||||
),
|
||||
Some(plaintext) => {
|
||||
credentials_supplied = true;
|
||||
api_key_supplied = true;
|
||||
Some(
|
||||
state
|
||||
.encrypt_catalog_secret_with_fallbacks(plaintext)
|
||||
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?,
|
||||
)
|
||||
}
|
||||
None => None,
|
||||
};
|
||||
}
|
||||
@@ -381,6 +390,7 @@ fn apply_imported_oauth_key_credentials(
|
||||
if raw_key.contains_key("auth_config") {
|
||||
record.encrypted_auth_config = match normalized_auth_config {
|
||||
Some(auth_config) => {
|
||||
credentials_supplied |= imported_oauth_auth_config_has_credentials(auth_config);
|
||||
let plaintext =
|
||||
serde_json::to_string(auth_config).map_err(|err| err.to_string())?;
|
||||
Some(
|
||||
@@ -392,14 +402,66 @@ fn apply_imported_oauth_key_credentials(
|
||||
None => None,
|
||||
};
|
||||
}
|
||||
record.expires_at_unix_secs = imported_oauth_expiry_after_import(
|
||||
record.expires_at_unix_secs,
|
||||
raw_key.contains_key("auth_config"),
|
||||
normalized_auth_config,
|
||||
api_key_supplied,
|
||||
);
|
||||
|
||||
// Importing OAuth credentials replaces the previous session state, so stale
|
||||
// expiry/invalid markers must not survive across the overwrite.
|
||||
record.expires_at_unix_secs = imported_oauth_expires_at_unix_secs(normalized_auth_config);
|
||||
record.oauth_invalid_at_unix_secs = None;
|
||||
record.oauth_invalid_reason = None;
|
||||
if credentials_supplied {
|
||||
record.oauth_invalid_at_unix_secs = None;
|
||||
record.oauth_invalid_reason = None;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
Ok(credentials_supplied)
|
||||
}
|
||||
|
||||
fn imported_oauth_auth_config_has_credentials(value: &Value) -> bool {
|
||||
const CREDENTIAL_FIELDS: &[&str] = &[
|
||||
"access_token",
|
||||
"accessToken",
|
||||
"api_key",
|
||||
"apiKey",
|
||||
"auth_token",
|
||||
"authToken",
|
||||
"cf_clearance",
|
||||
"cfClearance",
|
||||
"cf_cookies",
|
||||
"cfCookies",
|
||||
"cookie",
|
||||
"cookieHeader",
|
||||
"cookies",
|
||||
"id_token",
|
||||
"idToken",
|
||||
"refresh_token",
|
||||
"refreshToken",
|
||||
"session_token",
|
||||
"sessionToken",
|
||||
"sso_rw_token",
|
||||
"ssoRwToken",
|
||||
"sso_token",
|
||||
"ssoToken",
|
||||
"token",
|
||||
];
|
||||
|
||||
match value {
|
||||
Value::Object(object) => object.iter().any(|(key, value)| {
|
||||
(CREDENTIAL_FIELDS.contains(&key.as_str()) && imported_credential_value_present(value))
|
||||
|| imported_oauth_auth_config_has_credentials(value)
|
||||
}),
|
||||
Value::Array(items) => items.iter().any(imported_oauth_auth_config_has_credentials),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn imported_credential_value_present(value: &Value) -> bool {
|
||||
match value {
|
||||
Value::String(value) => !value.trim().is_empty(),
|
||||
Value::Array(items) => !items.is_empty(),
|
||||
Value::Object(object) => !object.is_empty(),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn imported_oauth_expires_at_unix_secs(normalized_auth_config: Option<&Value>) -> Option<u64> {
|
||||
@@ -425,6 +487,63 @@ fn imported_oauth_expires_at_unix_secs(normalized_auth_config: Option<&Value>) -
|
||||
None
|
||||
}
|
||||
|
||||
fn imported_oauth_expiry_after_import(
|
||||
current: Option<u64>,
|
||||
auth_config_present: bool,
|
||||
normalized_auth_config: Option<&Value>,
|
||||
api_key_supplied: bool,
|
||||
) -> Option<u64> {
|
||||
if auth_config_present {
|
||||
imported_oauth_expires_at_unix_secs(normalized_auth_config)
|
||||
} else if api_key_supplied {
|
||||
None
|
||||
} else {
|
||||
current
|
||||
}
|
||||
}
|
||||
|
||||
async fn seed_imported_oauth_pool_score(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
key: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey,
|
||||
now_unix_secs: u64,
|
||||
) -> Result<(), GatewayError> {
|
||||
let provider_id = provider_id.to_string();
|
||||
let provider = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
.await?
|
||||
.pop();
|
||||
let Some(provider) = provider else {
|
||||
return Ok(());
|
||||
};
|
||||
let Some(pool_config) = admin_provider_pool_config(&provider) else {
|
||||
return Ok(());
|
||||
};
|
||||
if !key.is_active || key.provider_id != provider.id {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let upsert = build_provider_key_pool_score_upsert(
|
||||
key,
|
||||
provider.provider_type.as_str(),
|
||||
None,
|
||||
now_unix_secs,
|
||||
pool_config.score_rules,
|
||||
);
|
||||
state
|
||||
.app()
|
||||
.data
|
||||
.upsert_pool_member_score_with_mode(upsert, PoolMemberScoreUpsertMode::OAuthRecovery)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
GatewayError::Internal(format!(
|
||||
"failed to recover OAuth pool score for key '{}': {error}",
|
||||
key.id
|
||||
))
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn build_import_provider_model_record(
|
||||
provider_id: &str,
|
||||
existing_id: Option<&str>,
|
||||
@@ -1528,6 +1647,11 @@ impl<'a> AdminAppState<'a> {
|
||||
.keys()
|
||||
.cloned()
|
||||
.collect::<BTreeSet<_>>();
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
|
||||
let imported_keys = routed!(parse_admin_system_config_nested_array::<
|
||||
ImportedProviderKey,
|
||||
@@ -1638,27 +1762,74 @@ impl<'a> AdminAppState<'a> {
|
||||
)
|
||||
.await
|
||||
);
|
||||
if auth_type == "oauth" {
|
||||
let oauth_credentials_supplied = if auth_type == "oauth" {
|
||||
invalid!(apply_imported_oauth_key_credentials(
|
||||
self,
|
||||
&raw_key,
|
||||
normalized_auth_config.as_ref(),
|
||||
&mut updated,
|
||||
));
|
||||
}
|
||||
))
|
||||
} else {
|
||||
false
|
||||
};
|
||||
updated.proxy =
|
||||
remap_import_proxy(imported_key.proxy.clone(), &node_id_map);
|
||||
updated.fingerprint = invalid!(normalize_json_object(
|
||||
imported_key.fingerprint.clone(),
|
||||
"fingerprint",
|
||||
));
|
||||
let Some(persisted) =
|
||||
let Some(mut persisted) =
|
||||
self.update_provider_catalog_key(&updated).await?
|
||||
else {
|
||||
return Ok(Err(invalid_request(format!(
|
||||
"更新 Provider '{provider_name}' 的 Key 失败"
|
||||
))));
|
||||
};
|
||||
if updated.learned_rpm_limit != existing_key.learned_rpm_limit {
|
||||
let Some(reloaded) = self
|
||||
.set_provider_catalog_key_learned_rpm_limit(
|
||||
&updated.id,
|
||||
updated.learned_rpm_limit,
|
||||
updated.updated_at_unix_secs,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(Err(invalid_request(format!(
|
||||
"更新 Provider '{provider_name}' 的 Key 失败"
|
||||
))));
|
||||
};
|
||||
persisted = reloaded;
|
||||
}
|
||||
if oauth_credentials_supplied {
|
||||
if !self
|
||||
.clear_provider_catalog_key_oauth_invalid_marker(&updated.id)
|
||||
.await?
|
||||
{
|
||||
return Ok(Err(invalid_request(format!(
|
||||
"更新 Provider '{provider_name}' 的 Key 失败"
|
||||
))));
|
||||
}
|
||||
let Some(reloaded) = self
|
||||
.reset_provider_catalog_key_recovery_state(&updated.id)
|
||||
.await?
|
||||
else {
|
||||
return Ok(Err(invalid_request(format!(
|
||||
"更新 Provider '{provider_name}' 的 Key 失败"
|
||||
))));
|
||||
};
|
||||
persisted = reloaded;
|
||||
let _ = self
|
||||
.app()
|
||||
.invalidate_local_oauth_refresh_entry(&updated.id)
|
||||
.await;
|
||||
seed_imported_oauth_pool_score(
|
||||
self,
|
||||
&provider.id,
|
||||
&persisted,
|
||||
now_unix_secs,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
existing_keys[existing_index] = persisted;
|
||||
stats.keys.updated += 1;
|
||||
}
|
||||
@@ -1676,14 +1847,16 @@ impl<'a> AdminAppState<'a> {
|
||||
self.build_admin_create_provider_key_record(&provider, payload)
|
||||
.await
|
||||
);
|
||||
if auth_type == "oauth" {
|
||||
let oauth_credentials_supplied = if auth_type == "oauth" {
|
||||
invalid!(apply_imported_oauth_key_credentials(
|
||||
self,
|
||||
&raw_key,
|
||||
normalized_auth_config.as_ref(),
|
||||
&mut record,
|
||||
));
|
||||
}
|
||||
))
|
||||
} else {
|
||||
false
|
||||
};
|
||||
record.is_active = imported_key.is_active;
|
||||
record.global_priority_by_format = invalid!(normalize_json_object(
|
||||
imported_key.global_priority_by_format.clone(),
|
||||
@@ -1699,6 +1872,10 @@ impl<'a> AdminAppState<'a> {
|
||||
"创建 Provider '{provider_name}' 的 Key 失败"
|
||||
))));
|
||||
};
|
||||
if oauth_credentials_supplied {
|
||||
seed_imported_oauth_pool_score(self, &provider.id, &created, now_unix_secs)
|
||||
.await?;
|
||||
}
|
||||
existing_keys.push(created);
|
||||
stats.keys.created += 1;
|
||||
}
|
||||
@@ -3283,16 +3460,27 @@ enum WalletOwner<'a> {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_data::repository::pool_scores::SqlitePoolMemberScoreRepository;
|
||||
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
build_imported_user_usage_total_aggregates, imported_optional_bool, imported_optional_f64,
|
||||
build_imported_user_usage_total_aggregates, imported_oauth_auth_config_has_credentials,
|
||||
imported_oauth_expiry_after_import, imported_optional_bool, imported_optional_f64,
|
||||
imported_optional_i32, imported_optional_u64, imported_rfc3339_to_unix_secs,
|
||||
imported_string_list_from_value, normalize_import_endpoint_format,
|
||||
normalize_import_key_formats, normalize_import_key_raw_payload,
|
||||
normalize_imported_wallet_target, validate_imported_system_users_export_version,
|
||||
ImportedProviderKey,
|
||||
normalize_imported_wallet_target, seed_imported_oauth_pool_score,
|
||||
validate_imported_system_users_export_version, ImportedProviderKey,
|
||||
};
|
||||
use crate::admin_api::AdminAppState;
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::AppState;
|
||||
|
||||
#[test]
|
||||
fn users_import_requires_supported_export_version() {
|
||||
@@ -3458,6 +3646,123 @@ mod tests {
|
||||
assert_eq!(payload["allow_auth_channel_mismatch_formats"], json!([]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oauth_import_only_treats_non_empty_secret_fields_as_credentials() {
|
||||
assert!(!imported_oauth_auth_config_has_credentials(&json!({})));
|
||||
assert!(!imported_oauth_auth_config_has_credentials(&json!({
|
||||
"provider_type": "codex",
|
||||
"expires_at": 4_102_444_800u64,
|
||||
"account_id": "acct-1",
|
||||
"refresh_token": " "
|
||||
})));
|
||||
assert!(imported_oauth_auth_config_has_credentials(&json!({
|
||||
"provider_type": "codex",
|
||||
"refresh_token": "refresh-1"
|
||||
})));
|
||||
assert!(imported_oauth_auth_config_has_credentials(&json!({
|
||||
"session": {"sso_token": "sso-1"}
|
||||
})));
|
||||
for field in [
|
||||
"sso_rw_token",
|
||||
"ssoRwToken",
|
||||
"cf_cookies",
|
||||
"cfCookies",
|
||||
"cf_clearance",
|
||||
"cfClearance",
|
||||
"cookieHeader",
|
||||
] {
|
||||
let mut config = serde_json::Map::new();
|
||||
config.insert(field.to_string(), json!("credential-1"));
|
||||
assert!(
|
||||
imported_oauth_auth_config_has_credentials(&serde_json::Value::Object(config)),
|
||||
"{field} is transport credential material"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oauth_import_expiry_tracks_the_supplied_credential_source() {
|
||||
let old_expiry = Some(1_700_000_000);
|
||||
assert_eq!(
|
||||
imported_oauth_expiry_after_import(old_expiry, false, None, true),
|
||||
None,
|
||||
"a new top-level api_key replaces the old session and clears its expiry"
|
||||
);
|
||||
assert_eq!(
|
||||
imported_oauth_expiry_after_import(old_expiry, false, None, false),
|
||||
old_expiry,
|
||||
"metadata-only imports preserve the current OAuth expiry"
|
||||
);
|
||||
assert_eq!(
|
||||
imported_oauth_expiry_after_import(
|
||||
old_expiry,
|
||||
true,
|
||||
Some(&json!({"expires_at": 4_102_444_800u64})),
|
||||
false,
|
||||
),
|
||||
Some(4_102_444_800),
|
||||
"an explicit auth_config owns the replacement expiry"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn oauth_pool_score_persistence_failure_is_propagated() {
|
||||
let mut provider = StoredProviderCatalogProvider::new(
|
||||
"provider-1".to_string(),
|
||||
"Provider One".to_string(),
|
||||
None,
|
||||
"codex".to_string(),
|
||||
)
|
||||
.expect("provider should build");
|
||||
provider.config = Some(json!({"pool_advanced": {}}));
|
||||
let key = StoredProviderCatalogKey::new(
|
||||
"key-1".to_string(),
|
||||
provider.id.clone(),
|
||||
"OAuth Key".to_string(),
|
||||
"oauth".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build");
|
||||
let provider_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
Vec::new(),
|
||||
vec![key.clone()],
|
||||
));
|
||||
let no_writer_app = AppState::new()
|
||||
.expect("app state should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_repository_for_tests(Arc::clone(
|
||||
&provider_repository,
|
||||
)),
|
||||
);
|
||||
seed_imported_oauth_pool_score(&AdminAppState::new(&no_writer_app), "provider-1", &key, 99)
|
||||
.await
|
||||
.expect("a disabled score writer remains an allowed no-op");
|
||||
|
||||
let pool = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite pool should connect");
|
||||
let score_repository = Arc::new(SqlitePoolMemberScoreRepository::new(pool.clone()));
|
||||
pool.close().await;
|
||||
let app = AppState::new()
|
||||
.expect("app state should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_repository_for_tests(provider_repository)
|
||||
.with_pool_score_repository_for_tests(score_repository),
|
||||
);
|
||||
|
||||
let error =
|
||||
seed_imported_oauth_pool_score(&AdminAppState::new(&app), "provider-1", &key, 100)
|
||||
.await
|
||||
.expect_err("closed pool must fail OAuth score recovery");
|
||||
assert!(error
|
||||
.into_message()
|
||||
.contains("failed to recover OAuth pool score for key 'key-1'"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn import_handles_legacy_string_scalars() {
|
||||
assert_eq!(
|
||||
|
||||
@@ -536,8 +536,7 @@ pub(crate) async fn build_admin_system_stats_payload(
|
||||
let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
|
||||
let usage_counter_snapshot = state
|
||||
.as_ref()
|
||||
.data
|
||||
.read_usage_counter_health()
|
||||
.read_cached_usage_counter_health()
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let usage_counter =
|
||||
|
||||
Reference in New Issue
Block a user