mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-07 18:07:47 +08:00
feat(pool): add bulk key configuration management
This commit is contained in:
@@ -288,6 +288,18 @@ pub(super) fn classify_admin_observability_family_route(
|
||||
"admin:pool",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::PATCH
|
||||
&& normalized_path_no_trailing.starts_with("/api/admin/pool/")
|
||||
&& normalized_path_no_trailing.ends_with("/keys/batch-update")
|
||||
&& normalized_path_no_trailing.matches('/').count() == 6
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"pool_manage",
|
||||
"batch_update_keys",
|
||||
"admin:pool",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
&& normalized_path_no_trailing.starts_with("/api/admin/pool/")
|
||||
&& normalized_path_no_trailing.ends_with("/keys/resolve-selection")
|
||||
|
||||
@@ -82,6 +82,17 @@ fn classifies_admin_pool_provider_key_routes_as_admin_proxy_route() {
|
||||
Some("batch_action_keys")
|
||||
);
|
||||
|
||||
let batch_update_uri: Uri = "/api/admin/pool/provider-1/keys/batch-update"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
let batch_update = classify_control_route(&http::Method::PATCH, &batch_update_uri, &headers)
|
||||
.expect("route should classify");
|
||||
assert_eq!(batch_update.route_family.as_deref(), Some("pool_manage"));
|
||||
assert_eq!(
|
||||
batch_update.route_kind.as_deref(),
|
||||
Some("batch_update_keys")
|
||||
);
|
||||
|
||||
let resolve_selection_uri: Uri = "/api/admin/pool/provider-1/keys/resolve-selection"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
|
||||
@@ -512,6 +512,20 @@ impl GatewayDataState {
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_keys(
|
||||
&self,
|
||||
keys: &[StoredProviderCatalogKey],
|
||||
) -> Result<Option<Vec<StoredProviderCatalogKey>>, DataLayerError> {
|
||||
let updated = match &self.provider_catalog_writer {
|
||||
Some(repository) => repository.update_keys(keys).await.map(Some),
|
||||
None => Ok(None),
|
||||
}?;
|
||||
if updated.as_ref().is_some_and(|keys| !keys.is_empty()) {
|
||||
self.clear_provider_catalog_cache();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_key_upstream_metadata(
|
||||
&self,
|
||||
key_id: &str,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use crate::handlers::admin::admin_provider_pool_config;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_update_key_id;
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdatePatch;
|
||||
use crate::handlers::admin::provider::write::keys::admin_provider_key_update_requires_immediate_model_fetch;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::maintenance::ensure_provider_key_pool_scores_for_keys;
|
||||
use crate::provider_key_auth::provider_key_effective_api_formats;
|
||||
@@ -84,12 +85,8 @@ pub(super) async fn maybe_handle(
|
||||
let Some(updated) = state.update_provider_catalog_key(&updated_record).await? else {
|
||||
return Ok(None);
|
||||
};
|
||||
let auto_fetch_filters_changed = existing_key.model_include_patterns
|
||||
!= updated.model_include_patterns
|
||||
|| existing_key.model_exclude_patterns != updated.model_exclude_patterns;
|
||||
// 自动获取开启后,调整过滤规则也要立即刷新 allowed_models。
|
||||
let should_overwrite_allowed_models_immediately = updated.auto_fetch_models
|
||||
&& (!existing_key.auto_fetch_models || auto_fetch_filters_changed);
|
||||
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 {
|
||||
let summary =
|
||||
perform_model_fetch_for_key(state.as_ref(), &provider.id, &updated.id).await?;
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
use super::{
|
||||
admin_pool_provider_id_from_path, build_admin_pool_error_response,
|
||||
ADMIN_POOL_PROVIDER_CATALOG_READER_UNAVAILABLE_DETAIL,
|
||||
ADMIN_POOL_PROVIDER_CATALOG_WRITER_UNAVAILABLE_DETAIL,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyBatchUpdateRequest;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
response::Response,
|
||||
};
|
||||
|
||||
pub(super) async fn build_admin_pool_batch_update_response(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
if !state.has_provider_catalog_data_reader() {
|
||||
return Ok(build_admin_pool_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
ADMIN_POOL_PROVIDER_CATALOG_READER_UNAVAILABLE_DETAIL,
|
||||
));
|
||||
}
|
||||
if !state.has_provider_catalog_data_writer() {
|
||||
return Ok(build_admin_pool_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
ADMIN_POOL_PROVIDER_CATALOG_WRITER_UNAVAILABLE_DETAIL,
|
||||
));
|
||||
}
|
||||
|
||||
let Some(provider_id) = admin_pool_provider_id_from_path(request_context.path()) else {
|
||||
return Ok(build_admin_pool_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Provider 不存在",
|
||||
));
|
||||
};
|
||||
let payload = match request_body.filter(|body| !body.is_empty()) {
|
||||
Some(body) => match serde_json::from_slice::<AdminProviderKeyBatchUpdateRequest>(body) {
|
||||
Ok(value) => value,
|
||||
Err(_) => {
|
||||
return Ok(build_admin_pool_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"请求体必须包含 key_ids 与 patch",
|
||||
));
|
||||
}
|
||||
},
|
||||
None => {
|
||||
return Ok(build_admin_pool_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"请求体必须包含 key_ids 与 patch",
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
state
|
||||
.build_admin_pool_batch_update_response(&provider_id, payload)
|
||||
.await
|
||||
}
|
||||
@@ -15,6 +15,8 @@ mod batch_import;
|
||||
mod batch_shared;
|
||||
#[path = "batch_routes/task_status.rs"]
|
||||
mod batch_task_status;
|
||||
#[path = "batch_routes/update.rs"]
|
||||
mod batch_update;
|
||||
pub(crate) mod payloads;
|
||||
#[path = "read_routes/keys.rs"]
|
||||
mod read_keys;
|
||||
@@ -152,6 +154,16 @@ pub(crate) async fn maybe_build_local_admin_pool_response(
|
||||
.await?,
|
||||
));
|
||||
}
|
||||
Some("batch_update_keys") => {
|
||||
return Ok(Some(
|
||||
batch_update::build_admin_pool_batch_update_response(
|
||||
state,
|
||||
request_context,
|
||||
request_body,
|
||||
)
|
||||
.await?,
|
||||
));
|
||||
}
|
||||
Some("batch_delete_task_status") => {
|
||||
return Ok(Some(
|
||||
batch_task_status::build_admin_pool_batch_delete_task_status_response(
|
||||
|
||||
@@ -213,6 +213,10 @@ pub(crate) fn is_admin_pool_route(request_context: &AdminRequestContext<'_>) ->
|
||||
&& path.starts_with("/api/admin/pool/")
|
||||
&& path.ends_with("/keys/batch-action")
|
||||
&& path.matches('/').count() == 6)
|
||||
|| (request_context.method() == http::Method::PATCH
|
||||
&& path.starts_with("/api/admin/pool/")
|
||||
&& path.ends_with("/keys/batch-update")
|
||||
&& path.matches('/').count() == 6)
|
||||
|| (request_context.method() == http::Method::POST
|
||||
&& path.starts_with("/api/admin/pool/")
|
||||
&& path.ends_with("/keys/resolve-selection")
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use super::payload::{
|
||||
provider_query_extract_api_key_id, provider_query_extract_force_refresh,
|
||||
provider_query_extract_model, provider_query_extract_provider_id,
|
||||
provider_query_extract_request_id,
|
||||
provider_query_extract_api_key_id, provider_query_extract_api_key_ids,
|
||||
provider_query_extract_force_refresh, provider_query_extract_model,
|
||||
provider_query_extract_provider_id, provider_query_extract_request_id,
|
||||
};
|
||||
use super::response::{
|
||||
build_admin_provider_query_bad_request_response, build_admin_provider_query_not_found_response,
|
||||
@@ -104,6 +104,26 @@ struct ProviderQueryKeyFetchResult {
|
||||
has_success: bool,
|
||||
}
|
||||
|
||||
fn provider_query_select_model_keys(
|
||||
keys: Vec<StoredProviderCatalogKey>,
|
||||
selected_key_ids: Option<&BTreeSet<String>>,
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, ()> {
|
||||
if selected_key_ids.is_some_and(|selected| {
|
||||
selected
|
||||
.iter()
|
||||
.any(|key_id| !keys.iter().any(|key| key.id == *key_id))
|
||||
}) {
|
||||
return Err(());
|
||||
}
|
||||
Ok(match selected_key_ids {
|
||||
Some(selected) => keys
|
||||
.into_iter()
|
||||
.filter(|key| selected.contains(&key.id))
|
||||
.collect(),
|
||||
None => keys.into_iter().filter(|key| key.is_active).collect(),
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_query_model_id(model: &Value) -> Option<&str> {
|
||||
model
|
||||
.get("id")
|
||||
@@ -622,21 +642,27 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
.into_response());
|
||||
}
|
||||
|
||||
let active_keys = keys
|
||||
.into_iter()
|
||||
.filter(|key| key.is_active)
|
||||
.collect::<Vec<_>>();
|
||||
if active_keys.is_empty() {
|
||||
let selected_key_ids = provider_query_extract_api_key_ids(payload);
|
||||
let query_keys = match provider_query_select_model_keys(keys, selected_key_ids.as_ref()) {
|
||||
Ok(keys) => keys,
|
||||
Err(()) => {
|
||||
return Ok(build_admin_provider_query_not_found_response(
|
||||
ADMIN_PROVIDER_QUERY_API_KEY_NOT_FOUND_DETAIL,
|
||||
));
|
||||
}
|
||||
};
|
||||
if query_keys.is_empty() {
|
||||
return Ok(build_admin_provider_query_bad_request_response(
|
||||
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
|
||||
));
|
||||
}
|
||||
let active_key_count = active_keys.len();
|
||||
let query_key_count = query_keys.len();
|
||||
|
||||
if provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("antigravity")
|
||||
if selected_key_ids.is_none()
|
||||
&& provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("antigravity")
|
||||
&& !force_refresh
|
||||
{
|
||||
if let Some(models) = provider_query_read_provider_cached_models(state, &provider.id).await
|
||||
@@ -649,8 +675,8 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
"error": serde_json::Value::Null,
|
||||
"warning": serde_json::Value::Null,
|
||||
"from_cache": true,
|
||||
"keys_total": active_key_count,
|
||||
"keys_cached": active_key_count,
|
||||
"keys_total": query_key_count,
|
||||
"keys_cached": query_key_count,
|
||||
"keys_fetched": 0,
|
||||
},
|
||||
"provider": provider_query_provider_payload(&provider),
|
||||
@@ -664,9 +690,9 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("antigravity")
|
||||
{
|
||||
provider_query_sort_antigravity_keys(state, &provider, &endpoints, active_keys).await?
|
||||
provider_query_sort_antigravity_keys(state, &provider, &endpoints, query_keys).await?
|
||||
} else {
|
||||
active_keys
|
||||
query_keys
|
||||
};
|
||||
|
||||
let mut all_models = Vec::new();
|
||||
@@ -709,10 +735,11 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
}
|
||||
|
||||
let models = aggregate_models_for_cache(&all_models);
|
||||
if provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("antigravity")
|
||||
if selected_key_ids.is_none()
|
||||
&& provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("antigravity")
|
||||
&& !models.is_empty()
|
||||
{
|
||||
provider_query_write_provider_cached_models(state, &provider.id, &models).await;
|
||||
@@ -742,7 +769,7 @@ pub(crate) async fn build_admin_provider_query_models_response(
|
||||
"error": error,
|
||||
"warning": warning,
|
||||
"from_cache": fetch_count == 0 && cache_hit_count > 0,
|
||||
"keys_total": active_key_count,
|
||||
"keys_total": query_key_count,
|
||||
"keys_cached": cache_hit_count,
|
||||
"keys_fetched": fetch_count,
|
||||
},
|
||||
@@ -770,6 +797,56 @@ mod tests {
|
||||
provider
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn selected_model_keys_use_the_explicit_batch_scope() {
|
||||
let mut first = StoredProviderCatalogKey::new(
|
||||
"key-a".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"A".to_string(),
|
||||
"api_key".to_string(),
|
||||
None,
|
||||
false,
|
||||
)
|
||||
.expect("key should build");
|
||||
first.is_active = false;
|
||||
let second = StoredProviderCatalogKey::new(
|
||||
"key-b".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"B".to_string(),
|
||||
"api_key".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build");
|
||||
|
||||
let selected = BTreeSet::from(["key-a".to_string()]);
|
||||
let keys =
|
||||
provider_query_select_model_keys(vec![first.clone(), second.clone()], Some(&selected))
|
||||
.expect("explicit selection should resolve");
|
||||
assert_eq!(keys.len(), 1);
|
||||
assert_eq!(keys[0].id, "key-a");
|
||||
|
||||
let active = provider_query_select_model_keys(vec![first, second], None)
|
||||
.expect("default selection should resolve");
|
||||
assert_eq!(active.len(), 1);
|
||||
assert_eq!(active[0].id, "key-b");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn selected_model_keys_reject_unknown_ids() {
|
||||
let key = StoredProviderCatalogKey::new(
|
||||
"key-a".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"A".to_string(),
|
||||
"api_key".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build");
|
||||
let selected = BTreeSet::from(["key-missing".to_string()]);
|
||||
assert!(provider_query_select_model_keys(vec![key], Some(&selected)).is_err());
|
||||
}
|
||||
|
||||
fn grok_key_with_quota(quota: Value) -> StoredProviderCatalogKey {
|
||||
let mut key = StoredProviderCatalogKey::new(
|
||||
"key-1".to_string(),
|
||||
|
||||
@@ -102,6 +102,12 @@ pub(crate) struct AdminProviderKeyUpdateRequest {
|
||||
|
||||
pub(crate) type AdminProviderKeyUpdatePatch = AdminTypedObjectPatch<AdminProviderKeyUpdateRequest>;
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub(crate) struct AdminProviderKeyBatchUpdateRequest {
|
||||
pub(crate) key_ids: Vec<String>,
|
||||
pub(crate) patch: serde_json::Value,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub(crate) struct AdminProviderKeyBatchDeleteRequest {
|
||||
pub(crate) ids: Vec<String>,
|
||||
|
||||
@@ -0,0 +1,102 @@
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdatePatch;
|
||||
use serde_json::{Map, Value};
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
const BATCH_EDITABLE_KEY_FIELDS: &[&str] = &[
|
||||
"allow_auth_channel_mismatch_formats",
|
||||
"allowed_models",
|
||||
"api_formats",
|
||||
"auth_type_by_format",
|
||||
"auto_fetch_models",
|
||||
"cache_ttl_minutes",
|
||||
"capabilities",
|
||||
"concurrent_limit",
|
||||
"global_priority_by_format",
|
||||
"internal_priority",
|
||||
"is_active",
|
||||
"locked_models",
|
||||
"max_probe_interval_minutes",
|
||||
"model_exclude_patterns",
|
||||
"model_include_patterns",
|
||||
"note",
|
||||
"proxy",
|
||||
"rate_multipliers",
|
||||
"rpm_limit",
|
||||
];
|
||||
|
||||
pub(crate) fn parse_admin_provider_key_batch_update_patch(
|
||||
value: Value,
|
||||
) -> Result<Map<String, Value>, String> {
|
||||
let Value::Object(patch) = value else {
|
||||
return Err("patch 必须是 JSON 对象".to_string());
|
||||
};
|
||||
if patch.is_empty() {
|
||||
return Err("patch 至少包含一个可编辑字段".to_string());
|
||||
}
|
||||
|
||||
let allowed = BATCH_EDITABLE_KEY_FIELDS
|
||||
.iter()
|
||||
.copied()
|
||||
.collect::<BTreeSet<_>>();
|
||||
let unsupported = patch
|
||||
.keys()
|
||||
.filter(|field| !allowed.contains(field.as_str()))
|
||||
.cloned()
|
||||
.collect::<BTreeSet<_>>();
|
||||
if !unsupported.is_empty() {
|
||||
return Err(format!(
|
||||
"批量编辑不支持字段: {}",
|
||||
unsupported.into_iter().collect::<Vec<_>>().join(", ")
|
||||
));
|
||||
}
|
||||
|
||||
AdminProviderKeyUpdatePatch::from_object(patch.clone())
|
||||
.map_err(|_| "patch 字段类型无效".to_string())?;
|
||||
Ok(patch)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::parse_admin_provider_key_batch_update_patch;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn accepts_shared_key_configuration_fields() {
|
||||
let patch = parse_admin_provider_key_batch_update_patch(json!({
|
||||
"api_formats": ["openai:responses"],
|
||||
"auto_fetch_models": true,
|
||||
"model_include_patterns": ["gpt-*"],
|
||||
"allowed_models": ["gpt-5.6-sol"],
|
||||
"rpm_limit": null
|
||||
}))
|
||||
.expect("batch patch should parse");
|
||||
|
||||
assert_eq!(patch.len(), 5);
|
||||
assert_eq!(patch["auto_fetch_models"], json!(true));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_identity_and_secret_fields() {
|
||||
let error = parse_admin_provider_key_batch_update_patch(json!({
|
||||
"name": "shared-name",
|
||||
"api_key": "sk-shared"
|
||||
}))
|
||||
.expect_err("identity fields must stay single-key only");
|
||||
|
||||
assert_eq!(error, "批量编辑不支持字段: api_key, name");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_empty_or_non_object_patch() {
|
||||
assert_eq!(
|
||||
parse_admin_provider_key_batch_update_patch(json!({}))
|
||||
.expect_err("empty patch should fail"),
|
||||
"patch 至少包含一个可编辑字段"
|
||||
);
|
||||
assert_eq!(
|
||||
parse_admin_provider_key_batch_update_patch(json!([]))
|
||||
.expect_err("array patch should fail"),
|
||||
"patch 必须是 JSON 对象"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -1,8 +1,14 @@
|
||||
pub(crate) use self::batch::parse_admin_provider_key_batch_update_patch;
|
||||
pub(crate) use self::create::build_admin_create_provider_key_record;
|
||||
pub(crate) use self::payload::build_admin_provider_keys_page_payload;
|
||||
pub(crate) use self::payload::build_admin_provider_keys_payload;
|
||||
pub(crate) use self::update::build_admin_update_provider_key_record;
|
||||
pub(crate) use self::update::{
|
||||
admin_provider_key_update_requires_immediate_model_fetch,
|
||||
build_admin_update_provider_key_record,
|
||||
build_admin_update_provider_key_record_with_existing_keys,
|
||||
};
|
||||
|
||||
mod batch;
|
||||
mod create;
|
||||
mod payload;
|
||||
mod update;
|
||||
|
||||
@@ -23,6 +23,27 @@ pub(crate) async fn build_admin_update_provider_key_record(
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
existing: &StoredProviderCatalogKey,
|
||||
patch: AdminProviderKeyUpdatePatch,
|
||||
) -> Result<StoredProviderCatalogKey, String> {
|
||||
let existing_keys = state
|
||||
.as_ref()
|
||||
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?;
|
||||
build_admin_update_provider_key_record_with_existing_keys(
|
||||
state,
|
||||
provider,
|
||||
existing,
|
||||
&existing_keys,
|
||||
patch,
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) fn build_admin_update_provider_key_record_with_existing_keys(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
existing: &StoredProviderCatalogKey,
|
||||
existing_keys: &[StoredProviderCatalogKey],
|
||||
patch: AdminProviderKeyUpdatePatch,
|
||||
) -> Result<StoredProviderCatalogKey, String> {
|
||||
let state = state.as_ref();
|
||||
let mut updated = existing.clone();
|
||||
@@ -57,11 +78,6 @@ pub(crate) async fn build_admin_update_provider_key_record(
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.cloned();
|
||||
|
||||
let existing_keys = state
|
||||
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?;
|
||||
|
||||
match target_auth_type.as_str() {
|
||||
"api_key" | "bearer" => {
|
||||
if let Some(api_key) = api_key_value
|
||||
@@ -308,7 +324,7 @@ pub(crate) async fn build_admin_update_provider_key_record(
|
||||
if let Some(auto_fetch_models) = payload.auto_fetch_models {
|
||||
updated.auto_fetch_models = auto_fetch_models;
|
||||
}
|
||||
if auto_fetch_disabled {
|
||||
if auto_fetch_disabled && !fields.contains("allowed_models") {
|
||||
updated.allowed_models = None;
|
||||
}
|
||||
if fields.contains("locked_models") {
|
||||
@@ -349,6 +365,17 @@ pub(crate) async fn build_admin_update_provider_key_record(
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) fn admin_provider_key_update_requires_immediate_model_fetch(
|
||||
existing: &StoredProviderCatalogKey,
|
||||
updated: &StoredProviderCatalogKey,
|
||||
) -> bool {
|
||||
let filters_changed = existing.model_include_patterns != updated.model_include_patterns
|
||||
|| existing.model_exclude_patterns != updated.model_exclude_patterns;
|
||||
let locked_models_changed = existing.locked_models != updated.locked_models;
|
||||
updated.auto_fetch_models
|
||||
&& (!existing.auto_fetch_models || filters_changed || locked_models_changed)
|
||||
}
|
||||
|
||||
fn raw_secret_auth_type(value: &str) -> bool {
|
||||
matches!(
|
||||
value.trim().to_ascii_lowercase().as_str(),
|
||||
|
||||
@@ -177,6 +177,16 @@ impl<'a> AdminAppState<'a> {
|
||||
self.app.update_provider_catalog_key(key).await
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_keys(
|
||||
&self,
|
||||
keys: &[aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey],
|
||||
) -> Result<
|
||||
Option<Vec<aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey>>,
|
||||
GatewayError,
|
||||
> {
|
||||
self.app.update_provider_catalog_keys(keys).await
|
||||
}
|
||||
|
||||
pub(crate) async fn create_provider_catalog_key(
|
||||
&self,
|
||||
key: &aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey,
|
||||
|
||||
@@ -13,6 +13,8 @@ use axum::{
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
impl<'a> AdminAppState<'a> {
|
||||
pub(crate) async fn clear_admin_provider_pool_cooldown(&self, provider_id: &str, key_id: &str) {
|
||||
@@ -466,7 +468,7 @@ impl<'a> AdminAppState<'a> {
|
||||
.into_response());
|
||||
}
|
||||
|
||||
let mut affected = 0usize;
|
||||
let mut updated_keys = Vec::with_capacity(keys.len());
|
||||
for mut key in keys {
|
||||
match plan.action {
|
||||
AdminPoolBatchActionKind::Enable => key.is_active = true,
|
||||
@@ -479,10 +481,13 @@ impl<'a> AdminAppState<'a> {
|
||||
}
|
||||
AdminPoolBatchActionKind::Delete => unreachable!(),
|
||||
}
|
||||
if self.update_provider_catalog_key(&key).await?.is_some() {
|
||||
affected = affected.saturating_add(1);
|
||||
}
|
||||
updated_keys.push(key);
|
||||
}
|
||||
let affected = self
|
||||
.update_provider_catalog_keys(&updated_keys)
|
||||
.await?
|
||||
.map(|keys| keys.len())
|
||||
.unwrap_or(0);
|
||||
|
||||
Ok(Json(
|
||||
admin_provider_pool_pure::build_admin_pool_batch_action_result_payload(
|
||||
@@ -492,4 +497,196 @@ impl<'a> AdminAppState<'a> {
|
||||
)
|
||||
.into_response())
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_pool_batch_update_response(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
payload: crate::handlers::admin::provider::shared::payloads::AdminProviderKeyBatchUpdateRequest,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
use crate::handlers::admin::provider::pool_admin::admin_provider_pool_config;
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdatePatch;
|
||||
use crate::handlers::admin::provider::write::keys::{
|
||||
admin_provider_key_update_requires_immediate_model_fetch,
|
||||
build_admin_update_provider_key_record_with_existing_keys,
|
||||
parse_admin_provider_key_batch_update_patch,
|
||||
};
|
||||
use crate::maintenance::ensure_provider_key_pool_scores_for_keys;
|
||||
use crate::model_fetch::perform_model_fetch_for_keys;
|
||||
|
||||
let Some(provider) = self
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id.to_string()))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok((
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
|
||||
)
|
||||
.into_response());
|
||||
};
|
||||
|
||||
let requested_key_ids = payload
|
||||
.key_ids
|
||||
.into_iter()
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
.collect::<BTreeSet<_>>();
|
||||
if requested_key_ids.is_empty() {
|
||||
return Ok((
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
Json(json!({ "detail": "key_ids 不能为空" })),
|
||||
)
|
||||
.into_response());
|
||||
}
|
||||
|
||||
let patch = match parse_admin_provider_key_batch_update_patch(payload.patch) {
|
||||
Ok(patch) => patch,
|
||||
Err(detail) => {
|
||||
return Ok((
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
Json(json!({ "detail": detail })),
|
||||
)
|
||||
.into_response());
|
||||
}
|
||||
};
|
||||
|
||||
let provider_ids = vec![provider.id.clone()];
|
||||
let existing_keys = self
|
||||
.list_provider_catalog_keys_by_provider_ids(&provider_ids)
|
||||
.await?;
|
||||
let keys_by_id = existing_keys
|
||||
.iter()
|
||||
.map(|key| (key.id.clone(), key))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
let missing_key_ids = requested_key_ids
|
||||
.iter()
|
||||
.filter(|key_id| !keys_by_id.contains_key(*key_id))
|
||||
.cloned()
|
||||
.collect::<Vec<_>>();
|
||||
if !missing_key_ids.is_empty() {
|
||||
return Ok((
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
Json(json!({
|
||||
"detail": format!(
|
||||
"以下密钥不存在或不属于当前 Provider: {}",
|
||||
missing_key_ids.join(", ")
|
||||
)
|
||||
})),
|
||||
)
|
||||
.into_response());
|
||||
}
|
||||
|
||||
let mut staged_updates = Vec::with_capacity(requested_key_ids.len());
|
||||
for key_id in &requested_key_ids {
|
||||
let existing = keys_by_id
|
||||
.get(key_id)
|
||||
.expect("validated provider key should exist");
|
||||
let typed_patch = AdminProviderKeyUpdatePatch::from_object(patch.clone())
|
||||
.expect("validated batch patch should remain parseable");
|
||||
let updated = match build_admin_update_provider_key_record_with_existing_keys(
|
||||
self,
|
||||
&provider,
|
||||
existing,
|
||||
&existing_keys,
|
||||
typed_patch,
|
||||
) {
|
||||
Ok(updated) => updated,
|
||||
Err(detail) => {
|
||||
return Ok((
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
Json(json!({
|
||||
"detail": format!("密钥 {} 配置无效: {detail}", existing.name)
|
||||
})),
|
||||
)
|
||||
.into_response());
|
||||
}
|
||||
};
|
||||
staged_updates.push(((*existing).clone(), updated));
|
||||
}
|
||||
|
||||
let model_fetch_key_ids = staged_updates
|
||||
.iter()
|
||||
.filter(|(existing, updated)| {
|
||||
admin_provider_key_update_requires_immediate_model_fetch(existing, updated)
|
||||
})
|
||||
.map(|(_, updated)| updated.id.clone())
|
||||
.collect::<BTreeSet<_>>();
|
||||
|
||||
let staged_records = staged_updates
|
||||
.into_iter()
|
||||
.map(|(_, updated)| updated)
|
||||
.collect::<Vec<_>>();
|
||||
let Some(updated_keys) = self.update_provider_catalog_keys(&staged_records).await? else {
|
||||
return Ok((
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
Json(json!({ "detail": "Provider 密钥写入能力不可用" })),
|
||||
)
|
||||
.into_response());
|
||||
};
|
||||
|
||||
let endpoints = self
|
||||
.list_provider_catalog_endpoints_by_provider_ids(&provider_ids)
|
||||
.await?;
|
||||
if let Some(pool_config) = admin_provider_pool_config(&provider) {
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
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(
|
||||
self.as_ref(),
|
||||
&provider,
|
||||
&pool_config,
|
||||
&endpoints,
|
||||
&updated_keys,
|
||||
now_unix_secs,
|
||||
score_ensure_budget,
|
||||
)
|
||||
.await
|
||||
{
|
||||
tracing::debug!(
|
||||
provider_id = %provider.id,
|
||||
updated_keys = updated_keys.len(),
|
||||
error = ?err,
|
||||
"gateway admin provider key batch update: failed to seed pool score rows"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let model_sync = if model_fetch_key_ids.is_empty() {
|
||||
serde_json::Value::Null
|
||||
} else {
|
||||
let requested = model_fetch_key_ids.len();
|
||||
match perform_model_fetch_for_keys(self.as_ref(), &provider.id, &model_fetch_key_ids)
|
||||
.await
|
||||
{
|
||||
Ok(summary) => json!({
|
||||
"requested": requested,
|
||||
"attempted": summary.attempted,
|
||||
"succeeded": summary.succeeded,
|
||||
"failed": summary.failed,
|
||||
"skipped": summary.skipped,
|
||||
}),
|
||||
Err(err) => json!({
|
||||
"requested": requested,
|
||||
"attempted": 0,
|
||||
"succeeded": 0,
|
||||
"failed": requested,
|
||||
"skipped": 0,
|
||||
"error": err.into_message(),
|
||||
}),
|
||||
}
|
||||
};
|
||||
|
||||
let affected = updated_keys.len();
|
||||
Ok(Json(json!({
|
||||
"affected": affected,
|
||||
"message": format!("已更新 {affected} 个密钥"),
|
||||
"model_sync": model_sync,
|
||||
}))
|
||||
.into_response())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,5 +5,6 @@ mod tests;
|
||||
pub(crate) use aether_model_fetch::ModelFetchRunSummary;
|
||||
pub(crate) use runtime::state::ModelFetchRuntimeState;
|
||||
pub(crate) use runtime::{
|
||||
perform_model_fetch_for_key, perform_model_fetch_once, spawn_model_fetch_worker,
|
||||
perform_model_fetch_for_key, perform_model_fetch_for_keys, perform_model_fetch_once,
|
||||
spawn_model_fetch_worker,
|
||||
};
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use std::collections::HashMap;
|
||||
use std::collections::{BTreeSet, HashMap};
|
||||
use std::time::Duration;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
@@ -70,7 +70,16 @@ pub(crate) async fn perform_model_fetch_for_key(
|
||||
provider_id: &str,
|
||||
key_id: &str,
|
||||
) -> Result<ModelFetchRunSummary, GatewayError> {
|
||||
perform_model_fetch_for_key_with_state(state, provider_id, key_id).await
|
||||
let key_ids = BTreeSet::from([key_id.to_string()]);
|
||||
perform_model_fetch_for_keys_with_state(state, provider_id, &key_ids).await
|
||||
}
|
||||
|
||||
pub(crate) async fn perform_model_fetch_for_keys(
|
||||
state: &AppState,
|
||||
provider_id: &str,
|
||||
key_ids: &BTreeSet<String>,
|
||||
) -> Result<ModelFetchRunSummary, GatewayError> {
|
||||
perform_model_fetch_for_keys_with_state(state, provider_id, key_ids).await
|
||||
}
|
||||
|
||||
async fn perform_model_fetch_once_with_state<S>(
|
||||
@@ -83,22 +92,22 @@ where
|
||||
execute_fetch_targets(state, targets).await
|
||||
}
|
||||
|
||||
async fn perform_model_fetch_for_key_with_state<S>(
|
||||
async fn perform_model_fetch_for_keys_with_state<S>(
|
||||
state: &S,
|
||||
provider_id: &str,
|
||||
key_id: &str,
|
||||
key_ids: &BTreeSet<String>,
|
||||
) -> Result<ModelFetchRunSummary, GatewayError>
|
||||
where
|
||||
S: ModelFetchRuntimeState + ?Sized,
|
||||
{
|
||||
let targets = collect_fetch_targets(state, Some(provider_id), Some(key_id)).await?;
|
||||
let targets = collect_fetch_targets(state, Some(provider_id), Some(key_ids)).await?;
|
||||
execute_fetch_targets(state, targets).await
|
||||
}
|
||||
|
||||
async fn collect_fetch_targets<S>(
|
||||
state: &S,
|
||||
provider_id_filter: Option<&str>,
|
||||
key_id_filter: Option<&str>,
|
||||
key_id_filter: Option<&BTreeSet<String>>,
|
||||
) -> Result<Vec<SelectedFetchTarget>, GatewayError>
|
||||
where
|
||||
S: ModelFetchRuntimeState + ?Sized,
|
||||
@@ -152,7 +161,7 @@ where
|
||||
.unwrap_or_default();
|
||||
let keys = keys_by_provider.remove(&provider.id).unwrap_or_default();
|
||||
for key in keys {
|
||||
if key_id_filter.is_some_and(|key_id| key.id != key_id) {
|
||||
if key_id_filter.is_some_and(|key_ids| !key_ids.contains(&key.id)) {
|
||||
continue;
|
||||
}
|
||||
if !key.is_active || !key.auto_fetch_models {
|
||||
|
||||
@@ -662,6 +662,21 @@ impl AppState {
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_keys(
|
||||
&self,
|
||||
keys: &[provider_catalog::StoredProviderCatalogKey],
|
||||
) -> Result<Option<Vec<provider_catalog::StoredProviderCatalogKey>>, GatewayError> {
|
||||
let updated = self
|
||||
.data
|
||||
.update_provider_catalog_keys(keys)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
if updated.as_ref().is_some_and(|keys| !keys.is_empty()) {
|
||||
self.invalidate_provider_routing_caches();
|
||||
}
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn update_provider_catalog_key_runtime_state(
|
||||
&self,
|
||||
key: &provider_catalog::StoredProviderCatalogKey,
|
||||
|
||||
@@ -886,6 +886,7 @@ fn admin_provider_pool_admin_mod_stays_thin() {
|
||||
"#[path = \"batch_routes/import.rs\"]",
|
||||
"#[path = \"batch_routes/shared.rs\"]",
|
||||
"#[path = \"batch_routes/task_status.rs\"]",
|
||||
"#[path = \"batch_routes/update.rs\"]",
|
||||
"#[path = \"read_routes/keys.rs\"]",
|
||||
"#[path = \"read_routes/overview.rs\"]",
|
||||
"#[path = \"read_routes/presets.rs\"]",
|
||||
|
||||
@@ -3405,6 +3405,127 @@ async fn gateway_handles_admin_pool_batch_action_locally_with_trusted_admin_prin
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_batch_updates_shared_pool_key_configuration() {
|
||||
let provider = sample_provider("provider-openai", "openai", 10).with_transport_fields(
|
||||
true,
|
||||
false,
|
||||
true,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
Some(json!({
|
||||
"pool_advanced": {
|
||||
"enabled": true
|
||||
}
|
||||
})),
|
||||
);
|
||||
let mut first_key = sample_key("key-openai-a", "provider-openai", "openai:chat", "sk-a");
|
||||
first_key.name = "alpha".to_string();
|
||||
first_key.auto_fetch_models = true;
|
||||
first_key.allowed_models = Some(json!(["legacy-model"]));
|
||||
let mut second_key = sample_key("key-openai-b", "provider-openai", "openai:chat", "sk-b");
|
||||
second_key.name = "beta".to_string();
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
Vec::new(),
|
||||
vec![first_key, second_key],
|
||||
));
|
||||
let state = AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_repository_for_tests(Arc::clone(
|
||||
&provider_catalog_repository,
|
||||
)),
|
||||
);
|
||||
|
||||
let response = local_admin_pool_response(
|
||||
&state,
|
||||
http::Method::PATCH,
|
||||
"/api/admin/pool/provider-openai/keys/batch-update",
|
||||
Some(json!({
|
||||
"key_ids": ["key-openai-b", "key-openai-a", "key-openai-a"],
|
||||
"patch": {
|
||||
"api_formats": ["openai:responses"],
|
||||
"internal_priority": 7,
|
||||
"rpm_limit": null,
|
||||
"auto_fetch_models": false,
|
||||
"allowed_models": ["gpt-5.6-sol", "gpt-5.6-luna"],
|
||||
"locked_models": [],
|
||||
"note": null
|
||||
}
|
||||
})),
|
||||
)
|
||||
.await;
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = serde_json::from_slice(
|
||||
&to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("body should read"),
|
||||
)
|
||||
.expect("json body should parse");
|
||||
assert_eq!(payload["affected"], json!(2));
|
||||
assert_eq!(payload["model_sync"], serde_json::Value::Null);
|
||||
|
||||
let stored = provider_catalog_repository
|
||||
.list_keys_by_ids(&["key-openai-a".to_string(), "key-openai-b".to_string()])
|
||||
.await
|
||||
.expect("keys should load");
|
||||
assert_eq!(stored.len(), 2);
|
||||
for key in stored {
|
||||
assert_eq!(key.api_formats, Some(json!(["openai:responses"])));
|
||||
assert_eq!(key.internal_priority, 7);
|
||||
assert_eq!(key.rpm_limit, None);
|
||||
assert!(!key.auto_fetch_models);
|
||||
assert_eq!(
|
||||
key.allowed_models,
|
||||
Some(json!(["gpt-5.6-sol", "gpt-5.6-luna"]))
|
||||
);
|
||||
assert_eq!(key.locked_models, None);
|
||||
assert_eq!(key.note, None);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_rejects_pool_batch_update_before_writing_any_key() {
|
||||
let provider = sample_provider("provider-openai", "openai", 10);
|
||||
let mut first_key = sample_key("key-openai-a", "provider-openai", "openai:chat", "sk-a");
|
||||
first_key.internal_priority = 3;
|
||||
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![provider],
|
||||
Vec::new(),
|
||||
vec![first_key],
|
||||
));
|
||||
let state = AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_provider_catalog_repository_for_tests(Arc::clone(
|
||||
&provider_catalog_repository,
|
||||
)),
|
||||
);
|
||||
|
||||
let response = local_admin_pool_response(
|
||||
&state,
|
||||
http::Method::PATCH,
|
||||
"/api/admin/pool/provider-openai/keys/batch-update",
|
||||
Some(json!({
|
||||
"key_ids": ["key-openai-a", "key-missing"],
|
||||
"patch": { "internal_priority": 9 }
|
||||
})),
|
||||
)
|
||||
.await;
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
|
||||
let stored = provider_catalog_repository
|
||||
.list_keys_by_ids(&["key-openai-a".to_string()])
|
||||
.await
|
||||
.expect("key should load");
|
||||
assert_eq!(stored[0].internal_priority, 3);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_handles_admin_pool_batch_delete_locally_with_trusted_admin_principal() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
|
||||
Reference in New Issue
Block a user