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:
elky
2026-07-22 02:11:08 +08:00
parent 7756c0913f
commit fc92c4f431
124 changed files with 36325 additions and 3217 deletions
@@ -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(
@@ -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(&current);
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 =