refactor: 大规模模块拆分与代码精简,新增 ai-pipeline/data-contracts 独立 crate

- 新增 aether-ai-pipeline 和 aether-data-contracts crate,将 pipeline 逻辑与数据契约从 gateway 中解耦
- 重构 admin handlers:拆分单体模块为 auth/billing/endpoint/features/model/observability/provider/system 等独立子模块
- 合并 chat/cli 重复代码路径:精简 conversion、finalize、planner 中的 sync/chat/cli 分支
- 重构 scheduler/executor/data 层,引入 facade 模式降低模块间耦合
- 移除冗余的 intent 模块,将 plan_fallback/policy/stream_path/sync_path 迁移至 executor
- 前端适配:调整 admin API 调用和 provider 模型测试对话框
This commit is contained in:
fawney19
2026-04-07 02:50:19 +08:00
parent 763ff03a7b
commit 5d96d6673b
732 changed files with 28589 additions and 20662 deletions
@@ -0,0 +1,5 @@
mod routes;
mod shared;
pub(crate) use routes::maybe_build_local_admin_providers_response;
use shared::*;
@@ -0,0 +1,723 @@
use super::shared::build_admin_providers_data_unavailable_response;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::delete_task::{
build_admin_provider_mapping_preview_payload, run_admin_provider_delete_task,
};
use crate::handlers::admin::provider::pool::{
build_admin_provider_pool_status_payload, clear_admin_provider_pool_cooldown,
reset_admin_provider_pool_cost,
};
use crate::handlers::admin::provider::shared::{
admin_provider_clear_pool_cooldown_parts, admin_provider_delete_task_parts,
admin_provider_id_for_health_monitor, admin_provider_id_for_manage_path,
admin_provider_id_for_mapping_preview, admin_provider_id_for_pool_status,
admin_provider_id_for_summary, admin_provider_reset_pool_cost_parts,
build_admin_provider_delete_task_payload, is_admin_providers_root,
put_admin_provider_delete_task, AdminProviderCreateRequest, AdminProviderUpdateRequest,
};
use crate::handlers::admin::provider::summary::{
build_admin_provider_health_monitor_payload, build_admin_provider_summary_payload,
build_admin_provider_summary_value, build_admin_providers_payload,
build_admin_providers_summary_payload,
};
use crate::handlers::admin::provider::write::{
build_admin_create_provider_record, build_admin_fixed_provider_endpoint_record,
build_admin_update_provider_record,
};
use crate::handlers::admin::shared::{
attach_admin_audit_response, query_param_optional_bool, query_param_value,
};
use crate::provider_transport::provider_types::fixed_provider_template;
use crate::{AppState, GatewayError, LocalProviderDeleteTaskState};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use tracing::warn;
use uuid::Uuid;
fn attach_admin_provider_delete_task_terminal_audit(
provider_id: &str,
task_id: &str,
task_status: &str,
response: Response<Body>,
) -> Response<Body> {
match task_status {
"completed" => attach_admin_audit_response(
response,
"admin_provider_delete_task_completed_viewed",
"view_provider_delete_task_terminal_state",
"provider_delete_task",
&format!("{provider_id}:{task_id}"),
),
"failed" => attach_admin_audit_response(
response,
"admin_provider_delete_task_failed_viewed",
"view_provider_delete_task_terminal_state",
"provider_delete_task",
&format!("{provider_id}:{task_id}"),
),
_ => response,
}
}
pub(crate) async fn maybe_build_local_admin_providers_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.control_decision.as_ref() else {
return Ok(None);
};
if decision.route_family.as_deref() == Some("providers_manage")
&& decision.route_kind.as_deref() == Some("create_provider")
&& request_context.request_method == http::Method::POST
&& matches!(
request_context.request_path.as_str(),
"/api/admin/providers" | "/api/admin/providers/"
)
{
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
if !state.has_provider_catalog_data_reader() || !state.has_provider_catalog_data_writer() {
return Ok(Some(build_admin_providers_data_unavailable_response()));
}
let payload = match serde_json::from_slice::<AdminProviderCreateRequest>(request_body) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let (record, shift_existing_priorities_from) =
match build_admin_create_provider_record(state, payload).await {
Ok(record) => record,
Err(message) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": message })),
)
.into_response(),
));
}
};
let Some(created_provider) = state
.create_provider_catalog_provider(&record, shift_existing_priorities_from)
.await?
else {
return Ok(Some(build_admin_providers_data_unavailable_response()));
};
if let Some((base_url, endpoint_signatures)) =
fixed_provider_template(&created_provider.provider_type)
{
for endpoint_signature in endpoint_signatures {
let endpoint = match build_admin_fixed_provider_endpoint_record(
&created_provider,
endpoint_signature,
base_url,
) {
Ok(endpoint) => endpoint,
Err(message) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": message })),
)
.into_response(),
));
}
};
let Some(_) = state.create_provider_catalog_endpoint(&endpoint).await? else {
return Ok(Some(build_admin_providers_data_unavailable_response()));
};
}
}
return Ok(Some(attach_admin_audit_response(
Json(json!({
"id": created_provider.id,
"name": created_provider.name,
"message": "提供商创建成功",
}))
.into_response(),
"admin_provider_created",
"create_provider",
"provider",
&created_provider.id,
)));
}
if decision.route_family.as_deref() == Some("providers_manage")
&& decision.route_kind.as_deref() == Some("list_providers")
&& is_admin_providers_root(&request_context.request_path)
{
if !state.has_provider_catalog_data_reader() {
return Ok(Some(build_admin_providers_data_unavailable_response()));
}
let skip = query_param_value(request_context.request_query_string.as_deref(), "skip")
.and_then(|value| value.parse::<usize>().ok())
.unwrap_or(0);
let limit = query_param_value(request_context.request_query_string.as_deref(), "limit")
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| *value > 0 && *value <= 500)
.unwrap_or(100);
let is_active =
query_param_optional_bool(request_context.request_query_string.as_deref(), "is_active");
let Some(payload) = build_admin_providers_payload(state, skip, limit, is_active).await
else {
return Ok(Some(build_admin_providers_data_unavailable_response()));
};
return Ok(Some(Json(payload).into_response()));
}
if decision.route_family.as_deref() == Some("providers_manage")
&& decision.route_kind.as_deref() == Some("update_provider")
&& request_context.request_method == http::Method::PATCH
&& request_context
.request_path
.starts_with("/api/admin/providers/")
{
let Some(provider_id) = admin_provider_id_for_manage_path(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
if !state.has_provider_catalog_data_reader() || !state.has_provider_catalog_data_writer() {
return Ok(Some(build_admin_providers_data_unavailable_response()));
}
let raw_value = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(value) => value,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let Some(raw_payload) = raw_value.as_object().cloned() else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
};
let payload = match serde_json::from_value::<AdminProviderUpdateRequest>(raw_value) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let Some(existing_provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
));
};
let updated_record = match build_admin_update_provider_record(
state,
&existing_provider,
&raw_payload,
payload,
)
.await
{
Ok(record) => record,
Err(detail) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
));
}
};
let Some(_updated) = state
.update_provider_catalog_provider(&updated_record)
.await?
else {
return Ok(Some(build_admin_providers_data_unavailable_response()));
};
return Ok(Some(
match build_admin_provider_summary_payload(state, &provider_id).await {
Some(payload) => attach_admin_audit_response(
Json(payload).into_response(),
"admin_provider_updated",
"update_provider",
"provider",
&provider_id,
),
None => (
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
},
));
}
if decision.route_family.as_deref() == Some("providers_manage")
&& decision.route_kind.as_deref() == Some("delete_provider")
&& request_context.request_method == http::Method::DELETE
{
let Some(provider_id) = admin_provider_id_for_manage_path(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "提供商不存在" })),
)
.into_response(),
));
};
let task_id = Uuid::new_v4().simple().to_string()[..16].to_string();
let pending_task = LocalProviderDeleteTaskState {
task_id: task_id.clone(),
provider_id: provider.id.clone(),
status: "pending".to_string(),
stage: "queued".to_string(),
total_keys: 0,
deleted_keys: 0,
total_endpoints: 0,
deleted_endpoints: 0,
message: "delete task submitted".to_string(),
};
put_admin_provider_delete_task(state, &pending_task);
if let Err(err) = run_admin_provider_delete_task(state, &provider.id, &task_id).await {
warn!(
"gateway admin provider delete task failed for provider {}: {:?}",
provider.id, err
);
put_admin_provider_delete_task(
state,
&LocalProviderDeleteTaskState {
task_id: task_id.clone(),
provider_id: provider.id.clone(),
status: "failed".to_string(),
stage: "failed".to_string(),
total_keys: 0,
deleted_keys: 0,
total_endpoints: 0,
deleted_endpoints: 0,
message: format!("provider delete failed: {err:?}"),
},
);
}
return Ok(Some(attach_admin_audit_response(
Json(json!({
"task_id": task_id,
"status": "pending",
"message": "删除任务已提交,提供商已进入后台删除队列",
}))
.into_response(),
"admin_provider_delete_queued",
"delete_provider",
"provider",
&provider.id,
)));
}
if decision.route_family.as_deref() == Some("providers_manage")
&& decision.route_kind.as_deref() == Some("summary_list")
&& request_context.request_path == "/api/admin/providers/summary"
{
if !state.has_provider_catalog_data_reader() {
return Ok(Some(build_admin_providers_data_unavailable_response()));
}
let page = query_param_value(request_context.request_query_string.as_deref(), "page")
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| *value > 0)
.unwrap_or(1);
let page_size =
query_param_value(request_context.request_query_string.as_deref(), "page_size")
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| *value > 0 && *value <= 10_000)
.unwrap_or(20);
let search = query_param_value(request_context.request_query_string.as_deref(), "search")
.unwrap_or_default();
let status = query_param_value(request_context.request_query_string.as_deref(), "status")
.unwrap_or_else(|| "all".to_string());
let api_format = query_param_value(
request_context.request_query_string.as_deref(),
"api_format",
)
.unwrap_or_else(|| "all".to_string());
let model_id =
query_param_value(request_context.request_query_string.as_deref(), "model_id")
.unwrap_or_else(|| "all".to_string());
let Some(payload) = build_admin_providers_summary_payload(
state,
page,
page_size,
&search,
&status,
&api_format,
&model_id,
)
.await
else {
return Ok(Some(build_admin_providers_data_unavailable_response()));
};
return Ok(Some(Json(payload).into_response()));
}
if decision.route_family.as_deref() == Some("providers_manage")
&& decision.route_kind.as_deref() == Some("provider_summary")
&& request_context
.request_path
.starts_with("/api/admin/providers/")
&& request_context.request_path.ends_with("/summary")
{
if !state.has_provider_catalog_data_reader() {
return Ok(Some(build_admin_providers_data_unavailable_response()));
}
let Some(provider_id) = admin_provider_id_for_summary(&request_context.request_path) else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
return Ok(Some(
match build_admin_provider_summary_payload(state, &provider_id).await {
Some(payload) => Json(payload).into_response(),
None => (
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
},
));
}
if decision.route_family.as_deref() == Some("providers_manage")
&& decision.route_kind.as_deref() == Some("delete_provider_task")
&& request_context.request_method == http::Method::GET
{
let Some((provider_id, task_id)) =
admin_provider_delete_task_parts(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Task not found" })),
)
.into_response(),
));
};
let Some(task) = state.get_provider_delete_task(&task_id) else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Task not found" })),
)
.into_response(),
));
};
if task.provider_id != provider_id {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Task not found" })),
)
.into_response(),
));
}
return Ok(Some(attach_admin_provider_delete_task_terminal_audit(
&provider_id,
&task_id,
task.status.as_str(),
Json(build_admin_provider_delete_task_payload(&task)).into_response(),
)));
}
if decision.route_family.as_deref() == Some("providers_manage")
&& decision.route_kind.as_deref() == Some("health_monitor")
&& request_context
.request_path
.starts_with("/api/admin/providers/")
&& request_context.request_path.ends_with("/health-monitor")
{
let Some(provider_id) = admin_provider_id_for_health_monitor(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let lookback_hours = query_param_value(
request_context.request_query_string.as_deref(),
"lookback_hours",
)
.and_then(|value| value.parse::<u64>().ok())
.filter(|value| (1..=72).contains(value))
.unwrap_or(6);
let per_endpoint_limit = query_param_value(
request_context.request_query_string.as_deref(),
"per_endpoint_limit",
)
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| (10..=200).contains(value))
.unwrap_or(48);
return Ok(Some(
match build_admin_provider_health_monitor_payload(
state,
&provider_id,
lookback_hours,
per_endpoint_limit,
)
.await
{
Some(payload) => Json(payload).into_response(),
None => (
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
},
));
}
if decision.route_family.as_deref() == Some("providers_manage")
&& decision.route_kind.as_deref() == Some("mapping_preview")
&& request_context
.request_path
.starts_with("/api/admin/providers/")
&& request_context.request_path.ends_with("/mapping-preview")
{
let Some(provider_id) =
admin_provider_id_for_mapping_preview(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
return Ok(Some(
match build_admin_provider_mapping_preview_payload(state, &provider_id).await {
Some(payload) => Json(payload).into_response(),
None => (
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
},
));
}
if decision.route_family.as_deref() == Some("providers_manage")
&& decision.route_kind.as_deref() == Some("pool_status")
&& request_context.request_method == http::Method::GET
&& request_context.request_path.ends_with("/pool-status")
{
let Some(provider_id) = admin_provider_id_for_pool_status(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
return Ok(Some(
match build_admin_provider_pool_status_payload(state, &provider_id).await {
Some(payload) => Json(payload).into_response(),
None => (
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
},
));
}
if decision.route_family.as_deref() == Some("providers_manage")
&& decision.route_kind.as_deref() == Some("clear_pool_cooldown")
&& request_context.request_method == http::Method::POST
{
let Some((provider_id, key_id)) =
admin_provider_clear_pool_cooldown_parts(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Key 不存在" })),
)
.into_response(),
));
};
let provider_exists = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
.is_some();
if !provider_exists {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
));
}
let key = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.find(|key| key.id == key_id);
let Some(key) = key else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Key {key_id} 不存在") })),
)
.into_response(),
));
};
clear_admin_provider_pool_cooldown(state, &provider_id, &key_id).await;
return Ok(Some(attach_admin_audit_response(
Json(json!({
"message": format!("已清除 Key {} 的冷却状态", key.name),
}))
.into_response(),
"admin_provider_pool_cooldown_cleared",
"clear_provider_pool_cooldown",
"provider_key",
&key_id,
)));
}
if decision.route_family.as_deref() == Some("providers_manage")
&& decision.route_kind.as_deref() == Some("reset_pool_cost")
&& request_context.request_method == http::Method::POST
{
let Some((provider_id, key_id)) =
admin_provider_reset_pool_cost_parts(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Key 不存在" })),
)
.into_response(),
));
};
let provider_exists = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
.is_some();
if !provider_exists {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
));
}
let key = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.find(|key| key.id == key_id);
let Some(key) = key else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Key {key_id} 不存在") })),
)
.into_response(),
));
};
reset_admin_provider_pool_cost(state, &provider_id, &key_id).await;
return Ok(Some(attach_admin_audit_response(
Json(json!({
"message": format!("已重置 Key {} 的成本窗口", key.name),
}))
.into_response(),
"admin_provider_pool_cost_reset",
"reset_provider_pool_cost",
"provider_key",
&key_id,
)));
}
Ok(None)
}
@@ -0,0 +1,18 @@
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) const ADMIN_PROVIDERS_DATA_UNAVAILABLE_DETAIL: &str =
"Admin provider catalog data unavailable";
pub(super) fn build_admin_providers_data_unavailable_response() -> Response<Body> {
(
http::StatusCode::SERVICE_UNAVAILABLE,
Json(json!({ "detail": ADMIN_PROVIDERS_DATA_UNAVAILABLE_DETAIL })),
)
.into_response()
}
@@ -0,0 +1,317 @@
use crate::handlers::admin::provider::shared::{
put_admin_provider_delete_task, ADMIN_PROVIDER_MAPPING_PREVIEW_FETCH_LIMIT,
ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_KEYS, ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_MODELS,
};
use crate::handlers::admin::shared::{decrypt_catalog_secret_with_fallbacks, json_string_list};
use crate::handlers::public::matches_model_mapping_for_models;
use crate::{AppState, GatewayError, LocalProviderDeleteTaskState};
use aether_data_contracts::repository::global_models::{
AdminProviderModelListQuery, PublicGlobalModelQuery, StoredPublicGlobalModel,
};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(crate) async fn run_admin_provider_delete_task(
state: &AppState,
provider_id: &str,
task_id: &str,
) -> Result<LocalProviderDeleteTaskState, GatewayError> {
let Some(mut provider) = state
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
.await?
.into_iter()
.next()
else {
return Err(GatewayError::Internal(format!(
"provider {provider_id} not found for delete task"
)));
};
let keys = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let models = state
.list_admin_provider_models(&AdminProviderModelListQuery {
provider_id: provider.id.clone(),
offset: 0,
limit: 10_000,
is_active: None,
})
.await
.unwrap_or_default();
let mut task = LocalProviderDeleteTaskState {
task_id: task_id.to_string(),
provider_id: provider.id.clone(),
status: "running".to_string(),
stage: "preparing".to_string(),
total_keys: keys.len(),
deleted_keys: 0,
total_endpoints: endpoints.len(),
deleted_endpoints: 0,
message: format!(
"preparing delete for {} keys and {} endpoints",
keys.len(),
endpoints.len()
),
};
put_admin_provider_delete_task(state, &task);
if provider.is_active {
provider.is_active = false;
provider.updated_at_unix_secs = Some(
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0),
);
let _ = state.update_provider_catalog_provider(&provider).await?;
}
task.stage = "disabling".to_string();
task.message = "provider disabled; starting cleanup".to_string();
put_admin_provider_delete_task(state, &task);
let endpoint_ids = endpoints
.iter()
.map(|item| item.id.clone())
.collect::<Vec<_>>();
let key_ids = keys.iter().map(|item| item.id.clone()).collect::<Vec<_>>();
state
.cleanup_deleted_provider_catalog_refs(&provider.id, &endpoint_ids, &key_ids)
.await?;
task.stage = "deleting_models".to_string();
task.message = format!("deleting {} provider models", models.len());
put_admin_provider_delete_task(state, &task);
for model in &models {
let _ = state
.delete_admin_provider_model(&provider.id, &model.id)
.await?;
}
task.stage = "deleting_keys".to_string();
task.message = "deleting provider keys".to_string();
put_admin_provider_delete_task(state, &task);
for key in &keys {
if state.delete_provider_catalog_key(&key.id).await? {
task.deleted_keys += 1;
task.message = format!("deleted {} / {} keys", task.deleted_keys, task.total_keys);
put_admin_provider_delete_task(state, &task);
}
}
task.stage = "deleting_endpoints".to_string();
task.message = "deleting provider endpoints".to_string();
put_admin_provider_delete_task(state, &task);
for endpoint in &endpoints {
if state.delete_provider_catalog_endpoint(&endpoint.id).await? {
task.deleted_endpoints += 1;
task.message = format!(
"deleted {} / {} endpoints",
task.deleted_endpoints, task.total_endpoints
);
put_admin_provider_delete_task(state, &task);
}
}
task.stage = "deleting_provider".to_string();
task.message = "deleting provider record".to_string();
put_admin_provider_delete_task(state, &task);
if !state.delete_provider_catalog_provider(&provider.id).await? {
task.status = "failed".to_string();
task.stage = "failed".to_string();
task.message = "provider delete failed".to_string();
put_admin_provider_delete_task(state, &task);
return Ok(task);
}
task.status = "completed".to_string();
task.stage = "completed".to_string();
task.message = format!(
"provider deleted: keys={}, endpoints={}",
task.deleted_keys, task.deleted_endpoints
);
put_admin_provider_delete_task(state, &task);
Ok(task)
}
pub(crate) fn public_global_model_mapping_patterns(model: &StoredPublicGlobalModel) -> Vec<String> {
model
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|config| config.get("model_mappings"))
.and_then(serde_json::Value::as_array)
.into_iter()
.flatten()
.filter_map(serde_json::Value::as_str)
.map(ToOwned::to_owned)
.collect()
}
pub(crate) fn mapping_preview_masked_catalog_api_key(
state: &AppState,
key: &StoredProviderCatalogKey,
) -> String {
let ciphertext = key.encrypted_api_key.trim();
if ciphertext.is_empty() {
return "***".to_string();
}
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)
.map(|value| {
if value.len() > 8 {
format!(
"{}***{}",
&value[..4],
&value[value.len().saturating_sub(4)..]
)
} else if value.len() >= 2 {
format!("{}***", &value[..2])
} else {
"***".to_string()
}
})
.unwrap_or_else(|| "***".to_string())
}
pub(crate) async fn build_admin_provider_mapping_preview_payload(
state: &AppState,
provider_id: &str,
) -> Option<serde_json::Value> {
if !state.has_provider_catalog_data_reader() || !state.has_global_model_data_reader() {
return None;
}
let provider = state
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
.await
.ok()
.and_then(|mut providers| providers.drain(..).next())?;
let mut keys = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
.await
.ok()
.unwrap_or_default()
.into_iter()
.filter(|key| key.allowed_models.is_some())
.collect::<Vec<_>>();
let total_keys_with_allowed_models = keys.len();
let truncated_keys =
total_keys_with_allowed_models.saturating_sub(ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_KEYS);
if keys.len() > ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_KEYS {
keys.truncate(ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_KEYS);
}
let public_models = state
.list_public_global_models(&PublicGlobalModelQuery {
offset: 0,
limit: ADMIN_PROVIDER_MAPPING_PREVIEW_FETCH_LIMIT,
is_active: None,
search: None,
})
.await
.ok()
.unwrap_or_else(|| {
aether_data_contracts::repository::global_models::StoredPublicGlobalModelPage {
items: Vec::new(),
total: 0,
}
})
.items;
let mut models_with_mappings = public_models
.into_iter()
.filter_map(|model| {
let mappings = public_global_model_mapping_patterns(&model);
(!mappings.is_empty()).then_some((model, mappings))
})
.collect::<Vec<_>>();
let total_models_with_mappings = models_with_mappings.len();
let truncated_models =
total_models_with_mappings.saturating_sub(ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_MODELS);
if models_with_mappings.len() > ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_MODELS {
models_with_mappings.truncate(ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_MODELS);
}
if models_with_mappings.is_empty() {
return Some(json!({
"provider_id": provider.id,
"provider_name": provider.name,
"keys": [],
"total_keys": 0,
"total_matches": 0,
"truncated": truncated_keys > 0 || truncated_models > 0,
"truncated_keys": truncated_keys,
"truncated_models": truncated_models,
}));
}
let mut key_payloads = Vec::new();
let mut total_matches = 0_u64;
for key in keys {
let allowed_models = json_string_list(key.allowed_models.as_ref());
if allowed_models.is_empty() {
continue;
}
let mut matching_global_models = Vec::new();
for (global_model, mappings) in &models_with_mappings {
let mut matched_models = Vec::new();
for allowed_model in &allowed_models {
for mapping_pattern in mappings {
if matches_model_mapping_for_models(mapping_pattern, allowed_model) {
matched_models.push(json!({
"allowed_model": allowed_model,
"mapping_pattern": mapping_pattern,
}));
break;
}
}
}
if !matched_models.is_empty() {
matching_global_models.push(json!({
"global_model_id": global_model.id,
"global_model_name": global_model.name,
"display_name": global_model
.display_name
.clone()
.unwrap_or_else(|| global_model.name.clone()),
"is_active": global_model.is_active,
"matched_models": matched_models,
}));
total_matches = total_matches.saturating_add(1);
}
}
if matching_global_models.is_empty() {
continue;
}
key_payloads.push(json!({
"key_id": key.id,
"key_name": key.name,
"masked_key": mapping_preview_masked_catalog_api_key(state, &key),
"is_active": key.is_active,
"allowed_models": allowed_models,
"matching_global_models": matching_global_models,
}));
}
Some(json!({
"provider_id": provider.id,
"provider_name": provider.name,
"keys": key_payloads,
"total_keys": key_payloads.len(),
"total_matches": total_matches,
"truncated": truncated_keys > 0 || truncated_models > 0,
"truncated_keys": truncated_keys,
"truncated_models": truncated_models,
}))
}
@@ -0,0 +1,755 @@
use super::oauth::{
normalize_string_id_list, refresh_antigravity_provider_quota_locally,
refresh_codex_provider_quota_locally, refresh_kiro_provider_quota_locally,
};
use super::write::{
build_admin_create_provider_key_record, build_admin_export_key_payload,
build_admin_provider_keys_payload, build_admin_reveal_key_payload,
build_admin_update_provider_key_record,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_clear_oauth_invalid_key_id, admin_export_key_id, admin_provider_id_for_keys,
admin_provider_id_for_refresh_quota, admin_reveal_key_id, admin_update_key_id,
AdminProviderKeyBatchDeleteRequest, AdminProviderKeyCreateRequest,
AdminProviderKeyUpdateRequest, AdminProviderQuotaRefreshRequest, OAUTH_ACCOUNT_BLOCK_PREFIX,
};
use crate::handlers::admin::shared::{
attach_admin_audit_response, build_admin_provider_key_response, query_param_value,
};
use crate::handlers::public::build_admin_keys_grouped_by_format_payload;
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::collections::BTreeSet;
use std::time::{SystemTime, UNIX_EPOCH};
pub(crate) async fn maybe_build_local_admin_endpoints_keys_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.control_decision.as_ref() else {
return Ok(None);
};
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("keys_grouped_by_format")
&& request_context.request_path == "/api/admin/endpoints/keys/grouped-by-format"
{
let Some(payload) = build_admin_keys_grouped_by_format_payload(state).await else {
return Ok(None);
};
return Ok(Some(Json(payload).into_response()));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("reveal_key")
&& request_context
.request_path
.starts_with("/api/admin/endpoints/keys/")
&& request_context.request_path.ends_with("/reveal")
{
let Some(key_id) = admin_reveal_key_id(&request_context.request_path) else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Key 不存在" })),
)
.into_response(),
));
};
let Some(key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Key {key_id} 不存在") })),
)
.into_response(),
));
};
return Ok(Some(match build_admin_reveal_key_payload(state, &key) {
Ok(payload) => attach_admin_audit_response(
Json(payload).into_response(),
"admin_provider_key_revealed",
"reveal_provider_key",
"provider_key",
&key_id,
),
Err(detail) => (
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
}));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("export_key")
&& request_context
.request_path
.starts_with("/api/admin/endpoints/keys/")
&& request_context.request_path.ends_with("/export")
{
let Some(key_id) = admin_export_key_id(&request_context.request_path) else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Key 不存在" })),
)
.into_response(),
));
};
let Some(key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Key {key_id} 不存在") })),
)
.into_response(),
));
};
return Ok(Some(
match build_admin_export_key_payload(state, &key).await {
Ok(payload) => attach_admin_audit_response(
Json(payload).into_response(),
"admin_provider_key_exported",
"export_provider_key",
"provider_key_export",
&key_id,
),
Err(detail) => (
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
},
));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("update_key")
&& request_context.request_method == http::Method::PUT
&& request_context
.request_path
.starts_with("/api/admin/endpoints/keys/")
{
let Some(key_id) = admin_update_key_id(&request_context.request_path) else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Key 不存在" })),
)
.into_response(),
));
};
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
if !state.has_provider_catalog_data_reader() {
return Ok(None);
}
let raw_value = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(value) => value,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let Some(raw_payload) = raw_value.as_object().cloned() else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
};
let payload = match serde_json::from_value::<AdminProviderKeyUpdateRequest>(raw_value) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let Some(existing_key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Key {key_id} 不存在") })),
)
.into_response(),
));
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&existing_key.provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {} 不存在", existing_key.provider_id) })),
)
.into_response(),
));
};
let updated_record = match build_admin_update_provider_key_record(
state,
&provider,
&existing_key,
&raw_payload,
payload,
)
.await
{
Ok(record) => record,
Err(detail) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
));
}
};
let Some(updated) = state.update_provider_catalog_key(&updated_record).await? else {
return Ok(None);
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
return Ok(Some(
Json(build_admin_provider_key_response(
state,
&updated,
now_unix_secs,
))
.into_response(),
));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("delete_key")
&& request_context.request_method == http::Method::DELETE
&& request_context
.request_path
.starts_with("/api/admin/endpoints/keys/")
{
let Some(key_id) = admin_update_key_id(&request_context.request_path) else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Key 不存在" })),
)
.into_response(),
));
};
let Some(_existing_key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Key {key_id} 不存在") })),
)
.into_response(),
));
};
if !state.delete_provider_catalog_key(&key_id).await? {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Key {key_id} 不存在") })),
)
.into_response(),
));
}
return Ok(Some(
Json(json!({
"message": format!("Key {key_id} 已删除")
}))
.into_response(),
));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("batch_delete_keys")
&& request_context.request_method == http::Method::POST
&& request_context.request_path == "/api/admin/endpoints/keys/batch-delete"
{
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
let payload =
match serde_json::from_slice::<AdminProviderKeyBatchDeleteRequest>(request_body) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
if payload.ids.len() > 100 {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "ids 最多 100 个" })),
)
.into_response(),
));
}
if payload.ids.is_empty() {
return Ok(Some(
Json(json!({
"success_count": 0,
"failed_count": 0,
"failed": []
}))
.into_response(),
));
}
let found_keys = state
.read_provider_catalog_keys_by_ids(&payload.ids)
.await?;
let found_ids = found_keys
.iter()
.map(|key| key.id.clone())
.collect::<BTreeSet<_>>();
let mut failed = payload
.ids
.iter()
.filter(|key_id| !found_ids.contains(*key_id))
.map(|key_id| json!({ "id": key_id, "error": "not found" }))
.collect::<Vec<_>>();
let mut success_count = 0usize;
for key_id in found_ids {
if state.delete_provider_catalog_key(&key_id).await? {
success_count += 1;
} else {
failed.push(json!({ "id": key_id, "error": "not found" }));
}
}
return Ok(Some(
Json(json!({
"success_count": success_count,
"failed_count": failed.len(),
"failed": failed,
}))
.into_response(),
));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("clear_oauth_invalid")
&& request_context.request_method == http::Method::POST
&& request_context
.request_path
.starts_with("/api/admin/endpoints/keys/")
&& request_context
.request_path
.ends_with("/clear-oauth-invalid")
{
let Some(key_id) = admin_clear_oauth_invalid_key_id(&request_context.request_path) else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Key 不存在" })),
)
.into_response(),
));
};
let Some(key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Key {key_id} 不存在") })),
)
.into_response(),
));
};
if key.oauth_invalid_at_unix_secs.is_none() {
return Ok(Some(
Json(json!({
"message": "该 Key 当前无失效标记,无需清除"
}))
.into_response(),
));
}
state
.clear_provider_catalog_key_oauth_invalid_marker(&key_id)
.await?;
return Ok(Some(
Json(json!({
"message": "已清除 OAuth 失效标记"
}))
.into_response(),
));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("refresh_quota")
&& request_context.request_method == http::Method::POST
&& request_context
.request_path
.starts_with("/api/admin/endpoints/providers/")
&& request_context.request_path.ends_with("/refresh-quota")
{
let Some(provider_id) = admin_provider_id_for_refresh_quota(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
));
};
let normalized_provider_type = provider.provider_type.trim().to_ascii_lowercase();
let payload = if let Some(request_body) = request_body {
match serde_json::from_slice::<AdminProviderQuotaRefreshRequest>(request_body) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
}
} else {
AdminProviderQuotaRefreshRequest { key_ids: None }
};
let raw_key_ids = payload.key_ids;
let selected_key_ids = normalize_string_id_list(raw_key_ids.clone());
let explicit_key_ids_requested = raw_key_ids.is_some();
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let endpoint = match normalized_provider_type.as_str() {
"codex" => endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("openai:cli")
}),
"antigravity" => endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& (endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:chat")
|| endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:cli"))
}),
"kiro" => endpoints
.iter()
.find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("claude:cli")
})
.cloned()
.or_else(|| endpoints.into_iter().find(|endpoint| endpoint.is_active)),
_ => return Ok(None),
};
let Some(endpoint) = endpoint else {
let detail = match normalized_provider_type.as_str() {
"codex" => "找不到有效的 openai:cli 端点",
"antigravity" => "找不到有效的 gemini:chat/gemini:cli 端点",
"kiro" => "找不到有效的 Kiro 端点",
_ => "找不到有效端点",
};
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
));
};
let mut keys = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
keys = if let Some(selected_key_ids) = selected_key_ids.as_ref() {
if selected_key_ids.is_empty() {
Vec::new()
} else {
let selected = selected_key_ids.iter().cloned().collect::<BTreeSet<_>>();
keys.into_iter()
.filter(|key| selected.contains(&key.id))
.collect()
}
} else {
keys.into_iter()
.filter(|key| {
key.is_active
|| key
.oauth_invalid_reason
.as_deref()
.map(str::trim)
.is_some_and(|value| value.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX))
})
.collect()
};
if explicit_key_ids_requested && selected_key_ids.is_none() {
return Ok(Some(
Json(json!({
"success": 0,
"failed": 0,
"total": 0,
"results": [],
"message": "未提供可刷新的 Key",
"auto_removed": 0,
}))
.into_response(),
));
}
if keys.is_empty() {
let message = if explicit_key_ids_requested {
"未提供可刷新的 Key"
} else {
"没有可刷新的 Key"
};
return Ok(Some(
Json(json!({
"success": 0,
"failed": 0,
"total": 0,
"results": [],
"message": message,
"auto_removed": 0,
}))
.into_response(),
));
}
let Some(payload) = (match normalized_provider_type.as_str() {
"codex" => {
refresh_codex_provider_quota_locally(state, &provider, &endpoint, keys).await?
}
"kiro" => {
refresh_kiro_provider_quota_locally(state, &provider, &endpoint, keys).await?
}
"antigravity" => {
refresh_antigravity_provider_quota_locally(state, &provider, &endpoint, keys)
.await?
}
_ => None,
}) else {
return Ok(None);
};
return Ok(Some(Json(payload).into_response()));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("create_provider_key")
&& request_context.request_method == http::Method::POST
&& request_context
.request_path
.starts_with("/api/admin/endpoints/providers/")
&& request_context.request_path.ends_with("/keys")
{
let Some(provider_id) = admin_provider_id_for_keys(&request_context.request_path) else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
if !state.has_provider_catalog_data_reader() {
return Ok(None);
}
let payload = match serde_json::from_slice::<AdminProviderKeyCreateRequest>(request_body) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
));
};
let record = match build_admin_create_provider_key_record(state, &provider, payload).await {
Ok(record) => record,
Err(detail) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
));
}
};
let Some(created) = state.create_provider_catalog_key(&record).await? else {
return Ok(None);
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
return Ok(Some(
Json(build_admin_provider_key_response(
state,
&created,
now_unix_secs,
))
.into_response(),
));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("list_provider_keys")
&& request_context
.request_path
.starts_with("/api/admin/endpoints/providers/")
&& request_context.request_path.ends_with("/keys")
{
let Some(provider_id) = admin_provider_id_for_keys(&request_context.request_path) else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let skip = query_param_value(request_context.request_query_string.as_deref(), "skip")
.and_then(|value| value.parse::<usize>().ok())
.unwrap_or(0);
let limit = query_param_value(request_context.request_query_string.as_deref(), "limit")
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| *value > 0)
.unwrap_or(100);
return Ok(Some(
match build_admin_provider_keys_payload(state, &provider_id, skip, limit).await {
Some(payload) => Json(payload).into_response(),
None => (
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
},
));
}
Ok(None)
}
@@ -0,0 +1,355 @@
use crate::api::ai::{admin_default_body_rules_for_signature, admin_endpoint_signature_parts};
use crate::handlers::public::{admin_requested_force_stream, normalize_admin_base_url};
use crate::provider_transport::provider_types::provider_type_is_fixed;
use crate::AppState;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
use uuid::Uuid;
use super::payloads::{build_admin_provider_endpoint_response, endpoint_key_counts_by_format};
use super::payloads::{AdminProviderEndpointCreateRequest, AdminProviderEndpointUpdateRequest};
pub(crate) async fn build_admin_provider_endpoints_payload(
state: &AppState,
provider_id: &str,
skip: usize,
limit: usize,
) -> Option<serde_json::Value> {
if !state.has_provider_catalog_data_reader() {
return None;
}
let provider = state
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
.await
.ok()
.and_then(|mut providers| providers.drain(..).next())?;
let mut endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await
.ok()
.unwrap_or_default();
endpoints.sort_by(|left, right| {
right
.created_at_unix_secs
.unwrap_or_default()
.cmp(&left.created_at_unix_secs.unwrap_or_default())
.then_with(|| left.id.cmp(&right.id))
});
let keys = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
.await
.ok()
.unwrap_or_default();
let (total_keys_by_format, active_keys_by_format) = endpoint_key_counts_by_format(&keys);
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
Some(serde_json::Value::Array(
endpoints
.into_iter()
.skip(skip)
.take(limit)
.map(|endpoint| {
build_admin_provider_endpoint_response(
&endpoint,
&provider.name,
total_keys_by_format
.get(endpoint.api_format.as_str())
.copied()
.unwrap_or(0),
active_keys_by_format
.get(endpoint.api_format.as_str())
.copied()
.unwrap_or(0),
now_unix_secs,
)
})
.collect(),
))
}
pub(crate) async fn build_admin_endpoint_payload(
state: &AppState,
endpoint_id: &str,
) -> Option<serde_json::Value> {
if !state.has_provider_catalog_data_reader() {
return None;
}
let endpoint = state
.read_provider_catalog_endpoints_by_ids(&[endpoint_id.to_string()])
.await
.ok()
.and_then(|mut endpoints| endpoints.drain(..).next())?;
let provider = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&endpoint.provider_id))
.await
.ok()
.and_then(|mut providers| providers.drain(..).next())?;
let keys = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&endpoint.provider_id))
.await
.ok()
.unwrap_or_default();
let (total_keys_by_format, active_keys_by_format) = endpoint_key_counts_by_format(&keys);
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
Some(build_admin_provider_endpoint_response(
&endpoint,
&provider.name,
total_keys_by_format
.get(endpoint.api_format.as_str())
.copied()
.unwrap_or(0),
active_keys_by_format
.get(endpoint.api_format.as_str())
.copied()
.unwrap_or(0),
now_unix_secs,
))
}
pub(crate) async fn build_admin_create_provider_endpoint_record(
state: &AppState,
provider: &StoredProviderCatalogProvider,
payload: AdminProviderEndpointCreateRequest,
) -> Result<StoredProviderCatalogEndpoint, String> {
if payload.provider_id.trim() != provider.id {
return Err("provider_id 不匹配".to_string());
}
if provider_type_is_fixed(&provider.provider_type) {
return Err("固定类型 Provider 不允许手动新增 Endpoint".to_string());
}
if !(0..=999).contains(&payload.max_retries) {
return Err("max_retries 必须在 0 到 999 之间".to_string());
}
let (normalized_api_format, api_family, endpoint_kind) =
admin_endpoint_signature_parts(&payload.api_format)
.ok_or_else(|| format!("无效的 api_format: {}", payload.api_format))?;
let base_url = normalize_admin_base_url(&payload.base_url)?;
let existing_endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await
.map_err(|err| format!("{err:?}"))?;
if existing_endpoints
.iter()
.any(|endpoint| endpoint.api_format == normalized_api_format)
{
return Err(format!(
"Provider {} 已存在 {} 格式的 Endpoint",
provider.name, normalized_api_format
));
}
let body_rules = match payload.body_rules {
Some(value) => Some(value),
None => admin_default_body_rules_for_signature(
normalized_api_format,
Some(provider.provider_type.as_str()),
)
.and_then(|(_, rules)| (!rules.is_empty()).then_some(serde_json::Value::Array(rules))),
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
StoredProviderCatalogEndpoint::new(
Uuid::new_v4().to_string(),
provider.id.clone(),
normalized_api_format.to_string(),
Some(api_family.to_string()),
Some(endpoint_kind.to_string()),
true,
)
.map_err(|err| err.to_string())?
.with_timestamps(Some(now_unix_secs), Some(now_unix_secs))
.with_transport_fields(
base_url,
payload.header_rules,
body_rules,
Some(payload.max_retries),
payload.custom_path.and_then(|value| {
let trimmed = value.trim().to_string();
(!trimmed.is_empty()).then_some(trimmed)
}),
payload.config,
payload.format_acceptance_config,
payload.proxy,
)
.map_err(|err| err.to_string())
}
pub(crate) async fn build_admin_update_provider_endpoint_record(
state: &AppState,
provider: &StoredProviderCatalogProvider,
existing_endpoint: &StoredProviderCatalogEndpoint,
raw_payload: &serde_json::Map<String, serde_json::Value>,
payload: AdminProviderEndpointUpdateRequest,
) -> Result<StoredProviderCatalogEndpoint, String> {
let mut updated = existing_endpoint.clone();
if provider_type_is_fixed(&provider.provider_type)
&& (raw_payload.contains_key("base_url") || raw_payload.contains_key("custom_path"))
{
return Err("固定类型 Provider 的 Endpoint 不允许修改 base_url/custom_path".to_string());
}
if let Some(value) = raw_payload.get("base_url") {
let Some(base_url) = payload.base_url.as_deref() else {
return Err(if value.is_null() {
"base_url 不能为空".to_string()
} else {
"base_url 必须是字符串".to_string()
});
};
updated.base_url = normalize_admin_base_url(base_url)?;
}
if raw_payload.contains_key("custom_path") {
updated.custom_path = payload.custom_path;
}
if let Some(value) = raw_payload.get("header_rules") {
if !value.is_null() && !value.is_array() {
return Err("header_rules 必须是数组或 null".to_string());
}
updated.header_rules = if value.is_null() {
None
} else {
payload.header_rules
};
}
if let Some(value) = raw_payload.get("body_rules") {
if !value.is_null() && !value.is_array() {
return Err("body_rules 必须是数组或 null".to_string());
}
updated.body_rules = if value.is_null() {
None
} else {
payload.body_rules
};
}
if let Some(value) = raw_payload.get("max_retries") {
let Some(max_retries) = payload.max_retries else {
return Err(if value.is_null() {
"max_retries 必须是 0 到 999 之间的整数".to_string()
} else {
"max_retries 必须是整数".to_string()
});
};
if !(0..=999).contains(&max_retries) {
return Err("max_retries 必须在 0 到 999 之间".to_string());
}
updated.max_retries = Some(max_retries);
}
if raw_payload.contains_key("is_active") {
let Some(is_active) = payload.is_active else {
return Err("is_active 必须是布尔值".to_string());
};
updated.is_active = is_active;
}
if let Some(value) = raw_payload.get("config") {
if !value.is_null() && !value.is_object() {
return Err("config 必须是对象或 null".to_string());
}
updated.config = if value.is_null() {
None
} else {
payload.config
};
}
if let Some(value) = raw_payload.get("proxy") {
if value.is_null() {
updated.proxy = None;
} else {
let Some(mut proxy) = payload.proxy.and_then(|value| value.as_object().cloned()) else {
return Err("proxy 必须是对象或 null".to_string());
};
if !proxy.contains_key("password") {
if let Some(old_password) = existing_endpoint
.proxy
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|proxy| proxy.get("password"))
.and_then(serde_json::Value::as_str)
.filter(|value| !value.is_empty())
{
proxy.insert("password".to_string(), json!(old_password));
}
}
updated.proxy = Some(serde_json::Value::Object(proxy));
}
}
if let Some(value) = raw_payload.get("format_acceptance_config") {
if !value.is_null() && !value.is_object() {
return Err("format_acceptance_config 必须是对象或 null".to_string());
}
updated.format_acceptance_config = if value.is_null() {
None
} else {
payload.format_acceptance_config
};
}
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if provider_type == "codex" && existing_endpoint.api_format == "openai:cli" {
let has_config_in_payload = raw_payload.contains_key("config");
let config_payload = if has_config_in_payload {
updated.config.clone().unwrap_or_else(|| json!({}))
} else {
existing_endpoint
.config
.clone()
.unwrap_or_else(|| json!({}))
};
let mut config = config_payload.as_object().cloned().unwrap_or_default();
let requested = config
.get("upstream_stream_policy")
.or_else(|| config.get("upstreamStreamPolicy"))
.or_else(|| config.get("upstream_stream"));
if has_config_in_payload
&& requested.is_some()
&& !admin_requested_force_stream(requested.expect("checked above"))
{
return Err("Codex OpenAI CLI 端点固定为强制流式,不允许修改".to_string());
}
config.remove("upstreamStreamPolicy");
config.remove("upstream_stream");
config.insert("upstream_stream_policy".to_string(), json!("force_stream"));
updated.config = Some(serde_json::Value::Object(config));
}
let (_, api_family, endpoint_kind) = admin_endpoint_signature_parts(&updated.api_format)
.ok_or_else(|| format!("无效的 api_format: {}", updated.api_format))?;
updated.api_family = Some(api_family.to_string());
updated.endpoint_kind = Some(endpoint_kind.to_string());
updated.updated_at_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs());
let _ = state;
Ok(updated)
}
@@ -0,0 +1,31 @@
pub(super) fn admin_provider_id_for_endpoints(request_path: &str) -> Option<String> {
let raw = request_path.strip_prefix("/api/admin/endpoints/providers/")?;
let raw = raw.strip_suffix("/endpoints")?;
let normalized = raw.trim().trim_matches('/');
if normalized.is_empty() {
None
} else {
Some(normalized.to_string())
}
}
pub(super) fn admin_endpoint_id(request_path: &str) -> Option<String> {
let raw = request_path.strip_prefix("/api/admin/endpoints/")?;
let normalized = raw.trim().trim_matches('/');
if normalized.is_empty() || normalized.contains('/') {
None
} else {
Some(normalized.to_string())
}
}
pub(super) fn admin_default_body_rules_api_format(request_path: &str) -> Option<String> {
let raw = request_path.strip_prefix("/api/admin/endpoints/defaults/")?;
let raw = raw.strip_suffix("/body-rules")?;
let normalized = raw.trim().trim_matches('/');
if normalized.is_empty() {
None
} else {
Some(normalized.to_string())
}
}
@@ -0,0 +1,478 @@
mod builders;
mod extractors;
mod payloads;
use self::builders::{
build_admin_create_provider_endpoint_record, build_admin_endpoint_payload,
build_admin_provider_endpoints_payload, build_admin_update_provider_endpoint_record,
};
use self::extractors::{
admin_default_body_rules_api_format, admin_endpoint_id, admin_provider_id_for_endpoints,
};
use self::payloads::endpoint_key_counts_by_format;
use self::payloads::{
build_admin_provider_endpoint_response, key_api_formats_without_entry,
AdminProviderEndpointCreateRequest, AdminProviderEndpointUpdateRequest,
};
use crate::api::ai::admin_default_body_rules_for_signature;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::shared::query_param_value;
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
const ADMIN_ENDPOINTS_DATA_UNAVAILABLE_DETAIL: &str = "Admin endpoint data unavailable";
fn build_admin_endpoints_data_unavailable_response() -> Response<Body> {
(
http::StatusCode::SERVICE_UNAVAILABLE,
Json(json!({ "detail": ADMIN_ENDPOINTS_DATA_UNAVAILABLE_DETAIL })),
)
.into_response()
}
pub(crate) async fn maybe_build_local_admin_endpoints_routes_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.control_decision.as_ref() else {
return Ok(None);
};
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("list_provider_endpoints")
&& request_context
.request_path
.starts_with("/api/admin/endpoints/providers/")
&& request_context.request_path.ends_with("/endpoints")
{
if !state.has_provider_catalog_data_reader() {
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
}
let Some(provider_id) = admin_provider_id_for_endpoints(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let skip = query_param_value(request_context.request_query_string.as_deref(), "skip")
.and_then(|value| value.parse::<usize>().ok())
.unwrap_or(0);
let limit = query_param_value(request_context.request_query_string.as_deref(), "limit")
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| *value > 0)
.unwrap_or(100);
return Ok(Some(
match build_admin_provider_endpoints_payload(state, &provider_id, skip, limit).await {
Some(payload) => Json(payload).into_response(),
None => (
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
},
));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("create_endpoint")
&& request_context.request_method == http::Method::POST
&& request_context
.request_path
.starts_with("/api/admin/endpoints/providers/")
&& request_context.request_path.ends_with("/endpoints")
{
if !state.has_provider_catalog_data_reader() || !state.has_provider_catalog_data_writer() {
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
}
let Some(provider_id) = admin_provider_id_for_endpoints(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
let payload =
match serde_json::from_slice::<AdminProviderEndpointCreateRequest>(request_body) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
));
};
let record =
match build_admin_create_provider_endpoint_record(state, &provider, payload).await {
Ok(record) => record,
Err(detail) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
));
}
};
let Some(created) = state.create_provider_catalog_endpoint(&record).await? else {
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
return Ok(Some(
Json(build_admin_provider_endpoint_response(
&created,
&provider.name,
0,
0,
now_unix_secs,
))
.into_response(),
));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("update_endpoint")
&& request_context.request_method == http::Method::PUT
&& request_context
.request_path
.starts_with("/api/admin/endpoints/")
{
if !state.has_provider_catalog_data_reader() || !state.has_provider_catalog_data_writer() {
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
}
let Some(endpoint_id) = admin_endpoint_id(&request_context.request_path) else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Endpoint 不存在" })),
)
.into_response(),
));
};
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
let raw_value = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(value) => value,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let Some(raw_payload) = raw_value.as_object().cloned() else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
};
let payload = match serde_json::from_value::<AdminProviderEndpointUpdateRequest>(raw_value)
{
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let Some(existing_endpoint) = state
.read_provider_catalog_endpoints_by_ids(std::slice::from_ref(&endpoint_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Endpoint {endpoint_id} 不存在") })),
)
.into_response(),
));
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(
&existing_endpoint.provider_id,
))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {} 不存在", existing_endpoint.provider_id) })),
)
.into_response(),
));
};
let updated_record = match build_admin_update_provider_endpoint_record(
state,
&provider,
&existing_endpoint,
&raw_payload,
payload,
)
.await
{
Ok(record) => record,
Err(detail) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
));
}
};
let Some(updated) = state
.update_provider_catalog_endpoint(&updated_record)
.await?
else {
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let keys = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
.await
.unwrap_or_default();
let (total_keys_by_format, active_keys_by_format) = endpoint_key_counts_by_format(&keys);
return Ok(Some(
Json(build_admin_provider_endpoint_response(
&updated,
&provider.name,
total_keys_by_format
.get(updated.api_format.as_str())
.copied()
.unwrap_or(0),
active_keys_by_format
.get(updated.api_format.as_str())
.copied()
.unwrap_or(0),
now_unix_secs,
))
.into_response(),
));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("delete_endpoint")
&& request_context.request_method == http::Method::DELETE
&& request_context
.request_path
.starts_with("/api/admin/endpoints/")
{
if !state.has_provider_catalog_data_reader() || !state.has_provider_catalog_data_writer() {
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
}
let Some(endpoint_id) = admin_endpoint_id(&request_context.request_path) else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Endpoint 不存在" })),
)
.into_response(),
));
};
let Some(existing_endpoint) = state
.read_provider_catalog_endpoints_by_ids(std::slice::from_ref(&endpoint_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Endpoint {endpoint_id} 不存在") })),
)
.into_response(),
));
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let keys = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(
&existing_endpoint.provider_id,
))
.await
.unwrap_or_default();
let mut affected_keys_count = 0usize;
for key in keys {
let Some(updated_formats) =
key_api_formats_without_entry(&key, existing_endpoint.api_format.as_str())
else {
continue;
};
let mut updated_key = key.clone();
updated_key.api_formats = Some(serde_json::Value::Array(
updated_formats
.into_iter()
.map(serde_json::Value::String)
.collect(),
));
updated_key.updated_at_unix_secs = Some(now_unix_secs);
if state
.update_provider_catalog_key(&updated_key)
.await?
.is_none()
{
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
}
affected_keys_count += 1;
}
if !state.delete_provider_catalog_endpoint(&endpoint_id).await? {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Endpoint {endpoint_id} 不存在") })),
)
.into_response(),
));
}
return Ok(Some(
Json(json!({
"message": format!("Endpoint {endpoint_id} 已删除"),
"affected_keys_count": affected_keys_count,
}))
.into_response(),
));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("get_endpoint")
&& request_context
.request_path
.starts_with("/api/admin/endpoints/")
{
if !state.has_provider_catalog_data_reader() {
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
}
let Some(endpoint_id) = admin_endpoint_id(&request_context.request_path) else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Endpoint 不存在" })),
)
.into_response(),
));
};
return Ok(Some(
match build_admin_endpoint_payload(state, &endpoint_id).await {
Some(payload) => Json(payload).into_response(),
None => (
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Endpoint {endpoint_id} 不存在") })),
)
.into_response(),
},
));
}
if decision.route_family.as_deref() == Some("endpoints_manage")
&& decision.route_kind.as_deref() == Some("default_body_rules")
&& request_context
.request_path
.starts_with("/api/admin/endpoints/defaults/")
&& request_context.request_path.ends_with("/body-rules")
{
let Some(api_format) = admin_default_body_rules_api_format(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "无效的 api_format" })),
)
.into_response(),
));
};
let provider_type = query_param_value(
request_context.request_query_string.as_deref(),
"provider_type",
);
return Ok(Some(
match admin_default_body_rules_for_signature(&api_format, provider_type.as_deref()) {
Some((normalized_api_format, body_rules)) => Json(json!({
"api_format": normalized_api_format,
"body_rules": body_rules,
}))
.into_response(),
None => (
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": format!("无效的 api_format: {api_format}") })),
)
.into_response(),
},
));
}
Ok(None)
}
@@ -0,0 +1,146 @@
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
};
use serde::Deserialize;
use serde_json::json;
use std::collections::BTreeMap;
pub(super) fn key_api_formats_without_entry(
key: &StoredProviderCatalogKey,
api_format: &str,
) -> Option<Vec<String>> {
let current_formats =
crate::handlers::admin::shared::json_string_list(key.api_formats.as_ref());
if !current_formats
.iter()
.any(|candidate| candidate == api_format)
{
return None;
}
Some(
current_formats
.into_iter()
.filter(|candidate| candidate != api_format)
.collect(),
)
}
pub(super) fn endpoint_key_counts_by_format(
keys: &[StoredProviderCatalogKey],
) -> (BTreeMap<String, usize>, BTreeMap<String, usize>) {
let mut total = BTreeMap::new();
let mut active = BTreeMap::new();
for key in keys {
let Some(formats) = key
.api_formats
.as_ref()
.and_then(serde_json::Value::as_array)
else {
continue;
};
for api_format in formats.iter().filter_map(serde_json::Value::as_str) {
*total.entry(api_format.to_string()).or_insert(0) += 1;
if key.is_active {
*active.entry(api_format.to_string()).or_insert(0) += 1;
}
}
}
(total, active)
}
pub(super) fn build_admin_provider_endpoint_response(
endpoint: &StoredProviderCatalogEndpoint,
provider_name: &str,
total_keys: usize,
active_keys: usize,
now_unix_secs: u64,
) -> serde_json::Value {
json!({
"id": endpoint.id,
"provider_id": endpoint.provider_id,
"provider_name": provider_name,
"api_format": endpoint.api_format,
"base_url": endpoint.base_url,
"custom_path": endpoint.custom_path,
"header_rules": endpoint.header_rules,
"body_rules": endpoint.body_rules,
"max_retries": endpoint.max_retries.unwrap_or(2),
"is_active": endpoint.is_active,
"config": endpoint.config,
"proxy": masked_proxy_value(endpoint.proxy.as_ref()),
"format_acceptance_config": endpoint.format_acceptance_config,
"total_keys": total_keys,
"active_keys": active_keys,
"created_at": endpoint_timestamp_or_now(endpoint.created_at_unix_secs, now_unix_secs),
"updated_at": endpoint_timestamp_or_now(endpoint.updated_at_unix_secs, now_unix_secs),
})
}
fn masked_proxy_value(proxy: Option<&serde_json::Value>) -> serde_json::Value {
let Some(proxy) = proxy.and_then(serde_json::Value::as_object) else {
return serde_json::Value::Null;
};
let mut masked = proxy.clone();
if masked
.get("password")
.and_then(serde_json::Value::as_str)
.is_some_and(|value| !value.trim().is_empty())
{
masked.insert("password".to_string(), json!("***"));
}
serde_json::Value::Object(masked)
}
fn endpoint_timestamp_or_now(value: Option<u64>, now_unix_secs: u64) -> serde_json::Value {
unix_secs_to_rfc3339(value.unwrap_or(now_unix_secs))
.map(serde_json::Value::String)
.unwrap_or(serde_json::Value::Null)
}
fn default_admin_endpoint_max_retries() -> i32 {
2
}
#[derive(Debug, Deserialize)]
pub(super) struct AdminProviderEndpointCreateRequest {
pub(super) provider_id: String,
pub(super) api_format: String,
pub(super) base_url: String,
#[serde(default)]
pub(super) custom_path: Option<String>,
#[serde(default)]
pub(super) header_rules: Option<serde_json::Value>,
#[serde(default)]
pub(super) body_rules: Option<serde_json::Value>,
#[serde(default = "default_admin_endpoint_max_retries")]
pub(super) max_retries: i32,
#[serde(default)]
pub(super) config: Option<serde_json::Value>,
#[serde(default)]
pub(super) proxy: Option<serde_json::Value>,
#[serde(default)]
pub(super) format_acceptance_config: Option<serde_json::Value>,
}
#[derive(Debug, Deserialize)]
pub(super) struct AdminProviderEndpointUpdateRequest {
#[serde(default)]
pub(super) base_url: Option<String>,
#[serde(default)]
pub(super) custom_path: Option<String>,
#[serde(default)]
pub(super) header_rules: Option<serde_json::Value>,
#[serde(default)]
pub(super) body_rules: Option<serde_json::Value>,
#[serde(default)]
pub(super) max_retries: Option<i32>,
#[serde(default)]
pub(super) is_active: Option<bool>,
#[serde(default)]
pub(super) config: Option<serde_json::Value>,
#[serde(default)]
pub(super) proxy: Option<serde_json::Value>,
#[serde(default)]
pub(super) format_acceptance_config: Option<serde_json::Value>,
}
@@ -0,0 +1,24 @@
pub(crate) mod endpoint_keys;
pub(crate) mod endpoints_admin;
pub(crate) mod oauth;
pub(crate) mod ops;
pub(crate) mod pool;
pub(crate) mod pool_admin;
pub(crate) mod shared;
pub(crate) mod write;
use super::auth::build_proxy_error_response;
mod crud;
mod delete_task;
mod models;
mod query;
mod strategy;
mod summary;
pub(crate) use self::crud::maybe_build_local_admin_providers_response;
pub(crate) use self::models::maybe_build_local_admin_provider_models_response;
pub(crate) use self::oauth::maybe_build_local_admin_provider_oauth_response;
pub(crate) use self::ops::maybe_build_local_admin_provider_ops_response;
pub(crate) use self::query::maybe_build_local_admin_provider_query_response;
pub(crate) use self::strategy::maybe_build_local_admin_provider_strategy_response;
@@ -0,0 +1,95 @@
use super::super::super::model::build_admin_batch_assign_global_models_payload;
use crate::control::GatewayControlDecision;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_assign_global_models_path, AdminBatchAssignGlobalModelsRequest,
};
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) async fn maybe_handle(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
decision: &GatewayControlDecision,
) -> Result<Option<Response<Body>>, GatewayError> {
if decision.route_family.as_deref() == Some("provider_models_manage")
&& decision.route_kind.as_deref() == Some("assign_global_models")
&& request_context.request_method == http::Method::POST
{
let Some(provider_id) =
admin_provider_assign_global_models_path(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let Some(_provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
));
};
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
let payload =
match serde_json::from_slice::<AdminBatchAssignGlobalModelsRequest>(request_body) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let payload = match build_admin_batch_assign_global_models_payload(
state,
&provider_id,
payload.global_model_ids,
)
.await
{
Ok(payload) => payload,
Err(detail) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
));
}
};
return Ok(Some(Json(payload).into_response()));
}
Ok(None)
}
@@ -0,0 +1,48 @@
use super::super::super::model::build_admin_provider_available_source_models_payload;
use crate::control::GatewayControlDecision;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::admin_provider_available_source_models_path;
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) async fn maybe_handle(
state: &AppState,
request_context: &GatewayPublicRequestContext,
_request_body: Option<&Bytes>,
decision: &GatewayControlDecision,
) -> Result<Option<Response<Body>>, GatewayError> {
if decision.route_family.as_deref() == Some("provider_models_manage")
&& decision.route_kind.as_deref() == Some("available_source_models")
&& request_context.request_method == http::Method::GET
{
let Some(provider_id) =
admin_provider_available_source_models_path(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
return Ok(Some(
match build_admin_provider_available_source_models_payload(state, &provider_id).await {
Some(payload) => Json(payload).into_response(),
None => (
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
},
));
}
Ok(None)
}
@@ -0,0 +1,150 @@
use super::super::super::model::{
admin_provider_model_name_exists, build_admin_provider_model_create_record,
build_admin_provider_model_response,
};
use crate::control::GatewayControlDecision;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_models_batch_path, AdminProviderModelCreateRequest,
};
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::collections::BTreeSet;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) async fn maybe_handle(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
decision: &GatewayControlDecision,
) -> Result<Option<Response<Body>>, GatewayError> {
if decision.route_family.as_deref() == Some("provider_models_manage")
&& decision.route_kind.as_deref() == Some("batch_create_provider_models")
&& request_context.request_method == http::Method::POST
&& request_context.request_path.ends_with("/models/batch")
{
let Some(provider_id) = admin_provider_models_batch_path(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let Some(_provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
));
};
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
let payloads =
match serde_json::from_slice::<Vec<AdminProviderModelCreateRequest>>(request_body) {
Ok(payloads) => payloads,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 数组" })),
)
.into_response(),
));
}
};
let mut created = Vec::new();
let mut seen = BTreeSet::new();
for payload in payloads {
let normalized_name = payload.provider_model_name.trim().to_string();
if normalized_name.is_empty() {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "provider_model_name 不能为空" })),
)
.into_response(),
));
}
if !seen.insert(normalized_name.clone()) {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": format!("批量请求中包含重复模型 {normalized_name}") })),
)
.into_response(),
));
}
if admin_provider_model_name_exists(state, &provider_id, &normalized_name, None).await?
{
continue;
}
let record = match build_admin_provider_model_create_record(
state,
&provider_id,
payload,
)
.await
{
Ok(record) => record,
Err(detail) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
));
}
};
let Some(model) = state.create_admin_provider_model(&record).await? else {
return Ok(Some(
(
http::StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({ "detail": "批量创建模型失败" })),
)
.into_response(),
));
};
created.push(model);
}
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
return Ok(Some(
Json(serde_json::Value::Array(
created
.iter()
.map(|model| build_admin_provider_model_response(model, now_unix_secs))
.collect(),
))
.into_response(),
));
}
Ok(None)
}
@@ -0,0 +1,110 @@
use super::super::super::model::{
build_admin_provider_model_create_record, build_admin_provider_model_response,
};
use crate::control::GatewayControlDecision;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_id_for_models_list, AdminProviderModelCreateRequest,
};
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) async fn maybe_handle(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
decision: &GatewayControlDecision,
) -> Result<Option<Response<Body>>, GatewayError> {
if decision.route_family.as_deref() == Some("provider_models_manage")
&& decision.route_kind.as_deref() == Some("create_provider_model")
&& request_context.request_method == http::Method::POST
&& request_context.request_path.ends_with("/models")
{
let Some(provider_id) = admin_provider_id_for_models_list(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let Some(_provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
));
};
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
let payload = match serde_json::from_slice::<AdminProviderModelCreateRequest>(request_body)
{
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let record =
match build_admin_provider_model_create_record(state, &provider_id, payload).await {
Ok(record) => record,
Err(detail) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
));
}
};
return Ok(Some(
match state.create_admin_provider_model(&record).await? {
Some(created) => {
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
Json(build_admin_provider_model_response(&created, now_unix_secs))
.into_response()
}
None => (
http::StatusCode::INTERNAL_SERVER_ERROR,
Json(json!({ "detail": "创建模型失败" })),
)
.into_response(),
},
));
}
Ok(None)
}
@@ -0,0 +1,68 @@
use crate::control::GatewayControlDecision;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::admin_provider_model_route_parts;
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) async fn maybe_handle(
state: &AppState,
request_context: &GatewayPublicRequestContext,
_request_body: Option<&Bytes>,
decision: &GatewayControlDecision,
) -> Result<Option<Response<Body>>, GatewayError> {
if decision.route_family.as_deref() == Some("provider_models_manage")
&& decision.route_kind.as_deref() == Some("delete_provider_model")
&& request_context.request_method == http::Method::DELETE
&& request_context.request_path.contains("/models/")
{
let Some((provider_id, model_id)) =
admin_provider_model_route_parts(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Model 不存在" })),
)
.into_response(),
));
};
let Some(existing) = state
.get_admin_provider_model(&provider_id, &model_id)
.await?
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Model {model_id} 不存在") })),
)
.into_response(),
));
};
if !state
.delete_admin_provider_model(&provider_id, &model_id)
.await?
{
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Model {model_id} 不存在") })),
)
.into_response(),
));
}
return Ok(Some(
Json(json!({
"message": format!("Model '{}' deleted successfully", existing.provider_model_name),
}))
.into_response(),
));
}
Ok(None)
}
@@ -0,0 +1,52 @@
use super::super::super::model::build_admin_provider_model_payload;
use crate::control::GatewayControlDecision;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::admin_provider_model_route_parts;
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) async fn maybe_handle(
state: &AppState,
request_context: &GatewayPublicRequestContext,
_request_body: Option<&Bytes>,
decision: &GatewayControlDecision,
) -> Result<Option<Response<Body>>, GatewayError> {
if decision.route_family.as_deref() == Some("provider_models_manage")
&& decision.route_kind.as_deref() == Some("get_provider_model")
&& request_context.request_method == http::Method::GET
&& request_context
.request_path
.starts_with("/api/admin/providers/")
&& request_context.request_path.contains("/models/")
{
let Some((provider_id, model_id)) =
admin_provider_model_route_parts(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Model 不存在" })),
)
.into_response(),
));
};
return Ok(Some(
match build_admin_provider_model_payload(state, &provider_id, &model_id).await {
Some(payload) => Json(payload).into_response(),
None => (
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Model {model_id} 不存在") })),
)
.into_response(),
},
));
}
Ok(None)
}
@@ -0,0 +1,89 @@
use super::super::super::model::build_admin_import_provider_models_payload;
use crate::control::GatewayControlDecision;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_import_models_path, AdminImportProviderModelsRequest,
};
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) async fn maybe_handle(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
decision: &GatewayControlDecision,
) -> Result<Option<Response<Body>>, GatewayError> {
if decision.route_family.as_deref() == Some("provider_models_manage")
&& decision.route_kind.as_deref() == Some("import_from_upstream")
&& request_context.request_method == http::Method::POST
{
let Some(provider_id) = admin_provider_import_models_path(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let Some(_provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
));
};
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
let payload = match serde_json::from_slice::<AdminImportProviderModelsRequest>(request_body)
{
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let payload =
match build_admin_import_provider_models_payload(state, &provider_id, payload).await {
Ok(payload) => payload,
Err(detail) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
));
}
};
return Ok(Some(Json(payload).into_response()));
}
Ok(None)
}
@@ -0,0 +1,63 @@
use super::super::super::model::build_admin_provider_models_payload;
use crate::control::GatewayControlDecision;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::admin_provider_id_for_models_list;
use crate::handlers::admin::shared::{query_param_optional_bool, query_param_value};
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) async fn maybe_handle(
state: &AppState,
request_context: &GatewayPublicRequestContext,
_request_body: Option<&Bytes>,
decision: &GatewayControlDecision,
) -> Result<Option<Response<Body>>, GatewayError> {
if decision.route_family.as_deref() == Some("provider_models_manage")
&& decision.route_kind.as_deref() == Some("list_provider_models")
&& request_context.request_method == http::Method::GET
&& request_context
.request_path
.starts_with("/api/admin/providers/")
&& request_context.request_path.ends_with("/models")
{
let Some(provider_id) = admin_provider_id_for_models_list(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let skip = query_param_value(request_context.request_query_string.as_deref(), "skip")
.and_then(|value| value.parse::<usize>().ok())
.unwrap_or(0);
let limit = query_param_value(request_context.request_query_string.as_deref(), "limit")
.and_then(|value| value.parse::<usize>().ok())
.filter(|value| *value > 0 && *value <= 500)
.unwrap_or(100);
let is_active =
query_param_optional_bool(request_context.request_query_string.as_deref(), "is_active");
return Ok(Some(
match build_admin_provider_models_payload(state, &provider_id, skip, limit, is_active)
.await
{
Some(payload) => Json(payload).into_response(),
None => (
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
)
.into_response(),
},
));
}
Ok(None)
}
@@ -0,0 +1,81 @@
use crate::control::GatewayControlDecision;
use crate::control::GatewayPublicRequestContext;
use crate::{AppState, GatewayError};
use axum::body::{Body, Bytes};
use axum::http::Response;
mod assign_global;
mod available_source;
mod batch;
mod create;
mod delete;
mod detail;
mod import;
mod list;
mod update;
pub(crate) async fn maybe_build_local_admin_provider_models_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.control_decision.as_ref() else {
return Ok(None);
};
if let Some(response) =
list::maybe_handle(state, request_context, request_body, decision).await?
{
return Ok(Some(response));
}
if let Some(response) =
detail::maybe_handle(state, request_context, request_body, decision).await?
{
return Ok(Some(response));
}
if let Some(response) =
create::maybe_handle(state, request_context, request_body, decision).await?
{
return Ok(Some(response));
}
if let Some(response) =
update::maybe_handle(state, request_context, request_body, decision).await?
{
return Ok(Some(response));
}
if let Some(response) =
delete::maybe_handle(state, request_context, request_body, decision).await?
{
return Ok(Some(response));
}
if let Some(response) =
batch::maybe_handle(state, request_context, request_body, decision).await?
{
return Ok(Some(response));
}
if let Some(response) =
available_source::maybe_handle(state, request_context, request_body, decision).await?
{
return Ok(Some(response));
}
if let Some(response) =
assign_global::maybe_handle(state, request_context, request_body, decision).await?
{
return Ok(Some(response));
}
if let Some(response) =
import::maybe_handle(state, request_context, request_body, decision).await?
{
return Ok(Some(response));
}
Ok(None)
}
@@ -0,0 +1,131 @@
use super::super::super::model::{
build_admin_provider_model_response, build_admin_provider_model_update_record,
};
use crate::control::GatewayControlDecision;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_model_route_parts, AdminProviderModelUpdateRequest,
};
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) async fn maybe_handle(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
decision: &GatewayControlDecision,
) -> Result<Option<Response<Body>>, GatewayError> {
if decision.route_family.as_deref() == Some("provider_models_manage")
&& decision.route_kind.as_deref() == Some("update_provider_model")
&& request_context.request_method == http::Method::PATCH
&& request_context.request_path.contains("/models/")
{
let Some((provider_id, model_id)) =
admin_provider_model_route_parts(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Model 不存在" })),
)
.into_response(),
));
};
let Some(existing) = state
.get_admin_provider_model(&provider_id, &model_id)
.await?
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Model {model_id} 不存在") })),
)
.into_response(),
));
};
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
let raw_value = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(value) => value,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let Some(raw_payload) = raw_value.as_object().cloned() else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
};
let payload = match serde_json::from_value::<AdminProviderModelUpdateRequest>(raw_value) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let record =
match build_admin_provider_model_update_record(state, &existing, &raw_payload, payload)
.await
{
Ok(record) => record,
Err(detail) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
));
}
};
return Ok(Some(
match state.update_admin_provider_model(&record).await? {
Some(updated) => {
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
Json(build_admin_provider_model_response(&updated, now_unix_secs))
.into_response()
}
None => (
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("Model {model_id} 不存在") })),
)
.into_response(),
},
));
}
Ok(None)
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,562 @@
use super::super::provider_oauth_quota::refresh_codex_provider_quota_locally;
use super::super::provider_oauth_refresh::{
build_internal_control_error_response, build_provider_oauth_auth_config_from_token_payload,
create_provider_oauth_catalog_key, find_duplicate_provider_oauth_key,
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
refresh_provider_oauth_account_state_after_update, update_existing_provider_oauth_catalog_key,
};
use super::super::provider_oauth_state::{
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
consume_provider_oauth_state, enrich_admin_provider_oauth_auth_config,
exchange_admin_provider_oauth_code, is_fixed_provider_type_for_provider_oauth,
json_non_empty_string, json_u64_value, parse_provider_oauth_callback_params,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_oauth_complete_key_id, admin_provider_oauth_complete_provider_id,
};
use crate::handlers::admin::shared::encrypt_catalog_secret_with_fallbacks;
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) async fn handle_admin_provider_oauth_complete_key(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
let Some(key_id) = admin_provider_oauth_complete_key_id(&request_context.request_path) else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
};
let Some(request_body) = request_body else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let raw_payload = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(serde_json::Value::Object(map)) => map,
_ => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
let callback_url = raw_payload
.get("callback_url")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.ok_or_else(|| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
)
});
let callback_url = match callback_url {
Ok(callback_url) => callback_url,
Err(response) => return Ok(response),
};
let params = parse_provider_oauth_callback_params(callback_url);
let Some(code) = params
.get("code")
.map(String::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
));
};
let Some(state_nonce) = params
.get("state")
.map(String::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
));
};
let state_data = match consume_provider_oauth_state(state, state_nonce).await {
Ok(Some(state_data)) => state_data,
Ok(None) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
));
}
};
if state_data.key_id != key_id {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next();
let Some(key) = key else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
};
if !key.auth_type.eq_ignore_ascii_case("oauth") {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Key 不是 oauth 认证类型",
));
}
if !state_data.provider_id.trim().is_empty() && state_data.provider_id != key.provider_id {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let provider_id = key.provider_id.clone();
let provider = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next();
let Some(provider) = provider else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不是固定类型,无法使用 provider-oauth",
));
}
if !state_data.provider_type.trim().is_empty()
&& !state_data
.provider_type
.eq_ignore_ascii_case(&provider_type)
{
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不支持 OAuth 授权",
));
};
let token_payload = match exchange_admin_provider_oauth_code(
state,
template,
code,
state_nonce,
state_data.pkce_verifier.as_deref(),
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(access_token) = json_non_empty_string(token_payload.get("access_token")) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
));
};
let refresh_token = json_non_empty_string(token_payload.get("refresh_token"));
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let expires_at = json_u64_value(token_payload.get("expires_in"))
.map(|expires_in| now_unix_secs.saturating_add(expires_in));
let mut auth_config = serde_json::Map::new();
auth_config.insert("provider_type".to_string(), json!(provider_type.clone()));
auth_config.insert("updated_at".to_string(), json!(now_unix_secs));
if let Some(token_type) = token_payload.get("token_type").cloned() {
auth_config.insert("token_type".to_string(), token_type);
}
if let Some(refresh_token) = refresh_token.as_ref() {
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
}
if let Some(expires_at) = expires_at {
auth_config.insert("expires_at".to_string(), json!(expires_at));
}
if let Some(scope) = token_payload.get("scope").cloned() {
auth_config.insert("scope".to_string(), scope);
}
enrich_admin_provider_oauth_auth_config(&provider_type, &mut auth_config, &token_payload);
let Some(encrypted_api_key) = encrypt_catalog_secret_with_fallbacks(state, &access_token)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth encryption unavailable",
));
};
let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let Some(encrypted_auth_config) =
encrypt_catalog_secret_with_fallbacks(state, &auth_config_json)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth encryption unavailable",
));
};
let updated = state
.update_provider_catalog_key_oauth_credentials(
&key_id,
&encrypted_api_key,
Some(&encrypted_auth_config),
expires_at,
)
.await?;
if !updated {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
}
let mut account_state_recheck_attempted = false;
let mut account_state_recheck_error = None::<String>;
if provider_type == "codex" {
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
if let Some(endpoint) = endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("openai:cli")
}) {
let refreshed_key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
.unwrap_or_else(|| key.clone());
if let Some(result) = refresh_codex_provider_quota_locally(
state,
&provider,
&endpoint,
vec![refreshed_key],
)
.await?
{
account_state_recheck_attempted = true;
let success = result
.get("success")
.and_then(serde_json::Value::as_u64)
.unwrap_or(0);
if success == 0 {
account_state_recheck_error = result
.get("results")
.and_then(serde_json::Value::as_array)
.and_then(|results| results.first())
.and_then(|value| value.get("message"))
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned);
}
}
}
}
Ok(Json(json!({
"provider_type": provider_type,
"expires_at": expires_at,
"has_refresh_token": refresh_token.is_some(),
"email": auth_config.get("email").cloned().unwrap_or(serde_json::Value::Null),
"account_state_recheck_attempted": account_state_recheck_attempted,
"account_state_recheck_error": account_state_recheck_error,
}))
.into_response())
}
pub(super) async fn handle_admin_provider_oauth_complete_provider(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
let Some(provider_id) =
admin_provider_oauth_complete_provider_id(&request_context.request_path)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let Some(request_body) = request_body else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let raw_payload = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(serde_json::Value::Object(map)) => map,
_ => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(callback_url) = raw_payload
.get("callback_url")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
));
};
let name = raw_payload
.get("name")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let proxy_node_id = raw_payload
.get("proxy_node_id")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let params = parse_provider_oauth_callback_params(callback_url);
let Some(code) = params
.get("code")
.map(String::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
));
};
let Some(state_nonce) = params
.get("state")
.map(String::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"callback_url 缺少 code/state",
));
};
let state_data = match consume_provider_oauth_state(state, state_nonce).await {
Ok(Some(state_data)) => state_data,
Ok(None) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
));
}
};
if !state_data.key_id.trim().is_empty() || state_data.provider_id != provider_id {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不是固定类型,无法使用 provider-oauth",
));
}
if provider_type == "kiro" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Kiro 不支持 OAuth 授权,请使用导入授权。",
));
}
if !state_data.provider_type.trim().is_empty()
&& !state_data
.provider_type
.eq_ignore_ascii_case(&provider_type)
{
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"state 无效或已过期",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不支持 OAuth 授权",
));
};
let token_payload = match exchange_admin_provider_oauth_code(
state,
template,
code,
state_nonce,
state_data.pkce_verifier.as_deref(),
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let (auth_config, access_token, refresh_token, expires_at) =
build_provider_oauth_auth_config_from_token_payload(&provider_type, &token_payload);
let Some(access_token) = access_token else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
));
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let api_formats = provider_oauth_active_api_formats(&endpoints);
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id.as_deref());
let duplicate =
match find_duplicate_provider_oauth_key(state, &provider_id, &auth_config, None).await {
Ok(duplicate) => duplicate,
Err(detail) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
detail,
));
}
};
let replaced = duplicate.is_some();
let persisted_key = if let Some(existing_key) = duplicate {
match update_existing_provider_oauth_catalog_key(
state,
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
} else {
let name = name
.or_else(|| {
auth_config
.get("email")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
.unwrap_or_else(|| {
format!(
"账号_{}",
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0)
)
});
match create_provider_oauth_catalog_key(
state,
&provider_id,
&name,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
};
let _ = refresh_provider_oauth_account_state_after_update(state, &provider, &persisted_key.id)
.await;
Ok(Json(json!({
"key_id": persisted_key.id,
"provider_type": provider_type,
"expires_at": expires_at,
"has_refresh_token": refresh_token.is_some(),
"email": auth_config.get("email").cloned().unwrap_or(serde_json::Value::Null),
"replaced": replaced,
}))
.into_response())
}
@@ -0,0 +1,550 @@
use super::super::provider_oauth_refresh::{
build_internal_control_error_response, build_provider_oauth_auth_config_from_token_payload,
create_provider_oauth_catalog_key, find_duplicate_provider_oauth_key,
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
refresh_provider_oauth_account_state_after_update, update_existing_provider_oauth_catalog_key,
};
use super::super::provider_oauth_state::{
build_admin_provider_oauth_backend_unavailable_response, build_kiro_device_key_name,
current_unix_secs, decode_jwt_claims, default_kiro_device_region,
default_kiro_device_start_url, generate_provider_oauth_nonce, json_non_empty_string,
json_u64_value, normalize_kiro_device_region, poll_admin_kiro_device_token,
read_provider_oauth_device_session, register_admin_kiro_device_oidc_client,
save_provider_oauth_device_session, start_admin_kiro_device_authorization,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_oauth_device_authorize_provider_id, admin_provider_oauth_device_poll_provider_id,
};
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::{AppState, GatewayError};
use aether_data::repository::provider_oauth::{
StoredAdminProviderOAuthDeviceSession, KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS,
};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde::Deserialize;
use serde_json::json;
#[derive(Debug, Deserialize)]
struct AdminProviderOAuthDeviceAuthorizePayload {
#[serde(default = "default_kiro_device_start_url")]
start_url: String,
#[serde(default = "default_kiro_device_region")]
region: String,
proxy_node_id: Option<String>,
}
#[derive(Debug, Deserialize)]
struct AdminProviderOAuthDevicePollPayload {
session_id: String,
}
pub(super) async fn handle_admin_provider_oauth_device_authorize(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) =
admin_provider_oauth_device_authorize_provider_id(&request_context.request_path)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let Some(request_body) = request_body else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let payload =
match serde_json::from_slice::<AdminProviderOAuthDeviceAuthorizePayload>(request_body) {
Ok(payload) => payload,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if provider_type != "kiro" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"设备授权仅支持 Kiro provider",
));
}
let region = normalize_kiro_device_region(Some(payload.region.as_str())).ok_or_else(|| {
build_internal_control_error_response(http::StatusCode::BAD_REQUEST, "region 格式无效")
});
let region = match region {
Ok(region) => region,
Err(response) => return Ok(response),
};
let start_url = payload.start_url.trim();
let start_url = if start_url.is_empty() {
default_kiro_device_start_url()
} else {
start_url.to_string()
};
let client_registration =
match register_admin_kiro_device_oidc_client(state, &region, &start_url).await {
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(client_id) = json_non_empty_string(client_registration.get("clientId")) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"注册 OIDC 客户端失败: unknown",
));
};
let Some(client_secret) = json_non_empty_string(client_registration.get("clientSecret")) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"注册 OIDC 客户端失败: unknown",
));
};
let device_authorization = match start_admin_kiro_device_authorization(
state,
&region,
&client_id,
&client_secret,
&start_url,
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let Some(device_code) = json_non_empty_string(
device_authorization
.get("deviceCode")
.or_else(|| device_authorization.get("device_code")),
) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"发起设备授权失败: unknown",
));
};
let user_code = json_non_empty_string(
device_authorization
.get("userCode")
.or_else(|| device_authorization.get("user_code")),
)
.unwrap_or_default();
let verification_uri = json_non_empty_string(
device_authorization
.get("verificationUri")
.or_else(|| device_authorization.get("verification_uri"))
.or_else(|| device_authorization.get("verificationUrl")),
)
.unwrap_or_default();
let verification_uri_complete = json_non_empty_string(
device_authorization
.get("verificationUriComplete")
.or_else(|| device_authorization.get("verification_uri_complete"))
.or_else(|| device_authorization.get("verificationUrlComplete")),
)
.unwrap_or_else(|| verification_uri.clone());
let expires_in = json_u64_value(
device_authorization
.get("expiresIn")
.or_else(|| device_authorization.get("expires_in")),
)
.unwrap_or(600);
let interval = json_u64_value(device_authorization.get("interval")).unwrap_or(5);
let now_unix_secs = current_unix_secs();
let session_id = generate_provider_oauth_nonce();
let session = StoredAdminProviderOAuthDeviceSession {
provider_id: provider_id.clone(),
region,
client_id,
client_secret,
device_code,
interval,
expires_at_unix_secs: now_unix_secs.saturating_add(expires_in),
status: "pending".to_string(),
proxy_node_id: payload
.proxy_node_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
created_at_unix_secs: now_unix_secs,
key_id: None,
email: None,
replaced: false,
error_msg: None,
};
if let Err(response) = save_provider_oauth_device_session(
state,
&session_id,
&session,
expires_in.saturating_add(KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS),
)
.await
{
return Ok(response);
}
Ok(Json(json!({
"session_id": session_id,
"user_code": user_code,
"verification_uri": verification_uri,
"verification_uri_complete": verification_uri_complete,
"expires_in": expires_in,
"interval": interval,
}))
.into_response())
}
pub(super) async fn handle_admin_provider_oauth_device_poll(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) =
admin_provider_oauth_device_poll_provider_id(&request_context.request_path)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let Some(request_body) = request_body else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let payload = match serde_json::from_slice::<AdminProviderOAuthDevicePollPayload>(request_body)
{
Ok(payload) => payload,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
let session_id = payload.session_id.trim();
if session_id.is_empty() {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"session_id 不能为空",
));
}
let Some(mut session) = read_provider_oauth_device_session(state, session_id).await? else {
return Ok(Json(json!({
"status": "expired",
"error": "会话不存在或已过期",
"replaced": false,
}))
.into_response());
};
if session.provider_id != provider_id {
return Ok(Json(json!({
"status": "error",
"error": "会话与 Provider 不匹配",
"replaced": false,
}))
.into_response());
}
if session.status == "authorized" {
return Ok(Json(json!({
"status": "authorized",
"key_id": session.key_id,
"email": session.email,
"replaced": session.replaced,
}))
.into_response());
}
if matches!(session.status.as_str(), "expired" | "error") {
return Ok(Json(json!({
"status": session.status,
"error": session.error_msg,
"replaced": session.replaced,
}))
.into_response());
}
if current_unix_secs() > session.expires_at_unix_secs {
session.status = "expired".to_string();
session.error_msg = Some("设备码已过期".to_string());
let _ = save_provider_oauth_device_session(state, session_id, &session, 30).await;
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"expired",
Json(json!({
"status": "expired",
"error": "设备码已过期",
"replaced": false,
}))
.into_response(),
));
}
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let token_result = match poll_admin_kiro_device_token(
state,
&session.region,
&session.client_id,
&session.client_secret,
&session.device_code,
)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
if token_result
.get("_error")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
{
let error_code = json_non_empty_string(token_result.get("error")).unwrap_or_default();
if error_code == "authorization_pending" {
return Ok(Json(json!({"status": "pending", "replaced": false})).into_response());
}
if error_code == "slow_down" {
return Ok(Json(json!({"status": "slow_down", "replaced": false})).into_response());
}
if error_code == "expired_token" {
session.status = "expired".to_string();
session.error_msg = Some("设备码已过期".to_string());
let _ = save_provider_oauth_device_session(state, session_id, &session, 30).await;
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"expired",
Json(json!({
"status": "expired",
"error": "设备码已过期",
"replaced": false,
}))
.into_response(),
));
}
if error_code == "access_denied" {
session.status = "error".to_string();
session.error_msg = Some("用户拒绝授权".to_string());
let _ = save_provider_oauth_device_session(state, session_id, &session, 30).await;
return Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"error",
Json(json!({
"status": "error",
"error": "用户拒绝授权",
"replaced": false,
}))
.into_response(),
));
}
let error_message = json_non_empty_string(token_result.get("error_description"))
.or_else(|| (!error_code.is_empty()).then_some(error_code.clone()))
.unwrap_or_else(|| "未知错误".to_string());
return Ok(Json(json!({
"status": "error",
"error": error_message,
"replaced": false,
}))
.into_response());
}
let Some(access_token) = json_non_empty_string(token_result.get("accessToken")) else {
return Ok(Json(json!({
"status": "error",
"error": "token 响应缺少 accessToken 或 refreshToken",
"replaced": false,
}))
.into_response());
};
let Some(refresh_token) = json_non_empty_string(token_result.get("refreshToken")) else {
return Ok(Json(json!({
"status": "error",
"error": "token 响应缺少 accessToken 或 refreshToken",
"replaced": false,
}))
.into_response());
};
let expires_at = json_u64_value(token_result.get("expiresIn"))
.map(|expires_in| current_unix_secs().saturating_add(expires_in))
.unwrap_or_else(|| current_unix_secs().saturating_add(3600));
let email = decode_jwt_claims(&access_token)
.and_then(|claims| claims.get("email").cloned())
.and_then(|value| value.as_str().map(ToOwned::to_owned));
let mut auth_config = serde_json::Map::new();
auth_config.insert("provider_type".to_string(), json!("kiro"));
auth_config.insert("auth_method".to_string(), json!("idc"));
auth_config.insert("refresh_token".to_string(), json!(refresh_token.clone()));
auth_config.insert("client_id".to_string(), json!(session.client_id.clone()));
auth_config.insert(
"client_secret".to_string(),
json!(session.client_secret.clone()),
);
auth_config.insert("region".to_string(), json!(session.region.clone()));
auth_config.insert("auth_region".to_string(), json!(session.region.clone()));
auth_config.insert("access_token".to_string(), json!(access_token.clone()));
auth_config.insert("expires_at".to_string(), json!(expires_at));
if let Some(email) = email.as_ref() {
auth_config.insert("email".to_string(), json!(email));
}
let duplicate =
match find_duplicate_provider_oauth_key(state, &provider_id, &auth_config, None).await {
Ok(duplicate) => duplicate,
Err(detail) => {
return Ok(Json(json!({
"status": "error",
"error": detail,
"replaced": false,
}))
.into_response());
}
};
let key_proxy = provider_oauth_key_proxy_value(session.proxy_node_id.as_deref());
let api_formats = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.filter(|endpoint| endpoint.is_active)
.map(|endpoint| endpoint.api_format)
.collect::<Vec<_>>();
let mut replaced = false;
let persisted_key = if let Some(existing_key) = duplicate {
replaced = true;
match update_existing_provider_oauth_catalog_key(
state,
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
Some(expires_at),
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
} else {
let key_name = build_kiro_device_key_name(email.as_deref(), Some(&refresh_token));
match create_provider_oauth_catalog_key(
state,
&provider_id,
&key_name,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
Some(expires_at),
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
};
let _ = refresh_provider_oauth_account_state_after_update(state, &provider, &persisted_key.id)
.await;
session.status = "authorized".to_string();
session.key_id = Some(persisted_key.id.clone());
session.email = email.clone();
session.replaced = replaced;
session.error_msg = None;
let _ = save_provider_oauth_device_session(state, session_id, &session, 60).await;
Ok(attach_admin_provider_oauth_device_poll_terminal_response(
session_id,
"authorized",
Json(json!({
"status": "authorized",
"key_id": persisted_key.id,
"email": email,
"replaced": replaced,
}))
.into_response(),
))
}
fn attach_admin_provider_oauth_device_poll_terminal_response(
session_id: &str,
status: &str,
response: Response<Body>,
) -> Response<Body> {
match status {
"authorized" => attach_admin_audit_response(
response,
"admin_provider_oauth_device_authorization_completed",
"poll_provider_oauth_device_authorization_terminal_state",
"provider_oauth_device_session",
session_id,
),
"expired" => attach_admin_audit_response(
response,
"admin_provider_oauth_device_authorization_expired",
"poll_provider_oauth_device_authorization_terminal_state",
"provider_oauth_device_session",
session_id,
),
"error" => attach_admin_audit_response(
response,
"admin_provider_oauth_device_authorization_failed",
"poll_provider_oauth_device_authorization_terminal_state",
"provider_oauth_device_session",
session_id,
),
_ => response,
}
}
@@ -0,0 +1,212 @@
use super::super::provider_oauth_refresh::{
build_internal_control_error_response, build_provider_oauth_auth_config_from_token_payload,
create_provider_oauth_catalog_key, find_duplicate_provider_oauth_key,
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
refresh_provider_oauth_account_state_after_update, update_existing_provider_oauth_catalog_key,
};
use super::super::provider_oauth_state::{
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
exchange_admin_provider_oauth_refresh_token, is_fixed_provider_type_for_provider_oauth,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::admin_provider_oauth_import_provider_id;
use crate::{AppState, GatewayError};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&axum::body::Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_provider_catalog_data_reader() {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
}
let Some(provider_id) = admin_provider_oauth_import_provider_id(&request_context.request_path)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let Some(request_body) = request_body else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
};
let raw_payload = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(serde_json::Value::Object(map)) => map,
_ => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"请求体必须是合法的 JSON 对象",
));
}
};
let Some(refresh_token_input) = raw_payload
.get("refresh_token")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Refresh Token 不能为空",
));
};
let name = raw_payload
.get("name")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let proxy_node_id = raw_payload
.get("proxy_node_id")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不是固定类型,无法使用 provider-oauth",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_admin_provider_oauth_backend_unavailable_response());
};
let token_payload =
match exchange_admin_provider_oauth_refresh_token(state, template, refresh_token_input)
.await
{
Ok(payload) => payload,
Err(response) => return Ok(response),
};
let (mut auth_config, access_token, returned_refresh_token, expires_at) =
build_provider_oauth_auth_config_from_token_payload(&provider_type, &token_payload);
let Some(access_token) = access_token else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token refresh 返回缺少 access_token",
));
};
let refresh_token = returned_refresh_token
.or_else(|| Some(refresh_token_input.to_string()))
.filter(|value| !value.trim().is_empty());
if let Some(refresh_token) = refresh_token.as_ref() {
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
}
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let api_formats = provider_oauth_active_api_formats(&endpoints);
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id.as_deref());
let duplicate =
match find_duplicate_provider_oauth_key(state, &provider_id, &auth_config, None).await {
Ok(duplicate) => duplicate,
Err(detail) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
detail,
));
}
};
let replaced = duplicate.is_some();
let persisted_key = if let Some(existing_key) = duplicate {
match update_existing_provider_oauth_catalog_key(
state,
&existing_key,
&access_token,
&auth_config,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
} else {
let name = name
.or_else(|| {
auth_config
.get("email")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
.unwrap_or_else(|| {
format!(
"账号_{}",
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0)
)
});
match create_provider_oauth_catalog_key(
state,
&provider_id,
&name,
&access_token,
&auth_config,
&api_formats,
key_proxy.clone(),
expires_at,
)
.await?
{
Some(key) => key,
None => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth write unavailable",
));
}
}
};
let _ = refresh_provider_oauth_account_state_after_update(state, &provider, &persisted_key.id)
.await;
Ok(Json(json!({
"key_id": persisted_key.id,
"provider_type": provider_type,
"expires_at": expires_at,
"has_refresh_token": refresh_token.is_some(),
"email": auth_config.get("email").cloned().unwrap_or(serde_json::Value::Null),
"replaced": replaced,
}))
.into_response())
}
@@ -0,0 +1,222 @@
use super::provider_oauth_state::{
build_admin_provider_oauth_backend_unavailable_response,
build_admin_provider_oauth_supported_types_payload,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_oauth_batch_import_provider_id,
admin_provider_oauth_batch_import_task_provider_id, admin_provider_oauth_complete_key_id,
admin_provider_oauth_complete_provider_id, admin_provider_oauth_device_authorize_provider_id,
admin_provider_oauth_import_provider_id, admin_provider_oauth_refresh_key_id,
admin_provider_oauth_start_key_id, admin_provider_oauth_start_provider_id,
};
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
mod batch;
mod complete;
mod device;
mod import;
mod refresh;
mod start;
mod tasks;
pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.control_decision.as_ref() else {
return Ok(None);
};
if decision.route_family.as_deref() != Some("provider_oauth_manage") {
return Ok(None);
}
let route_kind = decision.route_kind.as_deref();
let method = &request_context.request_method;
if route_kind == Some("supported_types")
&& *method == http::Method::GET
&& request_context.request_path == "/api/admin/provider-oauth/supported-types"
{
return Ok(Some(
Json(build_admin_provider_oauth_supported_types_payload()).into_response(),
));
}
if route_kind == Some("start_key_oauth") && *method == http::Method::POST {
let response = start::handle_admin_provider_oauth_start_key(state, request_context).await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_authorization_started",
"start_provider_oauth_for_key",
"provider_key",
admin_provider_oauth_start_key_id(&request_context.request_path),
)));
}
if route_kind == Some("start_provider_oauth") && *method == http::Method::POST {
let response =
start::handle_admin_provider_oauth_start_provider(state, request_context).await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_authorization_started",
"start_provider_oauth_for_provider",
"provider",
admin_provider_oauth_start_provider_id(&request_context.request_path),
)));
}
if route_kind == Some("get_batch_import_task_status") && *method == http::Method::GET {
return Ok(Some(
tasks::handle_admin_provider_oauth_batch_import_task_status(state, request_context)
.await?,
));
}
if route_kind == Some("complete_key_oauth") && *method == http::Method::POST {
let response = complete::handle_admin_provider_oauth_complete_key(
state,
request_context,
request_body,
)
.await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_completed",
"complete_provider_oauth_for_key",
"provider_key",
admin_provider_oauth_complete_key_id(&request_context.request_path),
)));
}
if route_kind == Some("refresh_key_oauth") && *method == http::Method::POST {
let response =
refresh::handle_admin_provider_oauth_refresh_key(state, request_context).await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_refreshed",
"refresh_provider_oauth_for_key",
"provider_key",
admin_provider_oauth_refresh_key_id(&request_context.request_path),
)));
}
if route_kind == Some("complete_provider_oauth") && *method == http::Method::POST {
let response = complete::handle_admin_provider_oauth_complete_provider(
state,
request_context,
request_body,
)
.await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_completed",
"complete_provider_oauth_for_provider",
"provider",
admin_provider_oauth_complete_provider_id(&request_context.request_path),
)));
}
if route_kind == Some("import_refresh_token") && *method == http::Method::POST {
let response = import::handle_admin_provider_oauth_import_refresh_token(
state,
request_context,
request_body,
)
.await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_refresh_token_imported",
"import_provider_oauth_refresh_token",
"provider",
admin_provider_oauth_import_provider_id(&request_context.request_path),
)));
}
if route_kind == Some("batch_import_oauth") && *method == http::Method::POST {
let response =
batch::handle_admin_provider_oauth_batch_import(state, request_context, request_body)
.await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_batch_import_completed",
"batch_import_provider_oauth",
"provider",
admin_provider_oauth_batch_import_provider_id(&request_context.request_path),
)));
}
if route_kind == Some("start_batch_import_oauth_task") && *method == http::Method::POST {
let response = batch::handle_admin_provider_oauth_start_batch_import_task(
state,
request_context,
request_body,
)
.await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_batch_import_started",
"start_provider_oauth_batch_import",
"provider",
admin_provider_oauth_batch_import_task_provider_id(&request_context.request_path),
)));
}
if route_kind == Some("device_authorize") && *method == http::Method::POST {
let response = device::handle_admin_provider_oauth_device_authorize(
state,
request_context,
request_body,
)
.await?;
return Ok(Some(attach_admin_provider_oauth_audit_response(
response,
"admin_provider_oauth_device_authorization_started",
"start_provider_oauth_device_authorization",
"provider",
admin_provider_oauth_device_authorize_provider_id(&request_context.request_path),
)));
}
if route_kind == Some("device_poll") && *method == http::Method::POST {
return Ok(Some(
device::handle_admin_provider_oauth_device_poll(state, request_context, request_body)
.await?,
));
}
if matches!(
route_kind,
Some("refresh_key_oauth" | "import_refresh_token")
) {
return Ok(Some(
build_admin_provider_oauth_backend_unavailable_response(),
));
}
Ok(None)
}
fn attach_admin_provider_oauth_audit_response(
response: Response<Body>,
event_name: &'static str,
action: &'static str,
target_type: &'static str,
target_id: Option<String>,
) -> Response<Body> {
if !response.status().is_success() {
return response;
}
let Some(target_id) = target_id else {
return response;
};
attach_admin_audit_response(response, event_name, action, target_type, &target_id)
}
@@ -0,0 +1,226 @@
use super::super::provider_oauth_quota::persist_provider_quota_refresh_state;
use super::super::provider_oauth_refresh::{
build_internal_control_error_response, merge_provider_oauth_refresh_failure_reason,
normalize_provider_oauth_refresh_error_message, provider_oauth_runtime_endpoint_for_provider,
refresh_provider_oauth_account_state_after_update,
};
use super::super::provider_oauth_state::is_fixed_provider_type_for_provider_oauth;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_oauth_refresh_key_id, OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
};
use crate::handlers::admin::shared::decrypt_catalog_secret_with_fallbacks;
use crate::{AppState, GatewayError};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(super) async fn handle_admin_provider_oauth_refresh_key(
state: &AppState,
request_context: &GatewayPublicRequestContext,
) -> Result<Response<Body>, GatewayError> {
let Some(key_id) = admin_provider_oauth_refresh_key_id(&request_context.request_path) else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
};
let Some(key) = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
};
if !key.auth_type.eq_ignore_ascii_case("oauth") {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Key 不是 oauth 认证类型",
));
}
let Some(encrypted_auth_config) = key.encrypted_auth_config.as_deref() else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"缺少 auth_config,无法 refresh",
));
};
let Some(decrypted_auth_config) =
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), encrypted_auth_config)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth encryption unavailable",
));
};
let parsed_auth_config = serde_json::from_str::<serde_json::Value>(&decrypted_auth_config)
.ok()
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
let has_refresh_token = parsed_auth_config
.get("refresh_token")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty());
if !has_refresh_token {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"缺少 refresh_token,需要重新授权",
));
}
let provider_id = key.provider_id.clone();
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不是固定类型,无法使用 provider-oauth",
));
}
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
let Some(endpoint) = provider_oauth_runtime_endpoint_for_provider(&provider_type, endpoints)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"找不到有效端点,无法 refresh",
));
};
let Some(transport) = state
.read_provider_transport_snapshot(&provider_id, &endpoint.id, &key_id)
.await?
else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Provider transport snapshot unavailable",
));
};
match state.force_local_oauth_refresh_entry(&transport).await {
Ok(Some(_)) => {}
Ok(None) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"缺少 refresh_token,需要重新授权",
));
}
Err(crate::provider_transport::LocalOAuthRefreshError::HttpStatus {
status_code,
body_excerpt,
..
}) => {
let error_reason = normalize_provider_oauth_refresh_error_message(
Some(status_code),
Some(body_excerpt.as_str()),
);
if matches!(status_code, 400 | 401 | 403) {
let merged_reason = merge_provider_oauth_refresh_failure_reason(
key.oauth_invalid_reason.as_deref(),
format!(
"{OAUTH_REFRESH_FAILED_PREFIX}Token 续期失败 ({status_code}): {error_reason}"
)
.as_str(),
);
if let Some(merged_reason) = merged_reason {
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let _ = persist_provider_quota_refresh_state(
state,
&key_id,
None,
Some(now_unix_secs),
Some(merged_reason),
None,
)
.await?;
}
}
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("Token 刷新失败:{error_reason}"),
));
}
Err(crate::provider_transport::LocalOAuthRefreshError::Transport { source, .. }) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
format!("Token 刷新失败:{}", source),
));
}
Err(crate::provider_transport::LocalOAuthRefreshError::InvalidResponse {
message, ..
}) => {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("Token 刷新失败:{message}"),
));
}
}
if !key
.oauth_invalid_reason
.as_deref()
.map(str::trim)
.is_some_and(|value| value.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX))
{
let _ = state
.clear_provider_catalog_key_oauth_invalid_marker(&key_id)
.await?;
}
let refreshed_key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next()
.unwrap_or(key);
let refreshed_auth_config = refreshed_key
.encrypted_auth_config
.as_deref()
.and_then(|ciphertext| {
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)
})
.and_then(|plaintext| serde_json::from_str::<serde_json::Value>(&plaintext).ok())
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
let (account_state_recheck_attempted, account_state_recheck_error) =
refresh_provider_oauth_account_state_after_update(state, &provider, &key_id).await?;
Ok(Json(json!({
"provider_type": provider_type,
"expires_at": refreshed_auth_config.get("expires_at").cloned().unwrap_or(serde_json::Value::Null),
"has_refresh_token": refreshed_auth_config
.get("refresh_token")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty()),
"email": refreshed_auth_config.get("email").cloned().unwrap_or(serde_json::Value::Null),
"account_state_recheck_attempted": account_state_recheck_attempted,
"account_state_recheck_error": account_state_recheck_error,
}))
.into_response())
}
@@ -0,0 +1,173 @@
use super::super::provider_oauth_refresh::build_internal_control_error_response;
use super::super::provider_oauth_state::{
admin_provider_oauth_template, build_provider_oauth_start_response,
generate_provider_oauth_pkce_verifier, is_fixed_provider_type_for_provider_oauth,
provider_oauth_pkce_s256, save_provider_oauth_state,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_oauth_start_key_id, admin_provider_oauth_start_provider_id,
};
use crate::{AppState, GatewayError};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
pub(super) async fn handle_admin_provider_oauth_start_key(
state: &AppState,
request_context: &GatewayPublicRequestContext,
) -> Result<Response<Body>, GatewayError> {
let Some(key_id) = admin_provider_oauth_start_key_id(&request_context.request_path) else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
};
let key = state
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
.await?
.into_iter()
.next();
let Some(key) = key else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Key 不存在",
));
};
if !key.auth_type.eq_ignore_ascii_case("oauth") {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Key 不是 oauth 认证类型",
));
}
let provider_id = key.provider_id.clone();
let provider = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next();
let Some(provider) = provider else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不是固定类型,无法使用 provider-oauth",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不支持 OAuth 授权",
));
};
let pkce_verifier = template
.use_pkce
.then(generate_provider_oauth_pkce_verifier);
let code_challenge = pkce_verifier.as_deref().map(provider_oauth_pkce_s256);
let nonce = match save_provider_oauth_state(
state,
&key_id,
&provider_id,
&provider_type,
pkce_verifier.as_deref(),
)
.await
{
Ok(nonce) => nonce,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
));
}
};
Ok(Json(build_provider_oauth_start_response(
template,
&nonce,
code_challenge.as_deref(),
))
.into_response())
}
pub(super) async fn handle_admin_provider_oauth_start_provider(
state: &AppState,
request_context: &GatewayPublicRequestContext,
) -> Result<Response<Body>, GatewayError> {
let Some(provider_id) = admin_provider_oauth_start_provider_id(&request_context.request_path)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next();
let Some(provider) = provider else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !is_fixed_provider_type_for_provider_oauth(&provider_type) {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不是固定类型,无法使用 provider-oauth",
));
}
if provider_type == "kiro" {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Kiro 不支持 OAuth 授权,请使用导入授权。",
));
}
let Some(template) = admin_provider_oauth_template(&provider_type) else {
return Ok(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"该 Provider 不支持 OAuth 授权",
));
};
let pkce_verifier = template
.use_pkce
.then(generate_provider_oauth_pkce_verifier);
let code_challenge = pkce_verifier.as_deref().map(provider_oauth_pkce_s256);
let nonce = match save_provider_oauth_state(
state,
"",
&provider_id,
&provider_type,
pkce_verifier.as_deref(),
)
.await
{
Ok(nonce) => nonce,
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
));
}
};
Ok(Json(build_provider_oauth_start_response(
template,
&nonce,
code_challenge.as_deref(),
))
.into_response())
}
@@ -0,0 +1,65 @@
use super::super::provider_oauth_refresh::build_internal_control_error_response;
use super::super::provider_oauth_state::read_provider_oauth_batch_task_payload;
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::admin_provider_oauth_batch_import_task_path;
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::{AppState, GatewayError};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
pub(super) async fn handle_admin_provider_oauth_batch_import_task_status(
state: &AppState,
request_context: &GatewayPublicRequestContext,
) -> Result<Response<Body>, GatewayError> {
let Some((provider_id, task_id)) =
admin_provider_oauth_batch_import_task_path(&request_context.request_path)
else {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"批量导入任务不存在",
));
};
let payload = match read_provider_oauth_batch_task_payload(state, &provider_id, &task_id).await
{
Ok(Some(payload)) => payload,
Ok(None) => {
return Ok(build_internal_control_error_response(
http::StatusCode::NOT_FOUND,
"批量导入任务不存在或已过期",
));
}
Err(_) => {
return Ok(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth batch task redis unavailable",
));
}
};
let status = payload
.get("status")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned)
.unwrap_or_default();
let response = Json(payload).into_response();
Ok(match status.as_str() {
"completed" => attach_admin_audit_response(
response,
"admin_provider_oauth_batch_task_completed_viewed",
"view_provider_oauth_batch_task_terminal_state",
"provider_oauth_batch_task",
&format!("{provider_id}:{task_id}"),
),
"failed" => attach_admin_audit_response(
response,
"admin_provider_oauth_batch_task_failed_viewed",
"view_provider_oauth_batch_task_terminal_state",
"provider_oauth_batch_task",
&format!("{provider_id}:{task_id}"),
),
_ => response,
})
}
@@ -0,0 +1,16 @@
mod dispatch;
pub(crate) mod quota;
pub(crate) mod refresh;
pub(crate) mod state;
pub(crate) use self::dispatch::maybe_build_local_admin_provider_oauth_response;
pub(crate) use self::quota as provider_oauth_quota;
pub(crate) use self::quota::{
normalize_string_id_list, refresh_antigravity_provider_quota_locally,
refresh_codex_provider_quota_locally, refresh_kiro_provider_quota_locally,
};
pub(crate) use self::refresh as provider_oauth_refresh;
pub(crate) use self::refresh::{
build_internal_control_error_response, normalize_provider_oauth_refresh_error_message,
};
pub(crate) use self::state as provider_oauth_state;
@@ -0,0 +1,340 @@
use super::{
coerce_json_f64, coerce_json_string, execute_provider_quota_plan,
extract_execution_error_message, persist_provider_quota_refresh_state,
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::provider::shared::ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH;
use crate::{AppState, GatewayError};
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
fn parse_antigravity_usage_response(
value: &serde_json::Value,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
let models = value.get("models")?.as_object()?;
let mut quota_by_model = serde_json::Map::new();
for (model_id, model_value) in models {
let mut payload = serde_json::Map::new();
if let Some(display_name) = coerce_json_string(
model_value
.get("displayName")
.or_else(|| model_value.get("display_name")),
) {
payload.insert("display_name".to_string(), json!(display_name));
}
let quota_info = model_value
.get("quotaInfo")
.and_then(serde_json::Value::as_object);
let remaining_fraction = quota_info
.and_then(|object| object.get("remainingFraction"))
.and_then(coerce_json_f64);
let used_percent = remaining_fraction
.map(|value| ((1.0 - value).max(0.0) * 100.0).min(100.0))
.unwrap_or(100.0);
payload.insert(
"remaining_fraction".to_string(),
json!(remaining_fraction.unwrap_or(0.0)),
);
payload.insert("used_percent".to_string(), json!(used_percent));
if let Some(reset_time) = quota_info
.and_then(|object| object.get("resetTime"))
.cloned()
.filter(|value| !value.is_null())
{
payload.insert("reset_time".to_string(), reset_time);
}
quota_by_model.insert(model_id.clone(), serde_json::Value::Object(payload));
}
Some(json!({
"updated_at": updated_at_unix_secs,
"is_forbidden": false,
"forbidden_reason": serde_json::Value::Null,
"forbidden_at": serde_json::Value::Null,
"models": quota_by_model,
}))
}
async fn execute_antigravity_quota_plan(
state: &AppState,
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
authorization: (String, String),
project_id: &str,
auth: &crate::provider_transport::antigravity::AntigravityRequestAuthSupport,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
let supported_auth = match auth {
crate::provider_transport::antigravity::AntigravityRequestAuthSupport::Supported(auth) => {
auth
}
crate::provider_transport::antigravity::AntigravityRequestAuthSupport::Unsupported(_) => {
return Ok(ProviderQuotaExecutionOutcome::Failure(
"缺少 OAuth 认证信息,请先授权/刷新 Token".to_string(),
));
}
};
let mut headers =
crate::provider_transport::antigravity::build_antigravity_static_identity_headers(
supported_auth,
);
headers.insert("authorization".to_string(), authorization.1);
headers.insert("content-type".to_string(), "application/json".to_string());
headers.insert("accept".to_string(), "application/json".to_string());
headers
.entry("user-agent".to_string())
.or_insert_with(|| "antigravity".to_string());
let body = json!({ "project": project_id });
let plan = ExecutionPlan {
request_id: format!("antigravity-quota:{}", transport.key.id),
candidate_id: None,
provider_name: Some("antigravity".to_string()),
provider_id: transport.provider.id.clone(),
endpoint_id: transport.endpoint.id.clone(),
key_id: transport.key.id.clone(),
method: "POST".to_string(),
url: format!(
"{}{}",
transport.endpoint.base_url.trim_end_matches('/'),
ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH
),
headers,
content_type: Some("application/json".to_string()),
content_encoding: None,
body: RequestBody {
json_body: Some(body),
body_bytes_b64: None,
body_ref: None,
},
stream: false,
client_api_format: "gemini:chat".to_string(),
provider_api_format: "antigravity:fetch_available_models".to_string(),
model_name: Some("fetchAvailableModels".to_string()),
proxy: crate::provider_transport::resolve_transport_proxy_snapshot_with_tunnel_affinity(
state, transport,
)
.await,
tls_profile: crate::provider_transport::resolve_transport_tls_profile(transport),
timeouts: crate::provider_transport::resolve_transport_execution_timeouts(transport).or(
Some(ExecutionTimeouts {
connect_ms: Some(30_000),
read_ms: Some(30_000),
write_ms: Some(30_000),
pool_ms: Some(30_000),
total_ms: Some(30_000),
..ExecutionTimeouts::default()
}),
),
};
execute_provider_quota_plan(state, transport, plan, "antigravity").await
}
pub(crate) async fn refresh_antigravity_provider_quota_locally(
state: &AppState,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
) -> Result<Option<serde_json::Value>, GatewayError> {
let mut results = Vec::new();
let mut success_count = 0usize;
let mut failed_count = 0usize;
for key in keys {
let transport = match state
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
.await?
{
Some(transport) => transport,
None => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Provider transport snapshot unavailable",
}));
continue;
}
};
let authorization = match state.resolve_local_oauth_request_auth(&transport).await? {
Some(crate::provider_transport::LocalResolvedOAuthRequestAuth::Header {
name,
value,
}) => (name, value),
_ => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 OAuth 认证信息,请先授权/刷新 Token",
}));
continue;
}
};
let antigravity_auth =
crate::provider_transport::antigravity::resolve_local_antigravity_request_auth(
&transport,
);
let project_id = match &antigravity_auth {
crate::provider_transport::antigravity::AntigravityRequestAuthSupport::Supported(
auth,
) => auth.project_id.clone(),
crate::provider_transport::antigravity::AntigravityRequestAuthSupport::Unsupported(
_,
) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 OAuth 认证信息,请先授权/刷新 Token",
}));
continue;
}
};
let result = match execute_antigravity_quota_plan(
state,
&transport,
authorization,
&project_id,
&antigravity_auth,
)
.await?
{
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": format!("fetchAvailableModels 请求执行失败: {detail}"),
"status_code": 502,
}));
continue;
}
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut metadata_update = None::<serde_json::Value>;
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
let mut status = "error".to_string();
let mut message = None::<String>;
if result.status_code == 200 {
if let Some(body_json) = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
{
metadata_update = parse_antigravity_usage_response(body_json, now_unix_secs)
.map(|metadata| json!({ "antigravity": metadata }));
if metadata_update.is_some() {
status = "success".to_string();
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含配额信息".to_string());
}
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含配额信息".to_string());
}
} else {
let err_msg = extract_execution_error_message(&result);
message = Some(match err_msg.as_deref() {
Some(detail) if !detail.is_empty() => {
format!(
"fetchAvailableModels 返回状态码 {}: {}",
result.status_code, detail
)
}
_ => format!("fetchAvailableModels 返回状态码 {}", result.status_code),
});
if result.status_code == 403 {
let reason = err_msg
.clone()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| "账户访问被禁止".to_string());
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason = Some(format!("账户访问被禁止: {reason}"));
metadata_update = Some(json!({
"antigravity": {
"is_forbidden": true,
"forbidden_reason": reason,
"forbidden_at": now_unix_secs,
"updated_at": now_unix_secs,
}
}));
status = "forbidden".to_string();
}
}
if !persist_provider_quota_refresh_state(
state,
&key.id,
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason,
None,
)
.await?
{
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Key 状态写入失败",
}));
continue;
}
if status == "success" {
success_count += 1;
} else {
failed_count += 1;
}
let mut payload = serde_json::Map::new();
payload.insert("key_id".to_string(), json!(key.id));
payload.insert("key_name".to_string(), json!(key.name));
payload.insert("status".to_string(), json!(status));
if let Some(message) = message {
payload.insert("message".to_string(), json!(message));
}
if let Some(metadata) = metadata_update
.as_ref()
.and_then(|value| value.get("antigravity"))
.cloned()
{
payload.insert("metadata".to_string(), metadata);
}
results.push(serde_json::Value::Object(payload));
}
Ok(Some(json!({
"success": success_count,
"failed": failed_count,
"total": success_count + failed_count,
"results": results,
"message": format!("已处理 {} 个 Key", success_count + failed_count),
"auto_removed": 0,
})))
}
@@ -0,0 +1,696 @@
use super::{
coerce_json_bool, coerce_json_f64, coerce_json_string, coerce_json_u64,
execute_provider_quota_plan, extract_execution_error_message, normalize_string_id_list,
persist_provider_quota_refresh_state, provider_auto_remove_banned_keys,
quota_refresh_success_invalid_state, should_auto_remove_structured_reason,
ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::provider::shared::{
CODEX_WHAM_USAGE_URL, OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX,
OAUTH_REQUEST_FAILED_PREFIX,
};
use crate::{AppState, GatewayError};
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use serde_json::json;
use std::collections::{BTreeMap, BTreeSet};
use std::time::{SystemTime, UNIX_EPOCH};
fn normalize_codex_plan_type(value: Option<&str>) -> Option<String> {
value
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase())
}
fn build_codex_quota_exhausted_fallback_metadata(
plan_type: Option<&str>,
updated_at_unix_secs: u64,
) -> serde_json::Value {
let mut object = serde_json::Map::new();
if let Some(plan_type) = normalize_codex_plan_type(plan_type) {
object.insert(
"plan_type".to_string(),
serde_json::Value::String(plan_type),
);
}
object.insert("updated_at".to_string(), json!(updated_at_unix_secs));
object.insert("primary_used_percent".to_string(), json!(100.0));
if normalize_codex_plan_type(plan_type) != Some("free".to_string()) {
object.insert("secondary_used_percent".to_string(), json!(100.0));
}
serde_json::Value::Object(object)
}
fn codex_write_window(
target: &mut serde_json::Map<String, serde_json::Value>,
source: &serde_json::Map<String, serde_json::Value>,
target_prefix: &str,
) {
if let Some(value) = source.get("used_percent").and_then(coerce_json_f64) {
target.insert(format!("{target_prefix}_used_percent"), json!(value));
}
if let Some(value) = source.get("reset_after_seconds").and_then(coerce_json_u64) {
target.insert(format!("{target_prefix}_reset_after_seconds"), json!(value));
}
if let Some(value) = source.get("reset_at").and_then(coerce_json_u64) {
target.insert(format!("{target_prefix}_reset_at"), json!(value));
}
if let Some(value) = source.get("window_minutes").and_then(coerce_json_u64) {
target.insert(format!("{target_prefix}_window_minutes"), json!(value));
}
}
fn parse_codex_wham_usage_response(
value: &serde_json::Value,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
let root = value.as_object()?;
if root.is_empty() {
return None;
}
let mut result = serde_json::Map::new();
let plan_type =
normalize_codex_plan_type(root.get("plan_type").and_then(serde_json::Value::as_str));
if let Some(plan_type) = plan_type.as_ref() {
result.insert("plan_type".to_string(), json!(plan_type));
}
let rate_limit = root
.get("rate_limit")
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
let primary_window = rate_limit
.get("primary_window")
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
let secondary_window = rate_limit
.get("secondary_window")
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
let use_paid_windows = !secondary_window.is_empty() && plan_type.as_deref() != Some("free");
if use_paid_windows {
codex_write_window(&mut result, &secondary_window, "primary");
codex_write_window(&mut result, &primary_window, "secondary");
} else {
codex_write_window(&mut result, &primary_window, "primary");
}
if let Some(credits) = root.get("credits").and_then(serde_json::Value::as_object) {
if let Some(value) = credits.get("has_credits").and_then(coerce_json_bool) {
result.insert("has_credits".to_string(), json!(value));
}
if let Some(value) = credits.get("balance").and_then(coerce_json_f64) {
result.insert("credits_balance".to_string(), json!(value));
}
if let Some(value) = credits.get("unlimited").and_then(coerce_json_bool) {
result.insert("credits_unlimited".to_string(), json!(value));
}
}
if result.is_empty() {
return None;
}
result.insert("updated_at".to_string(), json!(updated_at_unix_secs));
Some(serde_json::Value::Object(result))
}
fn parse_codex_usage_headers(
headers: &BTreeMap<String, String>,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
let mut result = serde_json::Map::new();
let normalized = headers
.iter()
.map(|(key, value)| (key.trim().to_ascii_lowercase(), value.trim().to_string()))
.collect::<BTreeMap<_, _>>();
if !normalized.keys().any(|key| key.starts_with("x-codex-")) {
return None;
}
let plan_type =
normalize_codex_plan_type(normalized.get("x-codex-plan-type").map(String::as_str));
if let Some(plan_type) = plan_type.as_ref() {
result.insert("plan_type".to_string(), json!(plan_type));
}
let read_window = |prefix: &str| -> serde_json::Map<String, serde_json::Value> {
let mut object = serde_json::Map::new();
let used_key = format!("x-codex-{prefix}-used-percent");
let reset_after_key = format!("x-codex-{prefix}-reset-after-seconds");
let reset_at_key = format!("x-codex-{prefix}-reset-at");
let window_minutes_key = format!("x-codex-{prefix}-window-minutes");
if let Some(value) = normalized
.get(&used_key)
.and_then(|value| value.parse::<f64>().ok())
{
object.insert("used_percent".to_string(), json!(value));
}
if let Some(value) = normalized
.get(&reset_after_key)
.and_then(|value| value.parse::<u64>().ok())
{
object.insert("reset_after_seconds".to_string(), json!(value));
}
if let Some(value) = normalized
.get(&reset_at_key)
.and_then(|value| value.parse::<u64>().ok())
{
object.insert("reset_at".to_string(), json!(value));
}
if let Some(value) = normalized
.get(&window_minutes_key)
.and_then(|value| value.parse::<u64>().ok())
{
object.insert("window_minutes".to_string(), json!(value));
}
object
};
let primary_window = read_window("primary");
let secondary_window = read_window("secondary");
let use_paid_windows = !secondary_window.is_empty() && plan_type.as_deref() != Some("free");
if use_paid_windows {
codex_write_window(&mut result, &secondary_window, "primary");
codex_write_window(&mut result, &primary_window, "secondary");
} else {
codex_write_window(&mut result, &primary_window, "primary");
}
if let Some(value) = normalized
.get("x-codex-primary-over-secondary-limit-percent")
.and_then(|value| value.parse::<f64>().ok())
{
result.insert(
"primary_over_secondary_limit_percent".to_string(),
json!(value),
);
}
if let Some(value) = normalized
.get("x-codex-credits-has-credits")
.and_then(|value| match value.to_ascii_lowercase().as_str() {
"true" | "1" => Some(true),
"false" | "0" => Some(false),
_ => None,
})
{
result.insert("has_credits".to_string(), json!(value));
}
if let Some(value) = normalized
.get("x-codex-credits-balance")
.and_then(|value| value.parse::<f64>().ok())
{
result.insert("credits_balance".to_string(), json!(value));
}
if let Some(value) = normalized
.get("x-codex-credits-unlimited")
.and_then(|value| match value.to_ascii_lowercase().as_str() {
"true" | "1" => Some(true),
"false" | "0" => Some(false),
_ => None,
})
{
result.insert("credits_unlimited".to_string(), json!(value));
}
if result.is_empty() {
return None;
}
result.insert("updated_at".to_string(), json!(updated_at_unix_secs));
Some(serde_json::Value::Object(result))
}
fn codex_current_invalid_reason(key: &StoredProviderCatalogKey) -> String {
key.oauth_invalid_reason
.as_deref()
.map(str::trim)
.unwrap_or_default()
.to_string()
}
fn codex_merge_invalid_reason(current: &str, candidate_reason: &str) -> String {
if current.is_empty() {
return candidate_reason.to_string();
}
if current.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX) {
return current.to_string();
}
if current.starts_with(OAUTH_EXPIRED_PREFIX)
&& candidate_reason.starts_with(OAUTH_REQUEST_FAILED_PREFIX)
{
return current.to_string();
}
candidate_reason.to_string()
}
fn codex_build_invalid_state(
key: &StoredProviderCatalogKey,
candidate_reason: String,
now_unix_secs: u64,
) -> (Option<u64>, Option<String>) {
let current_reason = codex_current_invalid_reason(key);
let merged_reason = codex_merge_invalid_reason(&current_reason, &candidate_reason);
if merged_reason == current_reason {
return (key.oauth_invalid_at_unix_secs, Some(merged_reason));
}
(Some(now_unix_secs), Some(merged_reason))
}
fn codex_looks_like_token_invalidated(message: Option<&str>) -> bool {
let lowered = message.unwrap_or_default().trim().to_ascii_lowercase();
lowered.contains("token invalid")
|| lowered.contains("token invalidated")
|| lowered.contains("session has expired")
|| lowered.contains("session expired")
}
fn codex_looks_like_account_deactivated(message: Option<&str>) -> bool {
let lowered = message.unwrap_or_default().trim().to_ascii_lowercase();
lowered.contains("account has been deactivated") || lowered.contains("account deactivated")
}
fn codex_looks_like_workspace_deactivated(message: Option<&str>) -> bool {
let lowered = message.unwrap_or_default().trim().to_ascii_lowercase();
lowered.contains("deactivated_workspace")
|| (lowered.contains("workspace") && lowered.contains("deactivated"))
}
fn codex_structured_invalid_reason(status_code: u16, upstream_message: Option<&str>) -> String {
let message = upstream_message.unwrap_or_default().trim();
if status_code == 402 && codex_looks_like_workspace_deactivated(Some(message)) {
return format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}工作区已停用 (deactivated_workspace)");
}
if codex_looks_like_account_deactivated(Some(message)) {
let detail = if message.is_empty() {
"OpenAI 账号已停用"
} else {
message
};
return format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}{detail}");
}
if codex_looks_like_token_invalidated(Some(message)) {
let detail = if message.is_empty() {
"Codex Token 无效或已过期"
} else {
message
};
return format!("{OAUTH_EXPIRED_PREFIX}{detail}");
}
if status_code == 401 {
let detail = if message.is_empty() {
"Codex Token 无效或已过期 (401)"
} else {
message
};
return format!("{OAUTH_EXPIRED_PREFIX}{detail}");
}
if status_code == 403 {
let detail = if message.is_empty() {
"Codex 账户访问受限 (403)"
} else {
message
};
return format!("{OAUTH_ACCOUNT_BLOCK_PREFIX}{detail}");
}
message.to_string()
}
fn codex_soft_request_failure_reason(status_code: u16, upstream_message: Option<&str>) -> String {
let detail = upstream_message
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.unwrap_or_else(|| format!("Codex 请求失败 ({status_code})"));
format!("{OAUTH_REQUEST_FAILED_PREFIX}{detail}")
}
fn build_codex_refresh_headers(
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
resolved_oauth_auth: Option<(String, String)>,
) -> Result<BTreeMap<String, String>, String> {
let mut headers = BTreeMap::new();
headers.insert("accept".to_string(), "application/json".to_string());
if let Some((name, value)) = resolved_oauth_auth {
headers.insert(name.to_ascii_lowercase(), value);
} else {
let decrypted_key = transport.key.decrypted_api_key.trim();
if decrypted_key.is_empty() || decrypted_key == "__placeholder__" {
return Err("缺少 OAuth 认证信息,请先授权/刷新 Token".to_string());
}
headers.insert(
"authorization".to_string(),
format!("Bearer {decrypted_key}"),
);
}
let auth_config = transport
.key
.decrypted_auth_config
.as_deref()
.and_then(|raw| serde_json::from_str::<serde_json::Value>(raw).ok());
let oauth_plan_type = normalize_codex_plan_type(
auth_config
.as_ref()
.and_then(|value| value.get("plan_type"))
.and_then(serde_json::Value::as_str),
);
let oauth_account_id = auth_config
.as_ref()
.and_then(|value| value.get("account_id"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty());
if oauth_account_id.is_some() && oauth_plan_type.as_deref() != Some("free") {
headers.insert(
"chatgpt-account-id".to_string(),
oauth_account_id.unwrap_or_default().to_string(),
);
}
Ok(headers)
}
async fn execute_codex_quota_plan(
state: &AppState,
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
headers: BTreeMap<String, String>,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
let plan = ExecutionPlan {
request_id: format!("codex-quota:{}", transport.key.id),
candidate_id: None,
provider_name: Some("codex".to_string()),
provider_id: transport.provider.id.clone(),
endpoint_id: transport.endpoint.id.clone(),
key_id: transport.key.id.clone(),
method: "GET".to_string(),
url: CODEX_WHAM_USAGE_URL.to_string(),
headers,
content_type: None,
content_encoding: None,
body: RequestBody {
json_body: None,
body_bytes_b64: None,
body_ref: None,
},
stream: false,
client_api_format: "openai:cli".to_string(),
provider_api_format: "openai:cli".to_string(),
model_name: Some("codex-wham-usage".to_string()),
proxy: crate::provider_transport::resolve_transport_proxy_snapshot_with_tunnel_affinity(
state, transport,
)
.await,
tls_profile: crate::provider_transport::resolve_transport_tls_profile(transport),
timeouts: crate::provider_transport::resolve_transport_execution_timeouts(transport).or(
Some(ExecutionTimeouts {
connect_ms: Some(30_000),
read_ms: Some(30_000),
write_ms: Some(30_000),
pool_ms: Some(30_000),
total_ms: Some(30_000),
..ExecutionTimeouts::default()
}),
),
};
execute_provider_quota_plan(state, transport, plan, "codex").await
}
pub(crate) async fn refresh_codex_provider_quota_locally(
state: &AppState,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
) -> Result<Option<serde_json::Value>, GatewayError> {
let auto_remove_abnormal_keys = provider_auto_remove_banned_keys(provider.config.as_ref());
let mut results = Vec::new();
let mut success_count = 0usize;
let mut failed_count = 0usize;
let mut auto_removed_count = 0usize;
for key in keys {
let transport = match state
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
.await?
{
Some(transport) => transport,
None => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Provider transport snapshot unavailable",
}));
continue;
}
};
let resolved_oauth_auth = if key.auth_type.trim().eq_ignore_ascii_case("oauth") {
match state.resolve_local_oauth_request_auth(&transport).await? {
Some(crate::provider_transport::LocalResolvedOAuthRequestAuth::Header {
name,
value,
}) => Some((name, value)),
_ => None,
}
} else {
None
};
let headers = match build_codex_refresh_headers(&transport, resolved_oauth_auth) {
Ok(headers) => headers,
Err(message) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": message,
}));
continue;
}
};
let result = match execute_codex_quota_plan(state, &transport, headers).await? {
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": format!("wham/usage 请求执行失败: {detail}"),
"status_code": 502,
}));
continue;
}
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut metadata_update = parse_codex_usage_headers(&result.headers, now_unix_secs)
.map(|metadata| json!({ "codex": metadata }));
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) = (None, None);
let mut status = "error".to_string();
let mut message = None::<String>;
let mut status_code = Some(result.status_code);
if result.status_code == 200 {
if let Some(body_json) = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
{
if let Some(parsed) = parse_codex_wham_usage_response(body_json, now_unix_secs) {
metadata_update = Some(json!({ "codex": parsed }));
(oauth_invalid_at_unix_secs, oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
status = "success".to_string();
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含限额信息".to_string());
}
} else {
message = Some("无法解析 wham/usage API 响应".to_string());
}
} else {
let err_msg = extract_execution_error_message(&result);
message = Some(match err_msg.as_deref() {
Some(detail) if !detail.is_empty() => {
format!(
"wham/usage API 返回状态码 {}: {}",
result.status_code, detail
)
}
_ => format!("wham/usage API 返回状态码 {}", result.status_code),
});
match result.status_code {
401 => {
let (at, reason) = codex_build_invalid_state(
&key,
codex_structured_invalid_reason(401, err_msg.as_deref()),
now_unix_secs,
);
oauth_invalid_at_unix_secs = at;
oauth_invalid_reason = reason;
status = "auth_invalid".to_string();
}
402 => {
if codex_looks_like_workspace_deactivated(err_msg.as_deref()) {
let mut codex_meta = metadata_update
.as_ref()
.and_then(|value| value.get("codex"))
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
codex_meta.insert("updated_at".to_string(), json!(now_unix_secs));
codex_meta.insert("account_disabled".to_string(), json!(true));
codex_meta.insert("reason".to_string(), json!("deactivated_workspace"));
codex_meta.insert(
"message".to_string(),
json!(err_msg
.clone()
.unwrap_or_else(|| "deactivated_workspace".to_string())),
);
let plan_type = transport
.key
.decrypted_auth_config
.as_deref()
.and_then(|raw| serde_json::from_str::<serde_json::Value>(raw).ok())
.and_then(|value| {
value
.get("plan_type")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned)
});
if let Some(plan_type) = plan_type {
codex_meta
.entry("plan_type".to_string())
.or_insert_with(|| json!(plan_type.to_ascii_lowercase()));
}
metadata_update = Some(json!({ "codex": codex_meta }));
let (at, reason) = codex_build_invalid_state(
&key,
codex_structured_invalid_reason(402, err_msg.as_deref()),
now_unix_secs,
);
oauth_invalid_at_unix_secs = at;
oauth_invalid_reason = reason;
status = "workspace_deactivated".to_string();
} else {
let plan_type = transport
.key
.decrypted_auth_config
.as_deref()
.and_then(|raw| serde_json::from_str::<serde_json::Value>(raw).ok())
.and_then(|value| {
value
.get("plan_type")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned)
});
metadata_update = Some(json!({
"codex": build_codex_quota_exhausted_fallback_metadata(
plan_type.as_deref(),
now_unix_secs,
)
}));
(oauth_invalid_at_unix_secs, oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
status = "quota_exhausted".to_string();
}
}
403 => {
let candidate_reason = if codex_looks_like_token_invalidated(err_msg.as_deref())
{
codex_structured_invalid_reason(403, err_msg.as_deref())
} else {
codex_soft_request_failure_reason(403, err_msg.as_deref())
};
let (at, reason) =
codex_build_invalid_state(&key, candidate_reason, now_unix_secs);
oauth_invalid_at_unix_secs = at;
oauth_invalid_reason = reason;
status = "forbidden".to_string();
}
_ => {}
}
}
let auto_removed = auto_remove_abnormal_keys
&& should_auto_remove_structured_reason(oauth_invalid_reason.as_deref());
if auto_removed {
if state.delete_provider_catalog_key(&key.id).await? {
auto_removed_count += 1;
}
} else if !persist_provider_quota_refresh_state(
state,
&key.id,
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason.clone(),
None,
)
.await?
{
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Key 状态写入失败",
}));
continue;
}
if status == "success" {
success_count += 1;
} else {
failed_count += 1;
}
let mut payload = serde_json::Map::new();
payload.insert("key_id".to_string(), json!(key.id));
payload.insert("key_name".to_string(), json!(key.name));
payload.insert("status".to_string(), json!(status));
if let Some(message) = message {
payload.insert("message".to_string(), json!(message));
}
if let Some(status_code) = status_code.take() {
if status_code != 200 {
payload.insert("status_code".to_string(), json!(status_code));
}
}
if let Some(metadata_update) = metadata_update
.as_ref()
.and_then(|value| value.get("codex"))
.cloned()
{
payload.insert("metadata".to_string(), metadata_update);
}
if auto_removed {
payload.insert("auto_removed".to_string(), json!(true));
}
results.push(serde_json::Value::Object(payload));
}
Ok(Some(json!({
"success": success_count,
"failed": failed_count,
"total": results.len(),
"results": results,
"auto_removed": auto_removed_count,
})))
}
@@ -0,0 +1,450 @@
use super::{
coerce_json_f64, execute_provider_quota_plan, extract_execution_error_message,
persist_provider_quota_refresh_state, quota_refresh_success_invalid_state,
ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::provider::shared::{KIRO_USAGE_LIMITS_PATH, KIRO_USAGE_SDK_VERSION};
use crate::handlers::admin::shared::encrypt_catalog_secret_with_fallbacks;
use crate::{AppState, GatewayError};
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use serde_json::json;
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
use url::form_urlencoded;
use uuid::Uuid;
fn compute_kiro_total_usage_limit(breakdown: &serde_json::Value) -> f64 {
let mut total = breakdown
.get("usageLimitWithPrecision")
.and_then(coerce_json_f64)
.unwrap_or(0.0);
if breakdown
.get("freeTrialInfo")
.and_then(serde_json::Value::as_object)
.is_some_and(|free_trial| {
free_trial
.get("freeTrialStatus")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|value| value.eq_ignore_ascii_case("ACTIVE"))
})
{
total += breakdown
.get("freeTrialInfo")
.and_then(|value| value.get("usageLimitWithPrecision"))
.and_then(coerce_json_f64)
.unwrap_or(0.0);
}
if let Some(bonuses) = breakdown
.get("bonuses")
.and_then(serde_json::Value::as_array)
{
for bonus in bonuses {
let is_active = bonus
.get("status")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|value| value.eq_ignore_ascii_case("ACTIVE"));
if is_active {
total += bonus
.get("usageLimit")
.and_then(coerce_json_f64)
.unwrap_or(0.0);
}
}
}
total
}
fn compute_kiro_current_usage(breakdown: &serde_json::Value) -> f64 {
let mut total = breakdown
.get("currentUsageWithPrecision")
.and_then(coerce_json_f64)
.unwrap_or(0.0);
if breakdown
.get("freeTrialInfo")
.and_then(serde_json::Value::as_object)
.is_some_and(|free_trial| {
free_trial
.get("freeTrialStatus")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|value| value.eq_ignore_ascii_case("ACTIVE"))
})
{
total += breakdown
.get("freeTrialInfo")
.and_then(|value| value.get("currentUsageWithPrecision"))
.and_then(coerce_json_f64)
.unwrap_or(0.0);
}
if let Some(bonuses) = breakdown
.get("bonuses")
.and_then(serde_json::Value::as_array)
{
for bonus in bonuses {
let is_active = bonus
.get("status")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|value| value.eq_ignore_ascii_case("ACTIVE"));
if is_active {
total += bonus
.get("currentUsage")
.and_then(coerce_json_f64)
.unwrap_or(0.0);
}
}
}
total
}
fn parse_kiro_usage_response(
value: &serde_json::Value,
updated_at_unix_secs: u64,
) -> Option<serde_json::Value> {
let root = value.as_object()?;
let breakdown = root
.get("usageBreakdownList")
.and_then(serde_json::Value::as_array)
.and_then(|items| items.first())?;
let usage_limit = compute_kiro_total_usage_limit(breakdown);
let current_usage = compute_kiro_current_usage(breakdown);
let remaining = (usage_limit - current_usage).max(0.0);
let usage_percentage = if usage_limit > 0.0 {
((current_usage / usage_limit) * 100.0).min(100.0)
} else {
0.0
};
let mut result = serde_json::Map::new();
result.insert("current_usage".to_string(), json!(current_usage));
result.insert("usage_limit".to_string(), json!(usage_limit));
result.insert("remaining".to_string(), json!(remaining));
result.insert("usage_percentage".to_string(), json!(usage_percentage));
result.insert("updated_at".to_string(), json!(updated_at_unix_secs));
if let Some(subscription_title) = root
.get("subscriptionInfo")
.and_then(|value| value.get("subscriptionTitle"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
result.insert("subscription_title".to_string(), json!(subscription_title));
}
if let Some(next_reset_at) = root
.get("nextDateReset")
.and_then(coerce_json_f64)
.or_else(|| breakdown.get("nextDateReset").and_then(coerce_json_f64))
{
result.insert("next_reset_at".to_string(), json!(next_reset_at));
}
let email = root
.get("desktopUserInfo")
.and_then(|value| value.get("email"))
.or_else(|| root.get("userInfo").and_then(|value| value.get("email")))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty());
if let Some(email) = email {
result.insert("email".to_string(), json!(email));
}
Some(serde_json::Value::Object(result))
}
fn build_kiro_usage_headers(
auth: &crate::provider_transport::kiro::KiroRequestAuth,
) -> BTreeMap<String, String> {
let kiro_version = auth.auth_config.effective_kiro_version();
let machine_id = auth.machine_id.trim();
let ide_tag = if machine_id.is_empty() {
format!("KiroIDE-{kiro_version}")
} else {
format!("KiroIDE-{kiro_version}-{machine_id}")
};
let host = format!(
"q.{}.amazonaws.com",
auth.auth_config.effective_api_region()
);
BTreeMap::from([
(
"x-amz-user-agent".to_string(),
format!("aws-sdk-js/{KIRO_USAGE_SDK_VERSION} {ide_tag}"),
),
(
"user-agent".to_string(),
format!(
"aws-sdk-js/{KIRO_USAGE_SDK_VERSION} ua/2.1 os/other#unknown lang/js md/nodejs#22.21.1 api/codewhispererruntime#1.0.0 m/N,E {ide_tag}"
),
),
("host".to_string(), host),
("amz-sdk-invocation-id".to_string(), Uuid::new_v4().to_string()),
("amz-sdk-request".to_string(), "attempt=1; max=1".to_string()),
("authorization".to_string(), auth.value.clone()),
("connection".to_string(), "close".to_string()),
])
}
fn build_kiro_usage_url(auth: &crate::provider_transport::kiro::KiroRequestAuth) -> String {
let host = format!(
"q.{}.amazonaws.com",
auth.auth_config.effective_api_region()
);
let mut serializer = form_urlencoded::Serializer::new(String::new());
serializer.append_pair("origin", "AI_EDITOR");
serializer.append_pair("resourceType", "AGENTIC_REQUEST");
serializer.append_pair("isEmailRequired", "true");
if let Some(profile_arn) = auth.auth_config.profile_arn_for_payload() {
serializer.append_pair("profileArn", profile_arn);
}
format!(
"https://{host}{KIRO_USAGE_LIMITS_PATH}?{}",
serializer.finish()
)
}
async fn execute_kiro_quota_plan(
state: &AppState,
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
auth: &crate::provider_transport::kiro::KiroRequestAuth,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
let plan = ExecutionPlan {
request_id: format!("kiro-quota:{}", transport.key.id),
candidate_id: None,
provider_name: Some("kiro".to_string()),
provider_id: transport.provider.id.clone(),
endpoint_id: transport.endpoint.id.clone(),
key_id: transport.key.id.clone(),
method: "GET".to_string(),
url: build_kiro_usage_url(auth),
headers: build_kiro_usage_headers(auth),
content_type: None,
content_encoding: None,
body: RequestBody {
json_body: None,
body_bytes_b64: None,
body_ref: None,
},
stream: false,
client_api_format: "claude:cli".to_string(),
provider_api_format: "kiro:usage".to_string(),
model_name: Some("kiro-usage-limits".to_string()),
proxy: crate::provider_transport::resolve_transport_proxy_snapshot_with_tunnel_affinity(
state, transport,
)
.await,
tls_profile: crate::provider_transport::resolve_transport_tls_profile(transport),
timeouts: crate::provider_transport::resolve_transport_execution_timeouts(transport).or(
Some(ExecutionTimeouts {
connect_ms: Some(30_000),
read_ms: Some(30_000),
write_ms: Some(30_000),
pool_ms: Some(30_000),
total_ms: Some(30_000),
..ExecutionTimeouts::default()
}),
),
};
execute_provider_quota_plan(state, transport, plan, "kiro").await
}
pub(crate) async fn refresh_kiro_provider_quota_locally(
state: &AppState,
provider: &StoredProviderCatalogProvider,
endpoint: &StoredProviderCatalogEndpoint,
keys: Vec<StoredProviderCatalogKey>,
) -> Result<Option<serde_json::Value>, GatewayError> {
let mut results = Vec::new();
let mut success_count = 0usize;
let mut failed_count = 0usize;
for key in keys {
let transport = match state
.read_provider_transport_snapshot(&provider.id, &endpoint.id, &key.id)
.await?
{
Some(transport) => transport,
None => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Provider transport snapshot unavailable",
}));
continue;
}
};
let Some(auth) = (match state.resolve_local_oauth_request_auth(&transport).await? {
Some(crate::provider_transport::LocalResolvedOAuthRequestAuth::Kiro(auth)) => {
Some(auth)
}
_ => None,
}) else {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 Kiro 认证配置 (auth_config)",
}));
continue;
};
let result = match execute_kiro_quota_plan(state, &transport, &auth).await? {
ProviderQuotaExecutionOutcome::Response(result) => result,
ProviderQuotaExecutionOutcome::Failure(detail) => {
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": format!("getUsageLimits 请求执行失败: {detail}"),
"status_code": 502,
}));
continue;
}
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut metadata_update = None::<serde_json::Value>;
let mut encrypted_auth_config = None::<String>;
let (mut oauth_invalid_at_unix_secs, mut oauth_invalid_reason) =
quota_refresh_success_invalid_state(&key);
let mut status = "error".to_string();
let mut message = None::<String>;
if result.status_code == 200 {
if let Some(body_json) = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
{
metadata_update = parse_kiro_usage_response(body_json, now_unix_secs)
.map(|metadata| json!({ "kiro": metadata }));
if metadata_update.is_some() {
let auth_config_json = auth.auth_config.to_json_value().to_string();
if let Some(auth_config_json) =
encrypt_catalog_secret_with_fallbacks(state, auth_config_json.as_str())
{
encrypted_auth_config = Some(auth_config_json);
}
status = "success".to_string();
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含限额信息".to_string());
}
} else {
status = "no_metadata".to_string();
message = Some("响应中未包含限额信息".to_string());
}
} else {
let err_msg = extract_execution_error_message(&result);
message = Some(match err_msg.as_deref() {
Some(detail) if !detail.is_empty() => {
format!(
"getUsageLimits 返回状态码 {}: {}",
result.status_code, detail
)
}
_ => format!("getUsageLimits 返回状态码 {}", result.status_code),
});
match result.status_code {
401 => {
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason = Some("Kiro Token 无效或已过期".to_string());
}
403 | 423 => {
let reason = err_msg
.clone()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| format!("HTTP {}", result.status_code));
oauth_invalid_at_unix_secs = Some(now_unix_secs);
oauth_invalid_reason = Some(format!("账户已封禁: {reason}"));
metadata_update = Some(json!({
"kiro": {
"is_banned": true,
"ban_reason": reason,
"banned_at": now_unix_secs,
"updated_at": now_unix_secs,
}
}));
status = "banned".to_string();
}
_ => {}
}
}
if !persist_provider_quota_refresh_state(
state,
&key.id,
metadata_update.as_ref(),
oauth_invalid_at_unix_secs,
oauth_invalid_reason,
encrypted_auth_config,
)
.await?
{
failed_count += 1;
results.push(json!({
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "Key 状态写入失败",
}));
continue;
}
if status == "success" {
success_count += 1;
} else {
failed_count += 1;
}
let mut payload = serde_json::Map::new();
payload.insert("key_id".to_string(), json!(key.id));
payload.insert("key_name".to_string(), json!(key.name));
payload.insert("status".to_string(), json!(status));
if let Some(message) = message {
payload.insert("message".to_string(), json!(message));
}
if let Some(metadata) = metadata_update
.as_ref()
.and_then(|value| value.get("kiro"))
.cloned()
{
payload.insert("metadata".to_string(), metadata);
}
results.push(serde_json::Value::Object(payload));
}
Ok(Some(json!({
"success": success_count,
"failed": failed_count,
"total": success_count + failed_count,
"results": results,
"message": format!("已处理 {} 个 Key", success_count + failed_count),
"auto_removed": 0,
})))
}
@@ -0,0 +1,15 @@
mod antigravity;
mod codex;
mod kiro;
mod shared;
pub(crate) use self::antigravity::refresh_antigravity_provider_quota_locally;
pub(crate) use self::codex::refresh_codex_provider_quota_locally;
pub(crate) use self::kiro::refresh_kiro_provider_quota_locally;
use self::shared::{
coerce_json_bool, coerce_json_f64, coerce_json_string, coerce_json_u64,
execute_provider_quota_plan, extract_execution_error_message, provider_auto_remove_banned_keys,
quota_refresh_success_invalid_state, should_auto_remove_structured_reason,
ProviderQuotaExecutionOutcome,
};
pub(crate) use self::shared::{normalize_string_id_list, persist_provider_quota_refresh_state};
@@ -0,0 +1,225 @@
use crate::handlers::admin::provider::shared::{
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
};
use crate::{AppState, GatewayError};
use aether_contracts::{ExecutionPlan, ExecutionResult};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use std::collections::BTreeSet;
use std::time::{SystemTime, UNIX_EPOCH};
use tracing::warn;
pub(super) enum ProviderQuotaExecutionOutcome {
Response(ExecutionResult),
Failure(String),
}
pub(super) fn provider_auto_remove_banned_keys(config: Option<&serde_json::Value>) -> bool {
config
.and_then(|value| value.get("pool_advanced"))
.and_then(serde_json::Value::as_object)
.and_then(|object| object.get("auto_remove_banned_keys"))
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
}
pub(super) fn should_auto_remove_structured_reason(reason: Option<&str>) -> bool {
reason
.map(str::trim)
.is_some_and(|value| value.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX))
}
pub(crate) fn normalize_string_id_list(values: Option<Vec<String>>) -> Option<Vec<String>> {
let mut out = Vec::new();
let mut seen = BTreeSet::new();
for value in values.into_iter().flatten() {
let trimmed = value.trim();
if trimmed.is_empty() || !seen.insert(trimmed.to_string()) {
continue;
}
out.push(trimmed.to_string());
}
(!out.is_empty()).then_some(out)
}
pub(super) fn coerce_json_u64(value: &serde_json::Value) -> Option<u64> {
match value {
serde_json::Value::Number(number) => number.as_u64(),
serde_json::Value::String(text) => text.trim().parse::<u64>().ok(),
_ => None,
}
}
pub(super) fn coerce_json_f64(value: &serde_json::Value) -> Option<f64> {
match value {
serde_json::Value::Number(number) => number.as_f64(),
serde_json::Value::String(text) => text.trim().parse::<f64>().ok(),
_ => None,
}
}
pub(super) fn coerce_json_bool(value: &serde_json::Value) -> Option<bool> {
match value {
serde_json::Value::Bool(value) => Some(*value),
serde_json::Value::String(text) => match text.trim().to_ascii_lowercase().as_str() {
"true" | "1" => Some(true),
"false" | "0" => Some(false),
_ => None,
},
_ => None,
}
}
fn merge_upstream_metadata(
current: Option<&serde_json::Value>,
updates: &serde_json::Value,
) -> serde_json::Value {
let mut merged = current
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
if let Some(update_object) = updates.as_object() {
for (key, value) in update_object {
merged.insert(key.clone(), value.clone());
}
}
serde_json::Value::Object(merged)
}
pub(super) fn extract_execution_error_message(result: &ExecutionResult) -> Option<String> {
if let Some(body_json) = result
.body
.as_ref()
.and_then(|body| body.json_body.as_ref())
.and_then(serde_json::Value::as_object)
{
if let Some(error) = body_json
.get("error")
.and_then(serde_json::Value::as_object)
{
if let Some(message) = error.get("message").and_then(serde_json::Value::as_str) {
let trimmed = message.trim();
if !trimmed.is_empty() {
return Some(trimmed.to_string());
}
}
}
if let Some(message) = body_json.get("message").and_then(serde_json::Value::as_str) {
let trimmed = message.trim();
if !trimmed.is_empty() {
return Some(trimmed.to_string());
}
}
}
result
.error
.as_ref()
.map(|error| error.message.trim().to_string())
.filter(|value| !value.is_empty())
}
pub(super) fn quota_refresh_success_invalid_state(
key: &StoredProviderCatalogKey,
) -> (Option<u64>, Option<String>) {
let current_reason = key
.oauth_invalid_reason
.as_deref()
.map(str::trim)
.unwrap_or_default();
if current_reason.starts_with(OAUTH_REFRESH_FAILED_PREFIX) {
return (
key.oauth_invalid_at_unix_secs,
(!current_reason.is_empty()).then_some(current_reason.to_string()),
);
}
(None, None)
}
pub(super) fn coerce_json_string(value: Option<&serde_json::Value>) -> Option<String> {
value
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
pub(crate) async fn persist_provider_quota_refresh_state(
state: &AppState,
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>,
) -> Result<bool, GatewayError> {
let Some(mut latest_key) = state
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
.await?
.into_iter()
.next()
else {
return Ok(false);
};
if let Some(metadata_update) = metadata_update {
latest_key.upstream_metadata = Some(merge_upstream_metadata(
latest_key.upstream_metadata.as_ref(),
metadata_update,
));
}
if let Some(encrypted_auth_config) = encrypted_auth_config {
latest_key.encrypted_auth_config = Some(encrypted_auth_config);
}
latest_key.oauth_invalid_at_unix_secs = oauth_invalid_at_unix_secs;
latest_key.oauth_invalid_reason = oauth_invalid_reason;
latest_key.updated_at_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs());
Ok(state
.update_provider_catalog_key(&latest_key)
.await?
.is_some())
}
pub(super) async fn execute_provider_quota_plan(
state: &AppState,
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
plan: ExecutionPlan,
quota_kind: &str,
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
match crate::execution_runtime::execute_execution_runtime_sync_plan(state, None, &plan).await {
Ok(result) => Ok(ProviderQuotaExecutionOutcome::Response(result)),
Err(err) => {
let error = match err {
GatewayError::UpstreamUnavailable { message, .. }
| GatewayError::ControlUnavailable { message, .. }
| GatewayError::Internal(message) => message,
};
let proxy_node_id = plan
.proxy
.as_ref()
.and_then(|proxy| proxy.node_id.as_deref())
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let proxy_url_present = plan
.proxy
.as_ref()
.and_then(|proxy| proxy.url.as_deref())
.map(str::trim)
.is_some_and(|value| !value.is_empty());
warn!(
key_id = %transport.key.id,
endpoint_id = %transport.endpoint.id,
url = %plan.url,
tls_profile = ?plan.tls_profile.as_deref(),
proxy_node_id = ?proxy_node_id,
proxy_url_present,
error = %error,
quota_kind = %quota_kind,
"gateway provider quota execution runtime request failed"
);
Ok(ProviderQuotaExecutionOutcome::Failure(error))
}
}
}
@@ -0,0 +1,647 @@
use super::provider_oauth_quota::{
persist_provider_quota_refresh_state, refresh_antigravity_provider_quota_locally,
refresh_codex_provider_quota_locally, refresh_kiro_provider_quota_locally,
};
use super::provider_oauth_state::{
enrich_admin_provider_oauth_auth_config, json_non_empty_string, json_u64_value,
};
use crate::handlers::admin::provider::shared::{
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
OAUTH_REQUEST_FAILED_PREFIX,
};
use crate::handlers::admin::shared::{
decrypt_catalog_secret_with_fallbacks, encrypt_catalog_secret_with_fallbacks,
parse_catalog_auth_config_json,
};
use crate::{AppState, GatewayError};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::collections::BTreeSet;
use std::time::{SystemTime, UNIX_EPOCH};
use uuid::Uuid;
pub(crate) fn build_internal_control_error_response(
status: http::StatusCode,
message: impl Into<String>,
) -> Response<Body> {
(status, Json(json!({ "detail": message.into() }))).into_response()
}
pub(crate) fn normalize_provider_oauth_refresh_error_message(
status_code: Option<u16>,
body_excerpt: Option<&str>,
) -> String {
let mut message = None::<String>;
let mut error_code = None::<String>;
let mut error_type = None::<String>;
if let Some(body_excerpt) = body_excerpt {
if let Ok(value) = serde_json::from_str::<serde_json::Value>(body_excerpt) {
if let Some(object) = value.as_object() {
if let Some(error_object) =
object.get("error").and_then(serde_json::Value::as_object)
{
message = error_object
.get("message")
.or_else(|| error_object.get("error_description"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
error_code = error_object
.get("code")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
error_type = error_object
.get("type")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
}
if message.is_none() {
message = object
.get("message")
.or_else(|| object.get("error_description"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
}
if error_code.is_none() {
error_code = object
.get("code")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
}
if error_type.is_none() {
error_type = object
.get("type")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.to_ascii_lowercase());
}
}
}
}
let message = message
.or_else(|| {
body_excerpt
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| value.chars().take(300).collect::<String>())
})
.unwrap_or_default();
let lowered = message.to_ascii_lowercase();
let error_code = error_code.unwrap_or_default();
let error_type = error_type.unwrap_or_default();
if error_code == "refresh_token_reused"
|| lowered.contains("already been used to generate a new access token")
{
return "refresh_token 已被使用并轮换,请重新登录授权".to_string();
}
if error_code == "invalid_grant"
|| error_code == "invalid_refresh_token"
|| (lowered.contains("refresh token")
&& ["expired", "revoked", "invalid"]
.iter()
.any(|keyword| lowered.contains(keyword)))
{
return "refresh_token 无效、已过期或已撤销,请重新登录授权".to_string();
}
if error_type == "invalid_request_error" && !message.is_empty() {
return message;
}
if !message.is_empty() {
return message;
}
status_code
.map(|status_code| format!("HTTP {status_code}"))
.unwrap_or_else(|| "未知错误".to_string())
}
pub(crate) fn merge_provider_oauth_refresh_failure_reason(
current_reason: Option<&str>,
refresh_reason: &str,
) -> Option<String> {
let current_reason = current_reason.map(str::trim).unwrap_or_default();
let refresh_reason = refresh_reason.trim();
if refresh_reason.is_empty() {
return (!current_reason.is_empty()).then(|| current_reason.to_string());
}
if current_reason.is_empty() {
return Some(refresh_reason.to_string());
}
if current_reason.starts_with(OAUTH_EXPIRED_PREFIX) {
return None;
}
if current_reason.starts_with(OAUTH_ACCOUNT_BLOCK_PREFIX) {
if let Some((head, _)) = current_reason.split_once("[REFRESH_FAILED]") {
return Some(
format!("{}\n{}", head.trim_end(), refresh_reason)
.trim()
.to_string(),
);
}
return Some(format!("{current_reason}\n{refresh_reason}"));
}
Some(refresh_reason.to_string())
}
pub(crate) fn provider_oauth_key_proxy_value(
proxy_node_id: Option<&str>,
) -> Option<serde_json::Value> {
proxy_node_id
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| json!({ "node_id": value, "enabled": true }))
}
pub(crate) fn provider_oauth_active_api_formats(
endpoints: &[StoredProviderCatalogEndpoint],
) -> Vec<String> {
let mut formats = Vec::new();
let mut seen = BTreeSet::new();
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
let api_format = endpoint.api_format.trim();
if api_format.is_empty() || !seen.insert(api_format.to_string()) {
continue;
}
formats.push(api_format.to_string());
}
formats
}
fn normalize_codex_plan_group_for_provider_oauth(
plan_type: Option<&serde_json::Value>,
) -> Option<String> {
let normalized = plan_type
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())?
.to_ascii_lowercase();
match normalized.as_str() {
"free" => Some("free".to_string()),
"team" | "plus" | "enterprise" => Some("team_plus_enterprise".to_string()),
_ => None,
}
}
fn normalize_provider_oauth_identity_value(value: Option<&serde_json::Value>) -> Option<String> {
value
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn is_codex_provider_oauth_provider_type(value: Option<&serde_json::Value>) -> bool {
value
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|provider_type| provider_type.eq_ignore_ascii_case("codex"))
}
fn match_codex_provider_oauth_identity(
new_auth_config: &serde_json::Map<String, serde_json::Value>,
existing_auth_config: &serde_json::Map<String, serde_json::Value>,
) -> Option<bool> {
let new_provider_type = new_auth_config.get("provider_type");
let existing_provider_type = existing_auth_config.get("provider_type");
if !is_codex_provider_oauth_provider_type(new_provider_type)
&& !is_codex_provider_oauth_provider_type(existing_provider_type)
{
return None;
}
let new_account_user_id =
normalize_provider_oauth_identity_value(new_auth_config.get("account_user_id"));
let existing_account_user_id =
normalize_provider_oauth_identity_value(existing_auth_config.get("account_user_id"));
if let (Some(new_account_user_id), Some(existing_account_user_id)) =
(new_account_user_id, existing_account_user_id)
{
return Some(new_account_user_id == existing_account_user_id);
}
let new_account_id = normalize_provider_oauth_identity_value(new_auth_config.get("account_id"));
let existing_account_id =
normalize_provider_oauth_identity_value(existing_auth_config.get("account_id"));
let new_user_id = normalize_provider_oauth_identity_value(new_auth_config.get("user_id"));
let existing_user_id =
normalize_provider_oauth_identity_value(existing_auth_config.get("user_id"));
let new_email = normalize_provider_oauth_identity_value(new_auth_config.get("email"));
let existing_email = normalize_provider_oauth_identity_value(existing_auth_config.get("email"));
if let (Some(new_account_id), Some(existing_account_id)) =
(new_account_id.as_deref(), existing_account_id.as_deref())
{
if new_account_id != existing_account_id {
return Some(false);
}
}
if let (
Some(new_account_id),
Some(existing_account_id),
Some(new_user_id),
Some(existing_user_id),
) = (
new_account_id.as_deref(),
existing_account_id.as_deref(),
new_user_id.as_deref(),
existing_user_id.as_deref(),
) {
return Some(new_account_id == existing_account_id && new_user_id == existing_user_id);
}
if let (
Some(new_account_id),
Some(existing_account_id),
Some(new_email),
Some(existing_email),
) = (
new_account_id.as_deref(),
existing_account_id.as_deref(),
new_email.as_deref(),
existing_email.as_deref(),
) {
return Some(new_account_id == existing_account_id && new_email == existing_email);
}
None
}
fn is_codex_cross_plan_group_non_duplicate(
new_auth_config: &serde_json::Map<String, serde_json::Value>,
existing_auth_config: &serde_json::Map<String, serde_json::Value>,
) -> bool {
let new_provider_type = new_auth_config.get("provider_type");
let existing_provider_type = existing_auth_config.get("provider_type");
if !is_codex_provider_oauth_provider_type(new_provider_type)
&& !is_codex_provider_oauth_provider_type(existing_provider_type)
{
return false;
}
let new_group = normalize_codex_plan_group_for_provider_oauth(new_auth_config.get("plan_type"));
let existing_group =
normalize_codex_plan_group_for_provider_oauth(existing_auth_config.get("plan_type"));
matches!(
(new_group.as_deref(), existing_group.as_deref()),
(Some(left), Some(right)) if left != right
)
}
pub(crate) async fn find_duplicate_provider_oauth_key(
state: &AppState,
provider_id: &str,
auth_config: &serde_json::Map<String, serde_json::Value>,
exclude_key_id: Option<&str>,
) -> Result<Option<StoredProviderCatalogKey>, String> {
let new_email = normalize_provider_oauth_identity_value(auth_config.get("email"));
let new_user_id = normalize_provider_oauth_identity_value(auth_config.get("user_id"));
let new_auth_method = normalize_provider_oauth_identity_value(auth_config.get("auth_method"));
if new_email.is_none() && new_user_id.is_none() {
return Ok(None);
}
let existing_keys = state
.list_provider_catalog_keys_by_provider_ids(&[provider_id.to_string()])
.await
.map_err(|err| format!("{err:?}"))?;
for existing_key in existing_keys.into_iter().filter(|key| {
key.auth_type.trim().eq_ignore_ascii_case("oauth")
&& exclude_key_id.is_none_or(|exclude| key.id != exclude)
}) {
let Some(existing_auth_config) = parse_catalog_auth_config_json(state, &existing_key)
else {
continue;
};
let existing_email =
normalize_provider_oauth_identity_value(existing_auth_config.get("email"));
let existing_user_id =
normalize_provider_oauth_identity_value(existing_auth_config.get("user_id"));
let existing_auth_method =
normalize_provider_oauth_identity_value(existing_auth_config.get("auth_method"));
let mut is_duplicate = false;
let codex_identity_match =
match_codex_provider_oauth_identity(auth_config, &existing_auth_config);
if let Some(codex_identity_match) = codex_identity_match {
is_duplicate = codex_identity_match;
}
if codex_identity_match.is_none()
&& !is_duplicate
&& new_user_id.is_some()
&& existing_user_id.is_some()
&& new_user_id == existing_user_id
&& !is_codex_cross_plan_group_non_duplicate(auth_config, &existing_auth_config)
{
is_duplicate = true;
}
if codex_identity_match.is_none()
&& !is_duplicate
&& new_email.is_some()
&& existing_email.is_some()
&& new_email == existing_email
{
let is_kiro = auth_config
.get("provider_type")
.and_then(serde_json::Value::as_str)
.is_some_and(|value| value.eq_ignore_ascii_case("kiro"))
|| existing_auth_config
.get("provider_type")
.and_then(serde_json::Value::as_str)
.is_some_and(|value| value.eq_ignore_ascii_case("kiro"));
if is_kiro {
if new_auth_method.is_some()
&& existing_auth_method.is_some()
&& new_auth_method
.as_deref()
.zip(existing_auth_method.as_deref())
.is_some_and(|(left, right)| left.eq_ignore_ascii_case(right))
{
is_duplicate = true;
}
} else if !is_codex_cross_plan_group_non_duplicate(auth_config, &existing_auth_config) {
is_duplicate = true;
}
}
if !is_duplicate {
continue;
}
if !existing_key.is_active {
return Ok(Some(existing_key));
}
let identifier =
normalize_provider_oauth_identity_value(auth_config.get("account_user_id"))
.or_else(|| normalize_provider_oauth_identity_value(auth_config.get("account_id")))
.or_else(|| new_email.clone())
.or_else(|| new_user_id.clone())
.unwrap_or_default();
return Err(format!(
"该 OAuth 账号 ({identifier}) 已存在于当前 Provider 中(名称: {})",
existing_key.name
));
}
Ok(None)
}
pub(crate) fn build_provider_oauth_auth_config_from_token_payload(
provider_type: &str,
token_payload: &serde_json::Value,
) -> (
serde_json::Map<String, serde_json::Value>,
Option<String>,
Option<String>,
Option<u64>,
) {
let access_token = json_non_empty_string(token_payload.get("access_token"));
let refresh_token = json_non_empty_string(token_payload.get("refresh_token"));
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let expires_at = json_u64_value(token_payload.get("expires_in"))
.map(|expires_in| now_unix_secs.saturating_add(expires_in));
let mut auth_config = serde_json::Map::new();
auth_config.insert("provider_type".to_string(), json!(provider_type));
auth_config.insert("updated_at".to_string(), json!(now_unix_secs));
if let Some(token_type) = token_payload.get("token_type").cloned() {
auth_config.insert("token_type".to_string(), token_type);
}
if let Some(refresh_token) = refresh_token.as_ref() {
auth_config.insert("refresh_token".to_string(), json!(refresh_token));
}
if let Some(expires_at) = expires_at {
auth_config.insert("expires_at".to_string(), json!(expires_at));
}
if let Some(scope) = token_payload.get("scope").cloned() {
auth_config.insert("scope".to_string(), scope);
}
enrich_admin_provider_oauth_auth_config(provider_type, &mut auth_config, token_payload);
(auth_config, access_token, refresh_token, expires_at)
}
pub(crate) async fn create_provider_oauth_catalog_key(
state: &AppState,
provider_id: &str,
name: &str,
access_token: &str,
auth_config: &serde_json::Map<String, serde_json::Value>,
api_formats: &[String],
proxy: Option<serde_json::Value>,
expires_at_unix_secs: Option<u64>,
) -> Result<Option<StoredProviderCatalogKey>, GatewayError> {
let Some(encrypted_api_key) = encrypt_catalog_secret_with_fallbacks(state, access_token) else {
return Ok(None);
};
let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let Some(encrypted_auth_config) =
encrypt_catalog_secret_with_fallbacks(state, &auth_config_json)
else {
return Ok(None);
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut record = StoredProviderCatalogKey::new(
Uuid::new_v4().to_string(),
provider_id.to_string(),
name.to_string(),
"oauth".to_string(),
None,
true,
)
.map_err(|err| GatewayError::Internal(err.to_string()))?
.with_transport_fields(
Some(json!(api_formats)),
encrypted_api_key,
Some(encrypted_auth_config),
None,
None,
None,
expires_at_unix_secs,
proxy,
None,
)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
record.internal_priority = 50;
record.cache_ttl_minutes = 5;
record.max_probe_interval_minutes = 32;
record.request_count = Some(0);
record.success_count = Some(0);
record.error_count = Some(0);
record.total_response_time_ms = Some(0);
record.health_by_format = Some(json!({}));
record.circuit_breaker_by_format = Some(json!({}));
record.created_at_unix_secs = Some(now_unix_secs);
record.updated_at_unix_secs = Some(now_unix_secs);
state.create_provider_catalog_key(&record).await
}
pub(crate) async fn update_existing_provider_oauth_catalog_key(
state: &AppState,
existing_key: &StoredProviderCatalogKey,
access_token: &str,
auth_config: &serde_json::Map<String, serde_json::Value>,
proxy: Option<serde_json::Value>,
expires_at_unix_secs: Option<u64>,
) -> Result<Option<StoredProviderCatalogKey>, GatewayError> {
let Some(encrypted_api_key) = encrypt_catalog_secret_with_fallbacks(state, access_token) else {
return Ok(None);
};
let auth_config_json = serde_json::to_string(&serde_json::Value::Object(auth_config.clone()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let Some(encrypted_auth_config) =
encrypt_catalog_secret_with_fallbacks(state, &auth_config_json)
else {
return Ok(None);
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut updated = existing_key.clone();
updated.encrypted_api_key = encrypted_api_key;
updated.encrypted_auth_config = Some(encrypted_auth_config);
updated.is_active = true;
updated.expires_at_unix_secs = expires_at_unix_secs;
updated.oauth_invalid_at_unix_secs = None;
updated.oauth_invalid_reason = None;
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);
state.update_provider_catalog_key(&updated).await
}
pub(crate) fn provider_oauth_runtime_endpoint_for_provider(
provider_type: &str,
endpoints: Vec<StoredProviderCatalogEndpoint>,
) -> Option<StoredProviderCatalogEndpoint> {
let provider_type = provider_type.trim().to_ascii_lowercase();
match provider_type.as_str() {
"codex" => endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("openai:cli")
}),
"antigravity" => endpoints.into_iter().find(|endpoint| {
endpoint.is_active
&& (endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:chat")
|| endpoint
.api_format
.trim()
.eq_ignore_ascii_case("gemini:cli"))
}),
"kiro" => endpoints
.iter()
.find(|endpoint| {
endpoint.is_active
&& endpoint
.api_format
.trim()
.eq_ignore_ascii_case("claude:cli")
})
.cloned()
.or_else(|| endpoints.into_iter().find(|endpoint| endpoint.is_active)),
_ => endpoints.into_iter().find(|endpoint| endpoint.is_active),
}
}
pub(crate) async fn refresh_provider_oauth_account_state_after_update(
state: &AppState,
provider: &StoredProviderCatalogProvider,
key_id: &str,
) -> Result<(bool, Option<String>), GatewayError> {
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
if !matches!(provider_type.as_str(), "codex" | "kiro" | "antigravity") {
return Ok((false, None));
}
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let Some(endpoint) = provider_oauth_runtime_endpoint_for_provider(&provider_type, endpoints)
else {
return Ok((false, None));
};
let Some(key) = state
.read_provider_catalog_keys_by_ids(&[key_id.to_string()])
.await?
.into_iter()
.next()
else {
return Ok((false, None));
};
let payload = match provider_type.as_str() {
"codex" => {
refresh_codex_provider_quota_locally(state, provider, &endpoint, vec![key]).await?
}
"kiro" => {
refresh_kiro_provider_quota_locally(state, provider, &endpoint, vec![key]).await?
}
"antigravity" => {
refresh_antigravity_provider_quota_locally(state, provider, &endpoint, vec![key])
.await?
}
_ => None,
};
let Some(payload) = payload else {
return Ok((false, None));
};
let success = payload
.get("success")
.and_then(serde_json::Value::as_u64)
.unwrap_or(0);
let error = if success == 0 {
payload
.get("results")
.and_then(serde_json::Value::as_array)
.and_then(|results| results.first())
.and_then(|value| value.get("message"))
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned)
} else {
None
};
Ok((true, error))
}
@@ -0,0 +1,996 @@
use super::{
build_internal_control_error_response, normalize_provider_oauth_refresh_error_message,
};
use crate::handlers::admin::provider::shared::ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL;
use crate::provider_transport::provider_types::{
provider_type_admin_oauth_template, provider_type_is_fixed_for_admin_oauth,
ProviderOAuthTemplate, ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES,
};
use crate::{AppState, GatewayError};
use aether_data::repository::provider_oauth::{
build_provider_oauth_batch_task_status_payload, provider_oauth_batch_task_storage_key,
provider_oauth_device_session_storage_key, provider_oauth_state_storage_key,
StoredAdminProviderOAuthDeviceSession, StoredAdminProviderOAuthState,
PROVIDER_OAUTH_BATCH_TASK_TTL_SECS, PROVIDER_OAUTH_STATE_TTL_SECS,
};
use axum::{body::Body, http, response::Response};
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use serde_json::json;
use sha2::{Digest, Sha256};
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
use url::{form_urlencoded, Url};
use uuid::Uuid;
pub(crate) fn is_fixed_provider_type_for_provider_oauth(provider_type: &str) -> bool {
provider_type_is_fixed_for_admin_oauth(provider_type)
}
pub(crate) fn admin_provider_oauth_template(provider_type: &str) -> Option<ProviderOAuthTemplate> {
provider_type_admin_oauth_template(provider_type)
}
pub(crate) fn build_admin_provider_oauth_supported_types_payload() -> Vec<serde_json::Value> {
ADMIN_PROVIDER_OAUTH_TEMPLATE_TYPES
.into_iter()
.filter_map(|provider_type| admin_provider_oauth_template(provider_type))
.map(|template| {
json!({
"provider_type": template.provider_type,
"display_name": template.display_name,
"scopes": template.scopes,
"redirect_uri": template.redirect_uri,
"authorize_url": template.authorize_url,
"token_url": template.token_url,
"use_pkce": template.use_pkce,
})
})
.collect()
}
pub(crate) fn build_admin_provider_oauth_backend_unavailable_response() -> Response<Body> {
build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL,
)
}
const KIRO_DEVICE_DEFAULT_START_URL: &str = "https://view.awsapps.com/start";
const KIRO_DEVICE_DEFAULT_REGION: &str = "us-east-1";
const KIRO_IDC_AMZ_USER_AGENT: &str =
"aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown api/sso-oidc#3.738.0 m/E KiroIDE";
pub(crate) fn default_kiro_device_start_url() -> String {
KIRO_DEVICE_DEFAULT_START_URL.to_string()
}
pub(crate) fn default_kiro_device_region() -> String {
KIRO_DEVICE_DEFAULT_REGION.to_string()
}
pub(crate) fn normalize_kiro_device_region(value: Option<&str>) -> Option<String> {
let value = value
.map(str::trim)
.filter(|value| !value.is_empty())
.unwrap_or(KIRO_DEVICE_DEFAULT_REGION);
value
.chars()
.all(|ch| ch.is_ascii_lowercase() || ch.is_ascii_digit() || ch == '-')
.then(|| value.to_string())
}
pub(crate) fn current_unix_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0)
}
pub(crate) async fn save_provider_oauth_device_session(
state: &AppState,
session_id: &str,
session: &StoredAdminProviderOAuthDeviceSession,
ttl_seconds: u64,
) -> Result<(), Response<Body>> {
let key = provider_oauth_device_session_storage_key(session_id);
let value = serde_json::to_string(session).map_err(|_| {
build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
)
})?;
if let Some(runner) = state.redis_kv_runner() {
runner
.setex(&key, &value, Some(ttl_seconds))
.await
.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
)
})?;
return Ok(());
}
if state.save_provider_oauth_device_session_for_tests(&key, &value) {
return Ok(());
}
Err(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
))
}
pub(crate) async fn read_provider_oauth_device_session(
state: &AppState,
session_id: &str,
) -> Result<Option<StoredAdminProviderOAuthDeviceSession>, GatewayError> {
let key = provider_oauth_device_session_storage_key(session_id);
let raw = if let Some(runner) = state.redis_kv_runner() {
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let namespaced_key = runner.keyspace().key(&key);
redis::cmd("GET")
.arg(&namespaced_key)
.query_async::<Option<String>>(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
} else {
state.load_provider_oauth_device_session_for_tests(&key)
};
raw.map(|value| {
serde_json::from_str::<StoredAdminProviderOAuthDeviceSession>(&value)
.map_err(|err| GatewayError::Internal(err.to_string()))
})
.transpose()
}
async fn post_kiro_device_oidc_json(
state: &AppState,
endpoint_key: &str,
default_url: String,
body: serde_json::Value,
) -> Result<serde_json::Value, Response<Body>> {
let url = state.provider_oauth_token_url(endpoint_key, &default_url);
let host = Url::parse(&url)
.ok()
.and_then(|value| value.host_str().map(ToOwned::to_owned))
.unwrap_or_default();
let response = state
.client
.post(url)
.header("Content-Type", "application/json")
.header("Accept", "*/*")
.header("User-Agent", "node")
.header("x-amz-user-agent", KIRO_IDC_AMZ_USER_AGENT)
.header("Host", host)
.json(&body)
.send()
.await
.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"发起设备授权失败: unknown",
)
})?;
let status = response.status();
let body_text = response.text().await.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"发起设备授权失败: unknown",
)
})?;
serde_json::from_str::<serde_json::Value>(&body_text).or_else(|_| {
Ok(json!({
"_error": !status.is_success(),
"error": body_text.trim(),
}))
})
}
pub(crate) async fn register_admin_kiro_device_oidc_client(
state: &AppState,
region: &str,
start_url: &str,
) -> Result<serde_json::Value, Response<Body>> {
let payload = post_kiro_device_oidc_json(
state,
"kiro_device_register",
format!("https://oidc.{region}.amazonaws.com/client/register"),
json!({
"clientName": "Aether Gateway",
"clientType": "public",
"scopes": [
"codewhisperer:completions",
"codewhisperer:analysis",
"codewhisperer:conversations",
"codewhisperer:transformations",
"codewhisperer:taskassist"
],
"grantTypes": [
"urn:ietf:params:oauth:grant-type:device_code",
"refresh_token"
],
"issuerUrl": start_url,
}),
)
.await?;
if payload
.get("_error")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
{
let error_desc = json_non_empty_string(payload.get("error_description"))
.or_else(|| json_non_empty_string(payload.get("error")))
.unwrap_or_else(|| "unknown".to_string());
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("注册 OIDC 客户端失败: {error_desc}"),
));
}
Ok(payload)
}
pub(crate) async fn start_admin_kiro_device_authorization(
state: &AppState,
region: &str,
client_id: &str,
client_secret: &str,
start_url: &str,
) -> Result<serde_json::Value, Response<Body>> {
let payload = post_kiro_device_oidc_json(
state,
"kiro_device_authorize",
format!("https://oidc.{region}.amazonaws.com/device_authorization"),
json!({
"clientId": client_id,
"clientSecret": client_secret,
"startUrl": start_url,
}),
)
.await?;
if payload
.get("_error")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
{
let error_desc = json_non_empty_string(payload.get("error_description"))
.or_else(|| json_non_empty_string(payload.get("error")))
.unwrap_or_else(|| "unknown".to_string());
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("发起设备授权失败: {error_desc}"),
));
}
Ok(payload)
}
pub(crate) async fn poll_admin_kiro_device_token(
state: &AppState,
region: &str,
client_id: &str,
client_secret: &str,
device_code: &str,
) -> Result<serde_json::Value, Response<Body>> {
post_kiro_device_oidc_json(
state,
"kiro_device_poll",
format!("https://oidc.{region}.amazonaws.com/token"),
json!({
"clientId": client_id,
"clientSecret": client_secret,
"grantType": "urn:ietf:params:oauth:grant-type:device_code",
"deviceCode": device_code,
}),
)
.await
}
pub(crate) fn build_kiro_device_key_name(
email: Option<&str>,
refresh_token: Option<&str>,
) -> String {
if let Some(email) = email.map(str::trim).filter(|value| !value.is_empty()) {
return format!("{email} (idc)");
}
let fallback = refresh_token
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| {
let digest = Sha256::digest(value.as_bytes());
digest[..3]
.iter()
.map(|byte| format!("{byte:02x}"))
.collect::<String>()
})
.unwrap_or_else(|| "unknown".to_string());
format!("kiro_{fallback} (idc)")
}
pub(crate) fn generate_provider_oauth_nonce() -> String {
format!("{}{}", Uuid::new_v4().simple(), Uuid::new_v4().simple())
}
pub(crate) fn generate_provider_oauth_pkce_verifier() -> String {
format!(
"{}{}{}",
Uuid::new_v4().simple(),
Uuid::new_v4().simple(),
Uuid::new_v4().simple()
)
}
pub(crate) fn provider_oauth_pkce_s256(verifier: &str) -> String {
let digest = Sha256::digest(verifier.as_bytes());
URL_SAFE_NO_PAD.encode(digest)
}
pub(crate) async fn save_provider_oauth_state(
state: &AppState,
key_id: &str,
provider_id: &str,
provider_type: &str,
pkce_verifier: Option<&str>,
) -> Result<String, GatewayError> {
let nonce = generate_provider_oauth_nonce();
let payload = json!({
"nonce": nonce,
"key_id": key_id,
"provider_id": provider_id,
"provider_type": provider_type,
"pkce_verifier": pkce_verifier,
"created_at": SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_secs())
.unwrap_or(0),
});
let key = provider_oauth_state_storage_key(&nonce);
let value = payload.to_string();
if let Some(runner) = state.redis_kv_runner() {
runner
.setex(&key, &value, Some(PROVIDER_OAUTH_STATE_TTL_SECS))
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(nonce);
}
if state.save_provider_oauth_state_for_tests(&key, &value) {
return Ok(nonce);
}
Err(GatewayError::Internal(
"provider oauth redis unavailable".to_string(),
))
}
pub(crate) fn parse_provider_oauth_callback_params(callback_url: &str) -> BTreeMap<String, String> {
let mut merged = BTreeMap::new();
let Ok(url) = Url::parse(callback_url.trim()) else {
return merged;
};
for (key, value) in url.query_pairs() {
merged.insert(key.into_owned(), value.into_owned());
}
if let Some(fragment) = url.fragment() {
for (key, value) in form_urlencoded::parse(fragment.as_bytes()) {
merged
.entry(key.into_owned())
.or_insert_with(|| value.into_owned());
}
}
if let Some(code) = merged.get("code").cloned() {
if let Some((code_part, state_part)) = code.split_once("#state=") {
merged.insert("code".to_string(), code_part.to_string());
merged
.entry("state".to_string())
.or_insert_with(|| state_part.to_string());
}
}
merged
}
pub(crate) async fn consume_provider_oauth_state(
state: &AppState,
nonce: &str,
) -> Result<Option<StoredAdminProviderOAuthState>, GatewayError> {
let key = provider_oauth_state_storage_key(nonce);
let raw = if let Some(runner) = state.redis_kv_runner() {
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let namespaced_key = runner.keyspace().key(&key);
redis::cmd("GETDEL")
.arg(&namespaced_key)
.query_async::<Option<String>>(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
} else {
state.take_provider_oauth_state_for_tests(&key)
};
raw.map(|value| {
serde_json::from_str::<StoredAdminProviderOAuthState>(&value)
.map_err(|err| GatewayError::Internal(err.to_string()))
})
.transpose()
}
pub(crate) fn json_non_empty_string(value: Option<&serde_json::Value>) -> Option<String> {
value
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
pub(crate) fn json_u64_value(value: Option<&serde_json::Value>) -> Option<u64> {
match value? {
serde_json::Value::Number(number) => number.as_u64(),
serde_json::Value::String(value) => value.trim().parse::<u64>().ok(),
_ => None,
}
}
pub(crate) fn decode_jwt_claims(token: &str) -> Option<serde_json::Map<String, serde_json::Value>> {
let payload = token.split('.').nth(1)?;
let bytes = URL_SAFE_NO_PAD.decode(payload.as_bytes()).ok()?;
serde_json::from_slice::<serde_json::Value>(&bytes)
.ok()?
.as_object()
.cloned()
}
fn merge_missing_auth_config_fields(
auth_config: &mut serde_json::Map<String, serde_json::Value>,
source: &serde_json::Map<String, serde_json::Value>,
fields: &[&str],
) {
for field in fields {
if auth_config.contains_key(*field) {
continue;
}
if let Some(value) = source.get(*field).cloned() {
auth_config.insert((*field).to_string(), value);
}
}
}
fn first_json_non_empty_string(
values: impl IntoIterator<Item = Option<serde_json::Value>>,
) -> Option<String> {
values.into_iter().find_map(|value| match value {
Some(serde_json::Value::String(value)) => {
let normalized = value.trim();
(!normalized.is_empty()).then(|| normalized.to_string())
}
_ => None,
})
}
fn extract_codex_auth_fields_from_object(
source: &serde_json::Map<String, serde_json::Value>,
) -> serde_json::Map<String, serde_json::Value> {
let auth = source
.get("https://api.openai.com/auth")
.and_then(serde_json::Value::as_object);
let mut result = serde_json::Map::new();
if let Some(email) = first_json_non_empty_string([
source.get("email").cloned(),
auth.and_then(|value| value.get("email")).cloned(),
]) {
result.insert("email".to_string(), json!(email));
}
if let Some(account_id) = first_json_non_empty_string([
auth.and_then(|value| value.get("chatgpt_account_id"))
.cloned(),
auth.and_then(|value| value.get("chatgptAccountId"))
.cloned(),
auth.and_then(|value| value.get("account_id")).cloned(),
auth.and_then(|value| value.get("accountId")).cloned(),
source.get("chatgpt_account_id").cloned(),
source.get("chatgptAccountId").cloned(),
source.get("account_id").cloned(),
source.get("accountId").cloned(),
]) {
result.insert("account_id".to_string(), json!(account_id));
}
if let Some(account_user_id) = first_json_non_empty_string([
auth.and_then(|value| value.get("chatgpt_account_user_id"))
.cloned(),
auth.and_then(|value| value.get("chatgptAccountUserId"))
.cloned(),
auth.and_then(|value| value.get("account_user_id")).cloned(),
auth.and_then(|value| value.get("accountUserId")).cloned(),
source.get("chatgpt_account_user_id").cloned(),
source.get("chatgptAccountUserId").cloned(),
source.get("account_user_id").cloned(),
source.get("accountUserId").cloned(),
]) {
result.insert("account_user_id".to_string(), json!(account_user_id));
}
if let Some(plan_type) = first_json_non_empty_string([
auth.and_then(|value| value.get("chatgpt_plan_type"))
.cloned(),
auth.and_then(|value| value.get("chatgptPlanType")).cloned(),
auth.and_then(|value| value.get("plan_type")).cloned(),
auth.and_then(|value| value.get("planType")).cloned(),
source.get("chatgpt_plan_type").cloned(),
source.get("chatgptPlanType").cloned(),
source.get("plan_type").cloned(),
source.get("planType").cloned(),
]) {
result.insert("plan_type".to_string(), json!(plan_type));
}
if let Some(user_id) = first_json_non_empty_string([
auth.and_then(|value| value.get("chatgpt_user_id")).cloned(),
auth.and_then(|value| value.get("chatgptUserId")).cloned(),
auth.and_then(|value| value.get("user_id")).cloned(),
auth.and_then(|value| value.get("userId")).cloned(),
source.get("chatgpt_user_id").cloned(),
source.get("chatgptUserId").cloned(),
source.get("user_id").cloned(),
source.get("userId").cloned(),
source.get("sub").cloned(),
]) {
result.insert("user_id".to_string(), json!(user_id));
}
if let Some(organizations) = auth
.and_then(|value| value.get("organizations"))
.and_then(serde_json::Value::as_array)
.filter(|value| !value.is_empty())
{
result.insert(
"organizations".to_string(),
serde_json::Value::Array(organizations.clone()),
);
}
result
}
pub(crate) fn enrich_admin_provider_oauth_auth_config(
provider_type: &str,
auth_config: &mut serde_json::Map<String, serde_json::Value>,
token_payload: &serde_json::Value,
) {
let Some(token_payload_object) = token_payload.as_object() else {
return;
};
merge_missing_auth_config_fields(
auth_config,
token_payload_object,
&[
"email",
"account_id",
"account_user_id",
"plan_type",
"user_id",
"account_name",
],
);
if !provider_type.eq_ignore_ascii_case("codex") {
return;
}
let codex_fields = extract_codex_auth_fields_from_object(token_payload_object);
merge_missing_auth_config_fields(
auth_config,
&codex_fields,
&[
"email",
"account_id",
"account_user_id",
"plan_type",
"user_id",
"organizations",
],
);
for token_field in ["id_token", "idToken", "access_token", "accessToken"] {
let Some(token) = json_non_empty_string(token_payload.get(token_field)) else {
continue;
};
let Some(claims) = decode_jwt_claims(&token) else {
continue;
};
merge_missing_auth_config_fields(
auth_config,
&claims,
&[
"email",
"account_id",
"account_user_id",
"plan_type",
"user_id",
"account_name",
],
);
let codex_claim_fields = extract_codex_auth_fields_from_object(&claims);
merge_missing_auth_config_fields(
auth_config,
&codex_claim_fields,
&[
"email",
"account_id",
"account_user_id",
"plan_type",
"user_id",
"organizations",
],
);
}
}
#[cfg(test)]
mod tests {
use super::enrich_admin_provider_oauth_auth_config;
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use serde_json::json;
fn sample_unsigned_jwt(payload: serde_json::Value) -> String {
let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"none","typ":"JWT"}"#);
let payload = URL_SAFE_NO_PAD.encode(payload.to_string());
format!("{header}.{payload}.sig")
}
#[test]
fn codex_enrichment_extracts_identity_from_nested_auth_claims() {
let access_token = sample_unsigned_jwt(json!({
"email": "[email protected]",
"https://api.openai.com/auth": {
"chatgpt_account_id": "acc-1",
"chatgpt_account_user_id": "user-1__acc-1",
"chatgpt_plan_type": "team",
"chatgpt_user_id": "user-1",
"organizations": [
{"id": "org-1", "title": "Personal", "is_default": true}
],
}
}));
let token_payload = json!({
"access_token": access_token,
});
let mut auth_config = serde_json::Map::new();
enrich_admin_provider_oauth_auth_config("codex", &mut auth_config, &token_payload);
assert_eq!(auth_config.get("email"), Some(&json!("[email protected]")));
assert_eq!(auth_config.get("account_id"), Some(&json!("acc-1")));
assert_eq!(
auth_config.get("account_user_id"),
Some(&json!("user-1__acc-1"))
);
assert_eq!(auth_config.get("plan_type"), Some(&json!("team")));
assert_eq!(auth_config.get("user_id"), Some(&json!("user-1")));
assert_eq!(
auth_config.get("organizations"),
Some(&json!([
{"id": "org-1", "title": "Personal", "is_default": true}
]))
);
}
#[test]
fn codex_enrichment_normalizes_direct_chatgpt_alias_fields() {
let token_payload = json!({
"email": "[email protected]",
"chatgpt_account_id": "acc-2",
"chatgpt_account_user_id": "user-2__acc-2",
"chatgpt_plan_type": "plus",
"chatgpt_user_id": "user-2",
});
let mut auth_config = serde_json::Map::new();
enrich_admin_provider_oauth_auth_config("codex", &mut auth_config, &token_payload);
assert_eq!(auth_config.get("email"), Some(&json!("[email protected]")));
assert_eq!(auth_config.get("account_id"), Some(&json!("acc-2")));
assert_eq!(
auth_config.get("account_user_id"),
Some(&json!("user-2__acc-2"))
);
assert_eq!(auth_config.get("plan_type"), Some(&json!("plus")));
assert_eq!(auth_config.get("user_id"), Some(&json!("user-2")));
}
}
pub(crate) async fn exchange_admin_provider_oauth_code(
state: &AppState,
template: ProviderOAuthTemplate,
code: &str,
state_nonce: &str,
pkce_verifier: Option<&str>,
) -> Result<serde_json::Value, Response<Body>> {
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
let request = state.client.post(token_url);
let response = if template.provider_type == "claude_code" {
let mut body = serde_json::Map::from_iter([
(
"grant_type".to_string(),
serde_json::Value::String("authorization_code".to_string()),
),
(
"client_id".to_string(),
serde_json::Value::String(template.client_id.to_string()),
),
(
"redirect_uri".to_string(),
serde_json::Value::String(template.redirect_uri.to_string()),
),
(
"code".to_string(),
serde_json::Value::String(code.to_string()),
),
(
"state".to_string(),
serde_json::Value::String(state_nonce.to_string()),
),
]);
if let Some(verifier) = pkce_verifier {
body.insert(
"code_verifier".to_string(),
serde_json::Value::String(verifier.to_string()),
);
}
request
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.json(&serde_json::Value::Object(body))
.send()
.await
} else {
let mut form = vec![
("grant_type", "authorization_code".to_string()),
("client_id", template.client_id.to_string()),
("redirect_uri", template.redirect_uri.to_string()),
("code", code.to_string()),
];
if !template.client_secret.trim().is_empty() {
form.push(("client_secret", template.client_secret.to_string()));
}
if let Some(verifier) = pkce_verifier {
form.push(("code_verifier", verifier.to_string()));
}
request
.header("Content-Type", "application/x-www-form-urlencoded")
.header("Accept", "application/json")
.form(&form)
.send()
.await
}
.map_err(|_| {
build_internal_control_error_response(http::StatusCode::BAD_REQUEST, "token exchange 失败")
})?;
if !response.status().is_success() {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 失败",
));
}
let payload = response.json::<serde_json::Value>().await.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
)
})?;
if json_non_empty_string(payload.get("access_token")).is_none() {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token exchange 返回缺少 access_token",
));
}
Ok(payload)
}
pub(crate) async fn exchange_admin_provider_oauth_refresh_token(
state: &AppState,
template: ProviderOAuthTemplate,
refresh_token: &str,
) -> Result<serde_json::Value, Response<Body>> {
let token_url = state.provider_oauth_token_url(template.provider_type, template.token_url);
let request = state.client.post(token_url);
let scope = template.scopes.join(" ");
let response = if template.provider_type == "claude_code" {
let mut body = serde_json::Map::from_iter([
(
"grant_type".to_string(),
serde_json::Value::String("refresh_token".to_string()),
),
(
"client_id".to_string(),
serde_json::Value::String(template.client_id.to_string()),
),
(
"refresh_token".to_string(),
serde_json::Value::String(refresh_token.to_string()),
),
]);
if !scope.trim().is_empty() {
body.insert("scope".to_string(), serde_json::Value::String(scope));
}
request
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.json(&serde_json::Value::Object(body))
.send()
.await
} else {
let mut form = vec![
("grant_type", "refresh_token".to_string()),
("client_id", template.client_id.to_string()),
("refresh_token", refresh_token.to_string()),
];
if !scope.trim().is_empty() {
form.push(("scope", scope));
}
if !template.client_secret.trim().is_empty() {
form.push(("client_secret", template.client_secret.to_string()));
}
request
.header("Content-Type", "application/x-www-form-urlencoded")
.header("Accept", "application/json")
.form(&form)
.send()
.await
}
.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Refresh Token 验证失败: token exchange 失败",
)
})?;
let status = response.status();
let body = response.text().await.map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"Refresh Token 验证失败: token exchange 失败",
)
})?;
if !status.is_success() {
let reason =
normalize_provider_oauth_refresh_error_message(Some(status.as_u16()), Some(&body));
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
format!("Refresh Token 验证失败: {reason}"),
));
}
let payload = serde_json::from_str::<serde_json::Value>(&body).map_err(|_| {
build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token refresh 返回缺少 access_token",
)
})?;
if json_non_empty_string(payload.get("access_token")).is_none() {
return Err(build_internal_control_error_response(
http::StatusCode::BAD_REQUEST,
"token refresh 返回缺少 access_token",
));
}
Ok(payload)
}
pub(crate) fn build_provider_oauth_start_response(
template: ProviderOAuthTemplate,
nonce: &str,
code_challenge: Option<&str>,
) -> serde_json::Value {
let mut serializer = form_urlencoded::Serializer::new(String::new());
serializer.append_pair("client_id", template.client_id);
serializer.append_pair("response_type", "code");
serializer.append_pair("redirect_uri", template.redirect_uri);
serializer.append_pair("scope", &template.scopes.join(" "));
serializer.append_pair("state", nonce);
if template.provider_type == "codex" {
serializer.append_pair("prompt", "login");
serializer.append_pair("id_token_add_organizations", "true");
serializer.append_pair("codex_cli_simplified_flow", "true");
}
if template.use_pkce {
if let Some(code_challenge) = code_challenge {
serializer.append_pair("code_challenge", code_challenge);
serializer.append_pair("code_challenge_method", "S256");
}
}
json!({
"authorization_url": format!("{}?{}", template.authorize_url, serializer.finish()),
"redirect_uri": template.redirect_uri,
"provider_type": template.provider_type,
"instructions": "1) 打开 authorization_url 完成授权\n2) 授权后会跳转到 redirect_uri(localhost)\n3) 复制浏览器地址栏完整 URL,调用 complete 接口粘贴 callback_url",
})
}
pub(crate) async fn save_provider_oauth_batch_task_payload(
state: &AppState,
task_id: &str,
task_state: &serde_json::Value,
) -> Result<(), GatewayError> {
let key = provider_oauth_batch_task_storage_key(task_id);
let serialized =
serde_json::to_string(task_state).map_err(|err| GatewayError::Internal(err.to_string()))?;
if let Some(runner) = state.redis_kv_runner() {
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
return Err(GatewayError::Internal(
"provider oauth batch task redis unavailable".to_string(),
));
};
let redis_key = runner.keyspace().key(&key);
redis::cmd("SET")
.arg(redis_key)
.arg(&serialized)
.arg("EX")
.arg(PROVIDER_OAUTH_BATCH_TASK_TTL_SECS)
.query_async::<()>(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(());
}
if state.save_provider_oauth_batch_task_for_tests(&key, &serialized) {
return Ok(());
}
Err(GatewayError::Internal(
"provider oauth batch task redis unavailable".to_string(),
))
}
pub(crate) async fn read_provider_oauth_batch_task_payload(
state: &AppState,
provider_id: &str,
task_id: &str,
) -> Result<Option<serde_json::Value>, GatewayError> {
let key = provider_oauth_batch_task_storage_key(task_id);
let raw = if let Some(runner) = state.redis_kv_runner() {
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
return Err(GatewayError::Internal(
"provider oauth batch task redis unavailable".to_string(),
));
};
let redis_key = runner.keyspace().key(&key);
redis::cmd("GET")
.arg(redis_key)
.query_async(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
} else {
state.load_provider_oauth_batch_task_for_tests(&key)
};
let Some(raw) = raw else {
return Ok(None);
};
let parsed = match serde_json::from_str::<serde_json::Value>(&raw) {
Ok(value) => value,
Err(_) => return Ok(None),
};
let Some(state) = parsed.as_object() else {
return Ok(None);
};
if state
.get("provider_id")
.and_then(serde_json::Value::as_str)
.unwrap_or_default()
!= provider_id
{
return Ok(None);
}
Ok(Some(build_provider_oauth_batch_task_status_payload(
provider_id,
state,
)))
}
File diff suppressed because one or more lines are too long
@@ -0,0 +1,83 @@
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_ops_architecture_id_from_path, is_admin_provider_ops_architectures_root,
};
use crate::GatewayError;
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::{json, Value};
static ADMIN_PROVIDER_OPS_ARCHITECTURES_ALL: std::sync::LazyLock<Vec<Value>> =
std::sync::LazyLock::new(|| {
serde_json::from_str(include_str!("architectures.all.json"))
.expect("admin provider ops architectures fixture should parse")
});
fn admin_provider_ops_architectures_list_payload() -> Vec<Value> {
ADMIN_PROVIDER_OPS_ARCHITECTURES_ALL
.iter()
.filter(|item| item.get("architecture_id").and_then(Value::as_str) != Some("generic_api"))
.cloned()
.collect()
}
fn admin_provider_ops_architecture_payload(architecture_id: &str) -> Option<Value> {
ADMIN_PROVIDER_OPS_ARCHITECTURES_ALL
.iter()
.find_map(|item| {
(item.get("architecture_id").and_then(Value::as_str) == Some(architecture_id))
.then(|| item.clone())
})
}
pub(super) async fn maybe_build_local_admin_provider_ops_architectures_response(
request_context: &GatewayPublicRequestContext,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.control_decision.as_ref() else {
return Ok(None);
};
if decision.route_family.as_deref() == Some("provider_ops_manage")
&& decision.route_kind.as_deref() == Some("list_architectures")
&& request_context.request_method == http::Method::GET
&& is_admin_provider_ops_architectures_root(&request_context.request_path)
{
return Ok(Some(
Json(admin_provider_ops_architectures_list_payload()).into_response(),
));
}
if decision.route_family.as_deref() == Some("provider_ops_manage")
&& decision.route_kind.as_deref() == Some("get_architecture")
&& request_context.request_method == http::Method::GET
{
let Some(architecture_id) =
admin_provider_ops_architecture_id_from_path(&request_context.request_path)
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "架构不存在" })),
)
.into_response(),
));
};
return Ok(Some(
match admin_provider_ops_architecture_payload(&architecture_id) {
Some(payload) => Json(payload).into_response(),
None => (
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": format!("架构 {architecture_id} 不存在") })),
)
.into_response(),
},
));
}
Ok(None)
}
@@ -0,0 +1,34 @@
use crate::control::GatewayPublicRequestContext;
use crate::{AppState, GatewayError};
use axum::body::{Body, Bytes};
use axum::http::Response;
mod architectures;
mod providers;
pub(crate) use self::providers::admin_provider_ops_local_action_response;
pub(crate) async fn maybe_build_local_admin_provider_ops_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
if let Some(response) =
architectures::maybe_build_local_admin_provider_ops_architectures_response(request_context)
.await?
{
return Ok(Some(response));
}
if let Some(response) = providers::maybe_build_local_admin_provider_ops_providers_response(
state,
request_context,
request_body,
)
.await?
{
return Ok(Some(response));
}
Ok(None)
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,384 @@
use super::{AdminProviderOpsSaveConfigRequest, ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS};
use crate::handlers::admin::shared::{
decrypt_catalog_secret_with_fallbacks, encrypt_catalog_secret_with_fallbacks,
};
use crate::AppState;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
};
use serde_json::json;
pub(super) fn admin_provider_ops_config_object(
provider: &StoredProviderCatalogProvider,
) -> Option<&serde_json::Map<String, serde_json::Value>> {
provider
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|config| config.get("provider_ops"))
.and_then(serde_json::Value::as_object)
}
pub(super) fn admin_provider_ops_connector_object(
provider_ops_config: &serde_json::Map<String, serde_json::Value>,
) -> Option<&serde_json::Map<String, serde_json::Value>> {
provider_ops_config
.get("connector")
.and_then(serde_json::Value::as_object)
}
fn admin_provider_ops_masked_secret(
state: &AppState,
field: &str,
ciphertext: &str,
) -> serde_json::Value {
let plaintext = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)
.unwrap_or_else(|| ciphertext.to_string());
if plaintext.is_empty() {
return serde_json::Value::String(String::new());
}
let masked = if field == "password" {
"********".to_string()
} else if plaintext.len() > 12 {
format!(
"{}****{}",
&plaintext[..4],
&plaintext[plaintext.len().saturating_sub(4)..]
)
} else if plaintext.len() > 8 {
format!(
"{}****{}",
&plaintext[..2],
&plaintext[plaintext.len().saturating_sub(2)..]
)
} else {
"*".repeat(plaintext.len())
};
serde_json::Value::String(masked)
}
fn admin_provider_ops_masked_credentials(
state: &AppState,
raw_credentials: Option<&serde_json::Value>,
) -> serde_json::Value {
let Some(credentials) = raw_credentials.and_then(serde_json::Value::as_object) else {
return json!({});
};
let mut masked = serde_json::Map::new();
for (key, value) in credentials {
if ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS.contains(&key.as_str()) {
if let Some(ciphertext) = value.as_str().filter(|value| !value.is_empty()) {
masked.insert(
key.clone(),
admin_provider_ops_masked_secret(state, key, ciphertext),
);
continue;
}
}
masked.insert(key.clone(), value.clone());
}
serde_json::Value::Object(masked)
}
fn admin_provider_ops_is_supported_auth_type(auth_type: &str) -> bool {
matches!(
auth_type,
"api_key" | "session_login" | "oauth" | "cookie" | "none"
)
}
pub(super) fn admin_provider_ops_uses_python_verify_fallback(
architecture_id: &str,
config: &serde_json::Map<String, serde_json::Value>,
) -> bool {
let _ = architecture_id;
config
.get("proxy_enabled")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
|| config
.get("proxy_node_id")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.is_some_and(|value| !value.is_empty())
}
pub(super) fn admin_provider_ops_decrypted_credentials(
state: &AppState,
raw_credentials: Option<&serde_json::Value>,
) -> serde_json::Map<String, serde_json::Value> {
let Some(credentials) = raw_credentials.and_then(serde_json::Value::as_object) else {
return serde_json::Map::new();
};
let mut decrypted = serde_json::Map::new();
for (key, value) in credentials {
if ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS.contains(&key.as_str()) {
if let Some(ciphertext) = value.as_str() {
let plaintext =
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)
.unwrap_or_else(|| ciphertext.to_string());
decrypted.insert(key.clone(), serde_json::Value::String(plaintext));
continue;
}
}
decrypted.insert(key.clone(), value.clone());
}
decrypted
}
fn admin_provider_ops_sensitive_placeholder_or_empty(value: Option<&serde_json::Value>) -> bool {
match value {
None | Some(serde_json::Value::Null) => true,
Some(serde_json::Value::String(raw)) => raw.is_empty() || raw.chars().all(|ch| ch == '*'),
Some(serde_json::Value::Array(items)) => items.is_empty(),
Some(serde_json::Value::Object(map)) => map.is_empty(),
_ => false,
}
}
pub(super) fn admin_provider_ops_merge_credentials(
state: &AppState,
provider: &StoredProviderCatalogProvider,
mut request_credentials: serde_json::Map<String, serde_json::Value>,
) -> serde_json::Map<String, serde_json::Value> {
let saved_credentials = admin_provider_ops_decrypted_credentials(
state,
admin_provider_ops_config_object(provider)
.and_then(admin_provider_ops_connector_object)
.and_then(|connector| connector.get("credentials")),
);
for field in ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS {
if admin_provider_ops_sensitive_placeholder_or_empty(request_credentials.get(*field))
&& saved_credentials.contains_key(*field)
{
if let Some(saved_value) = saved_credentials.get(*field) {
request_credentials.insert((*field).to_string(), saved_value.clone());
}
}
}
for (key, value) in saved_credentials {
if key.starts_with('_') && !request_credentials.contains_key(&key) {
request_credentials.insert(key, value);
}
}
request_credentials
}
fn admin_provider_ops_encrypt_credentials(
state: &AppState,
credentials: serde_json::Map<String, serde_json::Value>,
) -> Result<serde_json::Map<String, serde_json::Value>, String> {
let mut encrypted = serde_json::Map::new();
for (key, value) in credentials {
if ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS.contains(&key.as_str()) {
if let Some(plaintext) = value.as_str() {
if plaintext.is_empty() {
encrypted.insert(key, value);
} else {
let ciphertext = encrypt_catalog_secret_with_fallbacks(state, plaintext)
.ok_or_else(|| "gateway 未配置 Provider Ops 加密密钥".to_string())?;
encrypted.insert(key, serde_json::Value::String(ciphertext));
}
continue;
}
}
encrypted.insert(key, value);
}
Ok(encrypted)
}
pub(super) fn build_admin_provider_ops_saved_config_value(
state: &AppState,
provider: &StoredProviderCatalogProvider,
payload: AdminProviderOpsSaveConfigRequest,
) -> Result<serde_json::Value, String> {
let auth_type = payload.connector.auth_type.trim().to_string();
if auth_type.is_empty() || !admin_provider_ops_is_supported_auth_type(auth_type.as_str()) {
return Err("connector.auth_type 必须是合法的认证类型".to_string());
}
let merged_credentials =
admin_provider_ops_merge_credentials(state, provider, payload.connector.credentials);
let encrypted_credentials = admin_provider_ops_encrypt_credentials(state, merged_credentials)?;
let actions = payload
.actions
.into_iter()
.map(|(action_type, config)| {
(
action_type,
json!({
"enabled": config.enabled,
"config": config.config,
}),
)
})
.collect::<serde_json::Map<String, serde_json::Value>>();
Ok(json!({
"architecture_id": payload.architecture_id,
"base_url": payload.base_url,
"connector": {
"auth_type": auth_type,
"config": payload.connector.config,
"credentials": encrypted_credentials,
},
"actions": actions,
"schedule": payload.schedule,
}))
}
pub(super) fn resolve_admin_provider_ops_base_url(
provider: &StoredProviderCatalogProvider,
endpoints: &[StoredProviderCatalogEndpoint],
provider_ops_config: Option<&serde_json::Map<String, serde_json::Value>>,
) -> Option<String> {
let from_saved_config = provider_ops_config
.and_then(|config| config.get("base_url"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
if from_saved_config.is_some() {
return from_saved_config;
}
if let Some(base_url) = endpoints.iter().find_map(|endpoint| {
let value = endpoint.base_url.trim();
(!value.is_empty()).then(|| value.to_string())
}) {
return Some(base_url);
}
let from_provider_config = provider
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|config| config.get("base_url"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
if from_provider_config.is_some() {
return from_provider_config;
}
provider
.website
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
pub(super) fn build_admin_provider_ops_status_payload(
provider_id: &str,
provider: Option<&StoredProviderCatalogProvider>,
) -> serde_json::Value {
let provider_ops_config = provider.and_then(admin_provider_ops_config_object);
let auth_type = provider_ops_config
.and_then(admin_provider_ops_connector_object)
.and_then(|connector| connector.get("auth_type"))
.and_then(serde_json::Value::as_str)
.unwrap_or_else(|| {
if provider_ops_config.is_some() {
"api_key"
} else {
"none"
}
});
let mut enabled_actions = provider_ops_config
.and_then(|config| config.get("actions"))
.and_then(serde_json::Value::as_object)
.map(|actions| {
actions
.iter()
.filter_map(|(action_type, config)| {
let enabled = config
.as_object()
.and_then(|config| config.get("enabled"))
.and_then(serde_json::Value::as_bool)
.unwrap_or(true);
enabled.then(|| serde_json::Value::String(action_type.clone()))
})
.collect::<Vec<_>>()
})
.unwrap_or_default();
enabled_actions.sort_by(|left, right| left.as_str().cmp(&right.as_str()));
json!({
"provider_id": provider_id,
"is_configured": provider_ops_config.is_some(),
"architecture_id": provider_ops_config.map(|config| {
config
.get("architecture_id")
.and_then(serde_json::Value::as_str)
.unwrap_or("generic_api")
}),
"connection_status": {
"status": "disconnected",
"auth_type": auth_type,
"connected_at": serde_json::Value::Null,
"expires_at": serde_json::Value::Null,
"last_error": serde_json::Value::Null,
},
"enabled_actions": enabled_actions,
})
}
pub(super) fn build_admin_provider_ops_config_payload(
state: &AppState,
provider_id: &str,
provider: Option<&StoredProviderCatalogProvider>,
endpoints: &[StoredProviderCatalogEndpoint],
) -> serde_json::Value {
let Some(provider) = provider else {
return json!({
"provider_id": provider_id,
"is_configured": false,
});
};
let Some(provider_ops_config) = admin_provider_ops_config_object(provider) else {
return json!({
"provider_id": provider_id,
"is_configured": false,
});
};
let connector = admin_provider_ops_connector_object(provider_ops_config);
json!({
"provider_id": provider_id,
"is_configured": true,
"architecture_id": provider_ops_config
.get("architecture_id")
.and_then(serde_json::Value::as_str)
.unwrap_or("generic_api"),
"base_url": resolve_admin_provider_ops_base_url(
provider,
endpoints,
Some(provider_ops_config),
),
"connector": {
"auth_type": connector
.and_then(|connector| connector.get("auth_type"))
.and_then(serde_json::Value::as_str)
.unwrap_or("api_key"),
"config": connector
.and_then(|connector| connector.get("config"))
.filter(|value| value.is_object())
.cloned()
.unwrap_or_else(|| json!({})),
"credentials": admin_provider_ops_masked_credentials(
state,
connector.and_then(|connector| connector.get("credentials")),
),
},
})
}
@@ -0,0 +1,117 @@
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_id_for_provider_ops_balance, admin_provider_id_for_provider_ops_checkin,
admin_provider_id_for_provider_ops_config, admin_provider_id_for_provider_ops_connect,
admin_provider_id_for_provider_ops_disconnect, admin_provider_id_for_provider_ops_status,
admin_provider_id_for_provider_ops_verify, admin_provider_ops_action_route_parts,
};
use crate::handlers::admin::shared::{
decrypt_catalog_secret_with_fallbacks, encrypt_catalog_secret_with_fallbacks,
};
use crate::{AppState, GatewayError};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde::Deserialize;
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
mod actions;
mod config;
mod routes;
mod verify;
use self::actions::admin_provider_ops_is_valid_action_type;
pub(crate) use self::actions::admin_provider_ops_local_action_response;
use self::config::{
admin_provider_ops_config_object, admin_provider_ops_connector_object,
admin_provider_ops_decrypted_credentials, admin_provider_ops_merge_credentials,
admin_provider_ops_uses_python_verify_fallback, build_admin_provider_ops_config_payload,
build_admin_provider_ops_saved_config_value, build_admin_provider_ops_status_payload,
resolve_admin_provider_ops_base_url,
};
pub(super) use self::routes::maybe_build_local_admin_provider_ops_providers_response;
use self::verify::{
admin_provider_ops_local_verify_response, admin_provider_ops_normalized_verify_architecture_id,
admin_provider_ops_value_as_f64, admin_provider_ops_verify_failure,
admin_provider_ops_verify_headers,
};
const ADMIN_PROVIDER_OPS_SENSITIVE_FIELDS: &[&str] = &[
"api_key",
"password",
"refresh_token",
"session_token",
"session_cookie",
"token_cookie",
"auth_cookie",
"cookie_string",
"cookie",
];
const ADMIN_PROVIDER_OPS_CONNECT_RUST_ONLY_MESSAGE: &str =
"Provider 连接仅支持 Rust execution runtime";
const ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE: &str =
"Provider 操作仅支持 Rust execution runtime";
const ADMIN_PROVIDER_OPS_VERIFY_RUST_ONLY_MESSAGE: &str = "认证验证仅支持 Rust execution runtime";
#[derive(Debug, Deserialize)]
struct AdminProviderOpsSaveConfigRequest {
#[serde(default = "default_admin_provider_ops_architecture_id")]
architecture_id: String,
#[serde(default)]
base_url: Option<String>,
connector: AdminProviderOpsConnectorConfigRequest,
#[serde(default)]
actions: BTreeMap<String, AdminProviderOpsActionConfigRequest>,
#[serde(default)]
schedule: BTreeMap<String, String>,
}
#[derive(Debug, Deserialize)]
struct AdminProviderOpsConnectorConfigRequest {
auth_type: String,
#[serde(default)]
config: serde_json::Map<String, serde_json::Value>,
#[serde(default)]
credentials: serde_json::Map<String, serde_json::Value>,
}
#[derive(Debug, Deserialize)]
struct AdminProviderOpsActionConfigRequest {
#[serde(default = "default_admin_provider_ops_action_enabled")]
enabled: bool,
#[serde(default)]
config: serde_json::Map<String, serde_json::Value>,
}
#[derive(Debug, Deserialize)]
struct AdminProviderOpsConnectRequest {
#[serde(default)]
credentials: Option<serde_json::Map<String, serde_json::Value>>,
}
#[derive(Debug, Deserialize)]
struct AdminProviderOpsExecuteActionRequest {
#[serde(default)]
config: Option<serde_json::Map<String, serde_json::Value>>,
}
#[derive(Debug, Clone)]
struct AdminProviderOpsCheckinOutcome {
success: Option<bool>,
message: String,
cookie_expired: bool,
}
fn default_admin_provider_ops_architecture_id() -> String {
"generic_api".to_string()
}
fn default_admin_provider_ops_action_enabled() -> bool {
true
}
@@ -0,0 +1,697 @@
use super::{
admin_provider_ops_config_object, admin_provider_ops_connector_object,
admin_provider_ops_decrypted_credentials, admin_provider_ops_is_valid_action_type,
admin_provider_ops_local_action_response, admin_provider_ops_local_verify_response,
admin_provider_ops_merge_credentials, admin_provider_ops_normalized_verify_architecture_id,
admin_provider_ops_verify_failure, build_admin_provider_ops_config_payload,
build_admin_provider_ops_saved_config_value, build_admin_provider_ops_status_payload,
resolve_admin_provider_ops_base_url, AdminProviderOpsConnectRequest,
AdminProviderOpsExecuteActionRequest, AdminProviderOpsSaveConfigRequest,
ADMIN_PROVIDER_OPS_CONNECT_RUST_ONLY_MESSAGE,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_id_for_provider_ops_balance, admin_provider_id_for_provider_ops_checkin,
admin_provider_id_for_provider_ops_config, admin_provider_id_for_provider_ops_connect,
admin_provider_id_for_provider_ops_disconnect, admin_provider_id_for_provider_ops_status,
admin_provider_id_for_provider_ops_verify, admin_provider_ops_action_route_parts,
};
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
pub(crate) async fn maybe_build_local_admin_provider_ops_providers_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.control_decision.as_ref() else {
return Ok(None);
};
if decision.route_family.as_deref() != Some("provider_ops_manage") {
return Ok(None);
}
let route_kind = decision.route_kind.as_deref().unwrap_or_default();
if !state.has_provider_catalog_data_reader() && route_kind != "disconnect_provider" {
return Ok(None);
}
if route_kind == "batch_balance" {
let requested_provider_ids = match request_body {
Some(body) if !body.is_empty() => {
let raw_value = match serde_json::from_slice::<serde_json::Value>(body) {
Ok(value) => value,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是 provider_id 数组" })),
)
.into_response(),
));
}
};
let ids = if let Some(items) = raw_value.as_array() {
items
.iter()
.filter_map(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.collect::<Vec<_>>()
} else if let Some(items) = raw_value
.get("provider_ids")
.and_then(serde_json::Value::as_array)
{
items
.iter()
.filter_map(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.collect::<Vec<_>>()
} else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是 provider_id 数组" })),
)
.into_response(),
));
};
Some(ids)
}
_ => None,
};
let provider_ids = if let Some(provider_ids) = requested_provider_ids {
provider_ids
} else {
state
.list_provider_catalog_providers(true)
.await?
.into_iter()
.filter(|provider| {
provider
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.is_some_and(|config| config.contains_key("provider_ops"))
})
.map(|provider| provider.id)
.collect::<Vec<_>>()
};
if provider_ids.is_empty() {
return Ok(Some(Json(json!({})).into_response()));
}
let providers = state
.read_provider_catalog_providers_by_ids(&provider_ids)
.await?;
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(&provider_ids)
.await?;
let mut payload = serde_json::Map::new();
for provider_id in &provider_ids {
let provider = providers
.iter()
.find(|provider| provider.id == *provider_id);
let provider_endpoints = endpoints
.iter()
.filter(|endpoint| endpoint.provider_id == *provider_id)
.cloned()
.collect::<Vec<_>>();
let result = admin_provider_ops_local_action_response(
state,
provider_id,
provider,
&provider_endpoints,
"query_balance",
None,
)
.await;
payload.insert(provider_id.clone(), result);
}
return Ok(Some(
Json(serde_json::Value::Object(payload)).into_response(),
));
}
let action_route = if route_kind == "execute_provider_action" {
admin_provider_ops_action_route_parts(&request_context.request_path)
} else {
None
};
let provider_id = if matches!(
route_kind,
"get_provider_status"
| "get_provider_config"
| "save_provider_config"
| "delete_provider_config"
| "verify_provider"
| "connect_provider"
| "disconnect_provider"
| "get_provider_balance"
| "refresh_provider_balance"
| "provider_checkin"
| "execute_provider_action"
) {
admin_provider_id_for_provider_ops_config(&request_context.request_path)
.or_else(|| admin_provider_id_for_provider_ops_status(&request_context.request_path))
.or_else(|| admin_provider_id_for_provider_ops_verify(&request_context.request_path))
.or_else(|| admin_provider_id_for_provider_ops_connect(&request_context.request_path))
.or_else(|| admin_provider_id_for_provider_ops_balance(&request_context.request_path))
.or_else(|| admin_provider_id_for_provider_ops_checkin(&request_context.request_path))
.or_else(|| {
action_route
.as_ref()
.map(|(provider_id, _)| provider_id.clone())
})
.or_else(|| {
admin_provider_id_for_provider_ops_disconnect(&request_context.request_path)
})
} else {
None
};
let Some(provider_id) = provider_id else {
return Ok(None);
};
if decision.route_kind.as_deref() != Some(route_kind) {
return Ok(None);
}
if route_kind == "save_provider_config" {
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
let raw_value = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(value) => value,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
if !raw_value.is_object() {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
let payload = match serde_json::from_value::<AdminProviderOpsSaveConfigRequest>(raw_value) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let Some(existing_provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let provider_ops_config =
match build_admin_provider_ops_saved_config_value(state, &existing_provider, payload) {
Ok(config) => config,
Err(detail) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response(),
));
}
};
let mut updated_provider = existing_provider.clone();
let mut provider_config = updated_provider
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
provider_config.insert("provider_ops".to_string(), provider_ops_config);
updated_provider.config = Some(serde_json::Value::Object(provider_config));
updated_provider.updated_at_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs());
let Some(_updated) = state
.update_provider_catalog_provider(&updated_provider)
.await?
else {
return Ok(None);
};
return Ok(Some(
Json(json!({
"success": true,
"message": "配置保存成功",
}))
.into_response(),
));
}
if route_kind == "verify_provider" {
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
let raw_value = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(value) => value,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
if !raw_value.is_object() {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
let payload = match serde_json::from_value::<AdminProviderOpsSaveConfigRequest>(raw_value) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let existing_provider = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next();
let endpoints = if existing_provider.is_some() {
state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?
} else {
Vec::new()
};
let base_url = payload
.base_url
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.or_else(|| {
existing_provider.as_ref().and_then(|provider| {
resolve_admin_provider_ops_base_url(
provider,
&endpoints,
admin_provider_ops_config_object(provider),
)
})
});
let Some(base_url) = base_url else {
return Ok(Some(
Json(admin_provider_ops_verify_failure("请提供 API 地址")).into_response(),
));
};
let architecture_id =
admin_provider_ops_normalized_verify_architecture_id(&payload.architecture_id);
let credentials = existing_provider.as_ref().map_or_else(
|| payload.connector.credentials.clone(),
|provider| {
admin_provider_ops_merge_credentials(
state,
provider,
payload.connector.credentials.clone(),
)
},
);
let payload = admin_provider_ops_local_verify_response(
state,
&base_url,
architecture_id,
&payload.connector.config,
&credentials,
)
.await;
return Ok(Some(attach_admin_audit_response(
Json(payload).into_response(),
"admin_provider_ops_config_verified",
"verify_provider_ops_config",
"provider",
&provider_id,
)));
}
if route_kind == "connect_provider" {
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
let raw_value = match serde_json::from_slice::<serde_json::Value>(request_body) {
Ok(value) => value,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
if !raw_value.is_object() {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
let payload = match serde_json::from_value::<AdminProviderOpsConnectRequest>(raw_value) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let Some(existing_provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let Some(provider_ops_config) = admin_provider_ops_config_object(&existing_provider) else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "未配置操作设置" })),
)
.into_response(),
));
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?;
if resolve_admin_provider_ops_base_url(
&existing_provider,
&endpoints,
Some(provider_ops_config),
)
.is_none()
{
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "Provider 未配置 base_url" })),
)
.into_response(),
));
}
let actual_credentials = payload
.credentials
.filter(|value| !value.is_empty())
.unwrap_or_else(|| {
admin_provider_ops_decrypted_credentials(
state,
admin_provider_ops_config_object(&existing_provider)
.and_then(admin_provider_ops_connector_object)
.and_then(|connector| connector.get("credentials")),
)
});
if actual_credentials.is_empty() {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "未提供凭据" })),
)
.into_response(),
));
}
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": ADMIN_PROVIDER_OPS_CONNECT_RUST_ONLY_MESSAGE })),
)
.into_response(),
));
}
if matches!(
route_kind,
"get_provider_balance"
| "refresh_provider_balance"
| "provider_checkin"
| "execute_provider_action"
) {
let action_type = if route_kind == "provider_checkin" {
"checkin".to_string()
} else if matches!(
route_kind,
"get_provider_balance" | "refresh_provider_balance"
) {
"query_balance".to_string()
} else {
let Some((_, action_type)) = action_route else {
return Ok(None);
};
if !admin_provider_ops_is_valid_action_type(&action_type) {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": format!("无效的操作类型: {action_type}") })),
)
.into_response(),
));
}
action_type
};
let request_config = if route_kind == "execute_provider_action" {
match request_body {
Some(body) if !body.is_empty() => {
let raw_value = match serde_json::from_slice::<serde_json::Value>(body) {
Ok(value) => value,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
let payload = match serde_json::from_value::<AdminProviderOpsExecuteActionRequest>(
raw_value,
) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体必须是合法的 JSON 对象" })),
)
.into_response(),
));
}
};
payload.config
}
_ => None,
}
} else {
None
};
let providers = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?;
let provider = providers.first();
let endpoints = if provider.is_some() {
state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?
} else {
Vec::new()
};
let payload = admin_provider_ops_local_action_response(
state,
&provider_id,
provider,
&endpoints,
&action_type,
request_config.as_ref(),
)
.await;
return Ok(Some(Json(payload).into_response()));
}
if route_kind == "delete_provider_config" {
let Some(existing_provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(Some(
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider 不存在" })),
)
.into_response(),
));
};
let mut updated_provider = existing_provider.clone();
let mut provider_config = updated_provider
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
if provider_config.remove("provider_ops").is_some() {
updated_provider.config = Some(serde_json::Value::Object(provider_config));
updated_provider.updated_at_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs());
let Some(_updated) = state
.update_provider_catalog_provider(&updated_provider)
.await?
else {
return Ok(None);
};
}
return Ok(Some(
Json(json!({
"success": true,
"message": "配置已删除",
}))
.into_response(),
));
}
if route_kind == "disconnect_provider" {
return Ok(Some(
Json(json!({
"success": true,
"message": "已断开连接",
}))
.into_response(),
));
}
let providers = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?;
let provider = providers.first();
let endpoints = if route_kind == "get_provider_config" && provider.is_some() {
state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
.await?
} else {
Vec::new()
};
let payload = if route_kind == "get_provider_status" {
build_admin_provider_ops_status_payload(&provider_id, provider)
} else {
build_admin_provider_ops_config_payload(state, &provider_id, provider, &endpoints)
};
let response = Json(payload).into_response();
let response = if route_kind == "get_provider_config" {
attach_admin_audit_response(
response,
"admin_provider_ops_config_viewed",
"view_provider_ops_config",
"provider",
&provider_id,
)
} else if route_kind == "get_provider_status" {
attach_admin_audit_response(
response,
"admin_provider_ops_status_viewed",
"view_provider_ops_status",
"provider",
&provider_id,
)
} else {
response
};
Ok(Some(response))
}
@@ -0,0 +1,910 @@
use super::ADMIN_PROVIDER_OPS_VERIFY_RUST_ONLY_MESSAGE;
use crate::AppState;
use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine as _};
use regex::Regex;
use serde_json::json;
pub(super) fn admin_provider_ops_normalized_verify_architecture_id(architecture_id: &str) -> &str {
match architecture_id.trim() {
"" => "generic_api",
"generic_api" | "new_api" | "cubence" | "yescode" | "nekocode" | "anyrouter"
| "sub2api" => architecture_id.trim(),
_ => "generic_api",
}
}
fn admin_provider_ops_extract_cookie_value(cookie_input: &str, key: &str) -> String {
if cookie_input.contains(&format!("{key}=")) {
for part in cookie_input.split(';') {
let trimmed = part.trim();
if let Some(value) = trimmed.strip_prefix(&format!("{key}=")) {
return value.trim().to_string();
}
}
}
cookie_input.trim().to_string()
}
fn admin_provider_ops_yescode_cookie_header(cookie_input: &str) -> String {
if cookie_input.contains("yescode_auth=") {
let mut parts = Vec::new();
for part in cookie_input.split(';') {
let trimmed = part.trim();
if let Some(value) = trimmed.strip_prefix("yescode_auth=") {
parts.push(format!("yescode_auth={}", value.trim()));
} else if let Some(value) = trimmed.strip_prefix("yescode_csrf=") {
parts.push(format!("yescode_csrf={}", value.trim()));
}
}
return parts.join("; ");
}
format!("yescode_auth={}", cookie_input.trim())
}
const ADMIN_PROVIDER_OPS_ANYROUTER_XOR_KEY: &str = "3000176000856006061501533003690027800375";
const ADMIN_PROVIDER_OPS_ANYROUTER_UNSBOX_TABLE: [usize; 40] = [
0xF, 0x23, 0x1D, 0x18, 0x21, 0x10, 0x1, 0x26, 0xA, 0x9, 0x13, 0x1F, 0x28, 0x1B, 0x16, 0x17,
0x19, 0xD, 0x6, 0xB, 0x27, 0x12, 0x14, 0x8, 0xE, 0x15, 0x20, 0x1A, 0x2, 0x1E, 0x7, 0x4, 0x11,
0x5, 0x3, 0x1C, 0x22, 0x25, 0xC, 0x24,
];
fn admin_provider_ops_anyrouter_compute_acw_sc_v2(arg1: &str) -> Option<String> {
if arg1.len() != 40 || !arg1.chars().all(|ch| ch.is_ascii_hexdigit()) {
return None;
}
let chars = arg1.chars().collect::<Vec<_>>();
let unsboxed = ADMIN_PROVIDER_OPS_ANYROUTER_UNSBOX_TABLE
.iter()
.map(|index| chars.get(index.saturating_sub(1)).copied())
.collect::<Option<String>>()?;
let mut result = String::with_capacity(40);
for i in (0..40).step_by(2) {
let a = u8::from_str_radix(&unsboxed[i..i + 2], 16).ok()?;
let b = u8::from_str_radix(&ADMIN_PROVIDER_OPS_ANYROUTER_XOR_KEY[i..i + 2], 16).ok()?;
result.push_str(&format!("{:02x}", a ^ b));
}
Some(result)
}
fn admin_provider_ops_anyrouter_parse_session_user_id(cookie_input: &str) -> Option<String> {
let session_cookie = admin_provider_ops_extract_cookie_value(cookie_input, "session");
let decoded = URL_SAFE_NO_PAD.decode(session_cookie.as_bytes()).ok()?;
let text = String::from_utf8_lossy(&decoded);
let mut parts = text.split('|');
let _timestamp = parts.next()?;
let gob_b64 = parts.next()?;
let gob_data = URL_SAFE_NO_PAD.decode(gob_b64.as_bytes()).ok()?;
let id_pattern = b"\x02id\x03int";
let id_idx = gob_data
.windows(id_pattern.len())
.position(|window| window == id_pattern)?;
let value_start = id_idx + id_pattern.len() + 2;
let first_byte = *gob_data.get(value_start)?;
if first_byte != 0 {
return None;
}
let marker = *gob_data.get(value_start + 1)?;
if marker < 0x80 {
return None;
}
let length = 256usize.saturating_sub(marker as usize);
let end = value_start + 2 + length;
let bytes = gob_data.get(value_start + 2..end)?;
let val = bytes
.iter()
.fold(0u64, |acc, byte| (acc << 8) | (*byte as u64));
Some((val >> 1).to_string())
}
async fn admin_provider_ops_anyrouter_acw_cookie(
state: &AppState,
base_url: &str,
) -> Option<String> {
let response = state
.client
.get(base_url.trim_end_matches('/'))
.header(
reqwest::header::USER_AGENT,
"Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36",
)
.send()
.await
.ok()?;
let body = response.text().await.ok()?;
let compiled = Regex::new(r"var\s+arg1\s*=\s*'([0-9a-fA-F]{40})'").ok()?;
let captures = compiled.captures(&body)?;
let arg1 = captures.get(1)?.as_str();
admin_provider_ops_anyrouter_compute_acw_sc_v2(arg1).map(|value| format!("acw_sc__v2={value}"))
}
pub(super) fn admin_provider_ops_verify_failure(message: impl Into<String>) -> serde_json::Value {
json!({
"success": false,
"message": message.into(),
})
}
fn admin_provider_ops_verify_success(
data: serde_json::Value,
updated_credentials: Option<serde_json::Map<String, serde_json::Value>>,
) -> serde_json::Value {
let mut payload = serde_json::Map::from_iter([
("success".to_string(), serde_json::Value::Bool(true)),
("data".to_string(), data),
]);
if let Some(credentials) = updated_credentials.filter(|value| !value.is_empty()) {
payload.insert(
"updated_credentials".to_string(),
serde_json::Value::Object(credentials),
);
}
serde_json::Value::Object(payload)
}
fn admin_provider_ops_verify_user_payload(
username: Option<String>,
display_name: Option<String>,
email: Option<String>,
quota: Option<f64>,
extra: Option<serde_json::Map<String, serde_json::Value>>,
) -> serde_json::Value {
let resolved_username = username.filter(|value| !value.trim().is_empty());
let resolved_display_name = display_name
.filter(|value| !value.trim().is_empty())
.or_else(|| resolved_username.clone());
let mut payload = serde_json::Map::new();
payload.insert(
"username".to_string(),
resolved_username
.map(serde_json::Value::String)
.unwrap_or(serde_json::Value::Null),
);
payload.insert(
"display_name".to_string(),
resolved_display_name
.map(serde_json::Value::String)
.unwrap_or(serde_json::Value::Null),
);
payload.insert(
"email".to_string(),
email
.map(serde_json::Value::String)
.unwrap_or(serde_json::Value::Null),
);
payload.insert(
"quota".to_string(),
quota
.and_then(serde_json::Number::from_f64)
.map(serde_json::Value::Number)
.unwrap_or(serde_json::Value::Null),
);
if let Some(extra) = extra.filter(|value| !value.is_empty()) {
payload.insert("extra".to_string(), serde_json::Value::Object(extra));
}
serde_json::Value::Object(payload)
}
pub(super) fn admin_provider_ops_value_as_f64(value: Option<&serde_json::Value>) -> Option<f64> {
match value {
Some(serde_json::Value::Number(number)) => number.as_f64(),
Some(serde_json::Value::String(raw)) => raw.trim().parse::<f64>().ok(),
_ => None,
}
}
fn admin_provider_ops_json_object(
value: &serde_json::Value,
) -> Option<&serde_json::Map<String, serde_json::Value>> {
value.as_object()
}
fn admin_provider_ops_frontend_updated_credentials(
credentials: serde_json::Map<String, serde_json::Value>,
) -> Option<serde_json::Map<String, serde_json::Value>> {
let filtered = credentials
.into_iter()
.filter(|(key, value)| {
!key.starts_with('_')
&& !matches!(value, serde_json::Value::Null)
&& !value.as_str().is_some_and(|raw| raw.trim().is_empty())
})
.collect::<serde_json::Map<String, serde_json::Value>>();
(!filtered.is_empty()).then_some(filtered)
}
fn admin_provider_ops_generic_verify_payload(
status: http::StatusCode,
response_json: &serde_json::Value,
) -> serde_json::Value {
if status == http::StatusCode::UNAUTHORIZED {
return admin_provider_ops_verify_failure("认证失败:无效的凭据");
}
if status == http::StatusCode::FORBIDDEN {
return admin_provider_ops_verify_failure("认证失败:权限不足");
}
if status != http::StatusCode::OK {
return admin_provider_ops_verify_failure(format!("验证失败:HTTP {}", status.as_u16()));
}
let user_data = if response_json
.get("success")
.and_then(serde_json::Value::as_bool)
== Some(true)
&& response_json
.get("data")
.is_some_and(serde_json::Value::is_object)
{
response_json.get("data")
} else if response_json
.get("success")
.and_then(serde_json::Value::as_bool)
== Some(false)
{
return admin_provider_ops_verify_failure(
response_json
.get("message")
.and_then(serde_json::Value::as_str)
.unwrap_or("验证失败"),
);
} else {
Some(response_json)
};
let Some(user_data) = user_data.and_then(admin_provider_ops_json_object) else {
return admin_provider_ops_verify_failure("响应格式无效");
};
let mut extra = serde_json::Map::new();
for (key, value) in user_data {
if matches!(
key.as_str(),
"username" | "display_name" | "email" | "quota" | "used_quota" | "request_count"
) {
continue;
}
extra.insert(key.clone(), value.clone());
}
admin_provider_ops_verify_success(
admin_provider_ops_verify_user_payload(
user_data
.get("username")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned),
user_data
.get("display_name")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned),
user_data
.get("email")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned),
admin_provider_ops_value_as_f64(user_data.get("quota")),
Some(extra),
),
None,
)
}
fn admin_provider_ops_cubence_verify_payload(
status: http::StatusCode,
response_json: &serde_json::Value,
) -> serde_json::Value {
if status == http::StatusCode::UNAUTHORIZED {
return admin_provider_ops_verify_failure("Cookie 已失效,请重新配置");
}
if status == http::StatusCode::FORBIDDEN {
return admin_provider_ops_verify_failure("Cookie 已失效或无权限");
}
if status != http::StatusCode::OK {
return admin_provider_ops_verify_failure(format!("验证失败:HTTP {}", status.as_u16()));
}
let Some(payload) = admin_provider_ops_json_object(response_json) else {
return admin_provider_ops_verify_failure("响应格式无效");
};
let user_info = payload
.get("user")
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
let balance_info = payload
.get("balance")
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
let mut extra = serde_json::Map::new();
if let Some(role) = user_info.get("role") {
extra.insert("role".to_string(), role.clone());
}
if let Some(invite_code) = user_info.get("invite_code") {
extra.insert("invite_code".to_string(), invite_code.clone());
}
admin_provider_ops_verify_success(
admin_provider_ops_verify_user_payload(
user_info
.get("username")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned),
user_info
.get("username")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned),
None,
admin_provider_ops_value_as_f64(balance_info.get("total_balance_dollar")),
Some(extra),
),
None,
)
}
fn admin_provider_ops_yescode_verify_payload(
status: http::StatusCode,
response_json: &serde_json::Value,
) -> serde_json::Value {
if status == http::StatusCode::UNAUTHORIZED {
return admin_provider_ops_verify_failure("Cookie 已失效,请重新配置");
}
if status == http::StatusCode::FORBIDDEN {
return admin_provider_ops_verify_failure("Cookie 已失效或无权限");
}
if status != http::StatusCode::OK {
return admin_provider_ops_verify_failure(format!("验证失败:HTTP {}", status.as_u16()));
}
let Some(payload) = admin_provider_ops_json_object(response_json) else {
return admin_provider_ops_verify_failure("响应格式无效");
};
let Some(username) = payload
.get("username")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned)
else {
return admin_provider_ops_verify_failure("响应格式无效");
};
let pay_as_you_go =
admin_provider_ops_value_as_f64(payload.get("pay_as_you_go_balance")).unwrap_or(0.0);
let subscription =
admin_provider_ops_value_as_f64(payload.get("subscription_balance")).unwrap_or(0.0);
let plan = payload
.get("subscription_plan")
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default();
let weekly_limit = admin_provider_ops_value_as_f64(
payload
.get("weekly_limit")
.or_else(|| plan.get("weekly_limit")),
);
let weekly_spent = admin_provider_ops_value_as_f64(
payload
.get("weekly_spent_balance")
.or_else(|| payload.get("current_week_spend")),
)
.unwrap_or(0.0);
let subscription_available = weekly_limit
.map(|limit| (limit - weekly_spent).max(0.0).min(subscription))
.unwrap_or(subscription);
admin_provider_ops_verify_success(
admin_provider_ops_verify_user_payload(
Some(username.clone()),
Some(username),
payload
.get("email")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned),
Some(pay_as_you_go + subscription_available),
None,
),
None,
)
}
fn admin_provider_ops_nekocode_verify_payload(
status: http::StatusCode,
response_json: &serde_json::Value,
) -> serde_json::Value {
if status == http::StatusCode::UNAUTHORIZED {
return admin_provider_ops_verify_failure("Cookie 已失效,请重新配置");
}
if status == http::StatusCode::FORBIDDEN {
return admin_provider_ops_verify_failure("Cookie 已失效或无权限");
}
if status != http::StatusCode::OK {
return admin_provider_ops_verify_failure(format!("验证失败:HTTP {}", status.as_u16()));
}
let user_data = if response_json
.get("success")
.and_then(serde_json::Value::as_bool)
== Some(true)
&& response_json
.get("data")
.is_some_and(serde_json::Value::is_object)
{
response_json.get("data")
} else {
Some(response_json)
};
let Some(user_data) = user_data.and_then(admin_provider_ops_json_object) else {
return admin_provider_ops_verify_failure("响应格式无效");
};
admin_provider_ops_verify_success(
admin_provider_ops_verify_user_payload(
user_data
.get("username")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned),
user_data
.get("display_name")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned),
user_data
.get("email")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned),
admin_provider_ops_value_as_f64(user_data.get("balance")),
None,
),
None,
)
}
async fn admin_provider_ops_sub2api_exchange_token(
state: &AppState,
base_url: &str,
credentials: &serde_json::Map<String, serde_json::Value>,
) -> Result<(String, Option<serde_json::Map<String, serde_json::Value>>), String> {
let email = credentials
.get("email")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.unwrap_or_default();
let password = credentials
.get("password")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.unwrap_or_default();
let refresh_token = credentials
.get("refresh_token")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.unwrap_or_default();
let (path, body, default_error, previous_refresh_token) =
if !email.is_empty() && !password.is_empty() {
(
"/api/v1/auth/login",
json!({ "email": email, "password": password }),
"登录失败",
None,
)
} else if !refresh_token.is_empty() {
(
"/api/v1/auth/refresh",
json!({ "refresh_token": refresh_token }),
"Refresh Token 无效或已过期",
Some(refresh_token),
)
} else {
return Err("请填写账号密码或 Refresh Token".to_string());
};
let response = match state
.client
.post(format!("{}{path}", base_url.trim_end_matches('/')))
.json(&body)
.send()
.await
{
Ok(response) => response,
Err(err) if err.is_timeout() => return Err("连接超时".to_string()),
Err(err) if err.is_connect() => return Err(format!("连接失败: {err}")),
Err(err) => return Err(format!("验证失败: {err}")),
};
let status = response.status();
let response_json = match response.bytes().await {
Ok(bytes) => {
serde_json::from_slice::<serde_json::Value>(&bytes).unwrap_or_else(|_| json!({}))
}
Err(_) => json!({}),
};
let payload = response_json.as_object().cloned().unwrap_or_default();
if status != http::StatusCode::OK
|| payload
.get("code")
.and_then(serde_json::Value::as_i64)
.unwrap_or(-1)
!= 0
{
let message = payload
.get("message")
.and_then(serde_json::Value::as_str)
.unwrap_or(default_error);
return Err(message.to_string());
}
let Some(token_data) = payload.get("data").and_then(serde_json::Value::as_object) else {
return Err("响应格式无效".to_string());
};
let access_token = token_data
.get("access_token")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.ok_or_else(|| "响应格式无效".to_string())?;
let mut updated_credentials = serde_json::Map::new();
if let Some(new_refresh_token) = token_data
.get("refresh_token")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
if previous_refresh_token != Some(new_refresh_token) {
updated_credentials.insert(
"refresh_token".to_string(),
serde_json::Value::String(new_refresh_token.to_string()),
);
}
}
Ok((
access_token.to_string(),
admin_provider_ops_frontend_updated_credentials(updated_credentials),
))
}
fn admin_provider_ops_sub2api_verify_payload(
status: http::StatusCode,
response_json: &serde_json::Value,
updated_credentials: Option<serde_json::Map<String, serde_json::Value>>,
) -> serde_json::Value {
if status == http::StatusCode::UNAUTHORIZED {
return admin_provider_ops_verify_failure("认证失败:无效的凭据");
}
if status == http::StatusCode::FORBIDDEN {
return admin_provider_ops_verify_failure("认证失败:权限不足");
}
if status != http::StatusCode::OK {
return admin_provider_ops_verify_failure(format!("验证失败:HTTP {}", status.as_u16()));
}
let Some(payload) = admin_provider_ops_json_object(response_json) else {
return admin_provider_ops_verify_failure("响应格式无效");
};
if payload
.get("code")
.and_then(serde_json::Value::as_i64)
.unwrap_or(-1)
!= 0
{
return admin_provider_ops_verify_failure(
payload
.get("message")
.and_then(serde_json::Value::as_str)
.unwrap_or("验证失败"),
);
}
let Some(user_data) = payload.get("data").and_then(serde_json::Value::as_object) else {
return admin_provider_ops_verify_failure("响应格式无效");
};
let balance = admin_provider_ops_value_as_f64(user_data.get("balance")).unwrap_or(0.0);
let points = admin_provider_ops_value_as_f64(user_data.get("points")).unwrap_or(0.0);
let mut extra = serde_json::Map::new();
for key in ["balance", "points", "status", "concurrency"] {
if let Some(value) = user_data.get(key) {
extra.insert(key.to_string(), value.clone());
}
}
admin_provider_ops_verify_success(
admin_provider_ops_verify_user_payload(
user_data
.get("username")
.or_else(|| user_data.get("email"))
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned),
user_data
.get("username")
.or_else(|| user_data.get("email"))
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned),
user_data
.get("email")
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned),
Some(balance + points),
Some(extra),
),
updated_credentials,
)
}
async fn admin_provider_ops_local_sub2api_verify_response(
state: &AppState,
base_url: &str,
credentials: &serde_json::Map<String, serde_json::Value>,
) -> serde_json::Value {
let base_url = base_url.trim().trim_end_matches('/');
if base_url.is_empty() {
return admin_provider_ops_verify_failure("请提供 API 地址");
}
let (access_token, updated_credentials) =
match admin_provider_ops_sub2api_exchange_token(state, base_url, credentials).await {
Ok(value) => value,
Err(message) => return admin_provider_ops_verify_failure(message),
};
let response = match state
.client
.get(format!("{base_url}/api/v1/auth/me?timezone=Asia/Shanghai"))
.bearer_auth(access_token)
.send()
.await
{
Ok(response) => response,
Err(err) if err.is_timeout() => return admin_provider_ops_verify_failure("连接超时"),
Err(err) if err.is_connect() => {
return admin_provider_ops_verify_failure(format!("连接失败: {err}"));
}
Err(err) => return admin_provider_ops_verify_failure(format!("验证失败: {err}")),
};
let status = response.status();
let response_json = match response.bytes().await {
Ok(bytes) => {
serde_json::from_slice::<serde_json::Value>(&bytes).unwrap_or_else(|_| json!({}))
}
Err(_) => json!({}),
};
admin_provider_ops_sub2api_verify_payload(status, &response_json, updated_credentials)
}
fn admin_provider_ops_insert_header(
headers: &mut reqwest::header::HeaderMap,
name: &str,
value: &str,
) -> Result<(), String> {
let header_name = reqwest::header::HeaderName::from_bytes(name.as_bytes())
.map_err(|_| format!("无效的请求头: {name}"))?;
let header_value = reqwest::header::HeaderValue::from_str(value)
.map_err(|_| format!("无效的请求头值: {name}"))?;
headers.insert(header_name, header_value);
Ok(())
}
pub(super) fn admin_provider_ops_verify_headers(
architecture_id: &str,
config: &serde_json::Map<String, serde_json::Value>,
credentials: &serde_json::Map<String, serde_json::Value>,
) -> Result<reqwest::header::HeaderMap, String> {
let mut headers = reqwest::header::HeaderMap::new();
match architecture_id {
"generic_api" => {
let api_key = credentials
.get("api_key")
.and_then(serde_json::Value::as_str)
.unwrap_or_default()
.trim();
if !api_key.is_empty() {
let auth_method = config
.get("auth_method")
.and_then(serde_json::Value::as_str)
.unwrap_or("bearer");
if auth_method == "header" {
let header_name = config
.get("header_name")
.and_then(serde_json::Value::as_str)
.unwrap_or("X-API-Key");
admin_provider_ops_insert_header(&mut headers, header_name, api_key)?;
} else {
admin_provider_ops_insert_header(
&mut headers,
"Authorization",
&format!("Bearer {api_key}"),
)?;
}
}
}
"new_api" => {
for (name, value) in [
(
"User-Agent",
"Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/140.0.7339.249 Electron/38.7.0 Safari/537.36",
),
("Accept", "application/json"),
("Accept-Encoding", "gzip, deflate, br"),
("Accept-Language", "zh-CN"),
("sec-ch-ua", "\"Not=A?Brand\";v=\"24\", \"Chromium\";v=\"140\""),
("sec-ch-ua-mobile", "?0"),
("sec-ch-ua-platform", "\"macOS\""),
("Sec-Fetch-Site", "cross-site"),
("Sec-Fetch-Mode", "cors"),
("Sec-Fetch-Dest", "empty"),
] {
admin_provider_ops_insert_header(&mut headers, name, value)?;
}
if let Some(api_key) = credentials
.get("api_key")
.and_then(serde_json::Value::as_str)
{
if !api_key.trim().is_empty() {
admin_provider_ops_insert_header(
&mut headers,
"Authorization",
&format!("Bearer {}", api_key.trim()),
)?;
}
}
if let Some(user_id) = credentials
.get("user_id")
.and_then(serde_json::Value::as_str)
{
if !user_id.trim().is_empty() {
admin_provider_ops_insert_header(&mut headers, "New-Api-User", user_id.trim())?;
}
}
if let Some(cookie) = credentials
.get("cookie")
.and_then(serde_json::Value::as_str)
{
if !cookie.trim().is_empty() {
admin_provider_ops_insert_header(&mut headers, "Cookie", cookie.trim())?;
}
}
}
"cubence" => {
if let Some(token_cookie) = credentials
.get("token_cookie")
.and_then(serde_json::Value::as_str)
.filter(|value| !value.trim().is_empty())
{
let token = admin_provider_ops_extract_cookie_value(token_cookie, "token");
admin_provider_ops_insert_header(
&mut headers,
"Cookie",
&format!("token={token}"),
)?;
}
}
"yescode" => {
if let Some(auth_cookie) = credentials
.get("auth_cookie")
.and_then(serde_json::Value::as_str)
.filter(|value| !value.trim().is_empty())
{
admin_provider_ops_insert_header(
&mut headers,
"Cookie",
&admin_provider_ops_yescode_cookie_header(auth_cookie),
)?;
}
}
"nekocode" => {
if let Some(session_cookie) = credentials
.get("session_cookie")
.and_then(serde_json::Value::as_str)
.filter(|value| !value.trim().is_empty())
{
let session = admin_provider_ops_extract_cookie_value(session_cookie, "session");
admin_provider_ops_insert_header(
&mut headers,
"Cookie",
&format!("session={session}"),
)?;
}
}
"anyrouter" => {
let mut cookies = Vec::new();
if let Some(acw_cookie) = config
.get("acw_cookie")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
cookies.push(acw_cookie.to_string());
}
if let Some(session_cookie) = credentials
.get("session_cookie")
.and_then(serde_json::Value::as_str)
.filter(|value| !value.trim().is_empty())
{
let session = admin_provider_ops_extract_cookie_value(session_cookie, "session");
cookies.push(format!("session={session}"));
if let Some(user_id) =
admin_provider_ops_anyrouter_parse_session_user_id(session_cookie)
{
admin_provider_ops_insert_header(&mut headers, "New-Api-User", user_id.trim())?;
}
}
if !cookies.is_empty() {
admin_provider_ops_insert_header(&mut headers, "Cookie", &cookies.join("; "))?;
}
}
_ => {}
}
Ok(headers)
}
pub(super) async fn admin_provider_ops_local_verify_response(
state: &AppState,
base_url: &str,
architecture_id: &str,
config: &serde_json::Map<String, serde_json::Value>,
credentials: &serde_json::Map<String, serde_json::Value>,
) -> serde_json::Value {
if architecture_id == "sub2api" {
return admin_provider_ops_local_sub2api_verify_response(state, base_url, credentials)
.await;
}
let mut resolved_config = config.clone();
if architecture_id == "anyrouter" {
if let Some(acw_cookie) = admin_provider_ops_anyrouter_acw_cookie(state, base_url).await {
resolved_config.insert(
"acw_cookie".to_string(),
serde_json::Value::String(acw_cookie),
);
}
}
let verify_path = match architecture_id {
"anyrouter" => "/api/user/self",
"cubence" => "/api/v1/dashboard/overview",
"yescode" => "/api/v1/auth/profile",
"nekocode" => "/api/user/self",
"new_api" | "generic_api" => "/api/user/self",
_ => return admin_provider_ops_verify_failure(ADMIN_PROVIDER_OPS_VERIFY_RUST_ONLY_MESSAGE),
};
let base_url = base_url.trim().trim_end_matches('/');
if base_url.is_empty() {
return admin_provider_ops_verify_failure("请提供 API 地址");
}
let headers =
match admin_provider_ops_verify_headers(architecture_id, &resolved_config, credentials) {
Ok(headers) => headers,
Err(message) => return admin_provider_ops_verify_failure(message),
};
let response = match state
.client
.get(format!("{base_url}{verify_path}"))
.headers(headers)
.send()
.await
{
Ok(response) => response,
Err(err) if err.is_timeout() => return admin_provider_ops_verify_failure("连接超时"),
Err(err) if err.is_connect() => {
return admin_provider_ops_verify_failure(format!("连接失败: {err}"));
}
Err(err) => return admin_provider_ops_verify_failure(format!("验证失败: {err}")),
};
let status = response.status();
let response_json = match response.bytes().await {
Ok(bytes) => {
serde_json::from_slice::<serde_json::Value>(&bytes).unwrap_or_else(|_| json!({}))
}
Err(_) => json!({}),
};
match architecture_id {
"cubence" => admin_provider_ops_cubence_verify_payload(status, &response_json),
"yescode" => admin_provider_ops_yescode_verify_payload(status, &response_json),
"nekocode" => admin_provider_ops_nekocode_verify_payload(status, &response_json),
_ => admin_provider_ops_generic_verify_payload(status, &response_json),
}
}
@@ -0,0 +1,485 @@
use super::shared::{
AdminProviderPoolConfig, AdminProviderPoolRuntimeState, ADMIN_PROVIDER_POOL_SCAN_BATCH,
};
use crate::{AppState, GatewayError};
use aether_data::redis::{RedisKeyspace, RedisKvRunner};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider;
use serde_json::json;
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
use tracing::warn;
fn json_u64(value: &serde_json::Value) -> Option<u64> {
value
.as_u64()
.or_else(|| value.as_i64().and_then(|raw| u64::try_from(raw).ok()))
}
fn admin_provider_pool_lru_enabled(
raw_pool_advanced: &serde_json::Map<String, serde_json::Value>,
) -> bool {
if let Some(explicit) = raw_pool_advanced
.get("lru_enabled")
.and_then(serde_json::Value::as_bool)
{
return explicit;
}
let Some(presets) = raw_pool_advanced
.get("scheduling_presets")
.and_then(serde_json::Value::as_array)
else {
return false;
};
let Some(first) = presets.first() else {
return false;
};
if first.is_string() {
return raw_pool_advanced
.get("lru_enabled")
.and_then(serde_json::Value::as_bool)
.unwrap_or(true);
}
presets
.iter()
.filter_map(serde_json::Value::as_object)
.any(|item| {
item.get("preset")
.and_then(serde_json::Value::as_str)
.is_some_and(|preset| preset.eq_ignore_ascii_case("lru"))
&& item
.get("enabled")
.and_then(serde_json::Value::as_bool)
.unwrap_or(true)
})
}
pub(crate) fn admin_provider_pool_config(
provider: &StoredProviderCatalogProvider,
) -> Option<AdminProviderPoolConfig> {
let raw_pool_advanced = provider
.config
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|config| config.get("pool_advanced"))?;
let Some(pool_advanced) = raw_pool_advanced.as_object() else {
return Some(AdminProviderPoolConfig {
lru_enabled: false,
cost_window_seconds: 18_000,
cost_limit_per_key_tokens: None,
});
};
Some(AdminProviderPoolConfig {
lru_enabled: admin_provider_pool_lru_enabled(pool_advanced),
cost_window_seconds: pool_advanced
.get("cost_window_seconds")
.and_then(json_u64)
.filter(|value| *value > 0)
.unwrap_or(18_000),
cost_limit_per_key_tokens: pool_advanced
.get("cost_limit_per_key_tokens")
.and_then(json_u64),
})
}
fn pool_sticky_pattern(keyspace: &RedisKeyspace, provider_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:sticky:*"))
}
fn pool_lru_key(keyspace: &RedisKeyspace, provider_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:lru"))
}
fn pool_cooldown_key(keyspace: &RedisKeyspace, provider_id: &str, key_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:cooldown:{key_id}"))
}
fn pool_cooldown_index_key(keyspace: &RedisKeyspace, provider_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:cooldown_idx"))
}
fn pool_cost_key(keyspace: &RedisKeyspace, provider_id: &str, key_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:cost:{key_id}"))
}
fn parse_pool_cost_member(member: &str) -> u64 {
member
.rsplit_once(':')
.and_then(|(_, suffix)| suffix.parse::<u64>().ok())
.unwrap_or(0)
}
async fn scan_redis_keys(
connection: &mut redis::aio::MultiplexedConnection,
pattern: &str,
) -> Result<Vec<String>, GatewayError> {
let mut cursor = 0u64;
let mut keys = Vec::new();
loop {
let (next_cursor, batch): (u64, Vec<String>) = redis::cmd("SCAN")
.arg(cursor)
.arg("MATCH")
.arg(pattern)
.arg("COUNT")
.arg(ADMIN_PROVIDER_POOL_SCAN_BATCH)
.query_async(connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
keys.extend(batch);
if next_cursor == 0 {
break;
}
cursor = next_cursor;
}
Ok(keys)
}
fn pool_cooldown_keys(
keyspace: &RedisKeyspace,
provider_id: &str,
key_ids: &[String],
) -> Vec<String> {
key_ids
.iter()
.map(|key_id| pool_cooldown_key(keyspace, provider_id, key_id))
.collect()
}
fn pool_cost_keys(keyspace: &RedisKeyspace, provider_id: &str, key_ids: &[String]) -> Vec<String> {
key_ids
.iter()
.map(|key_id| pool_cost_key(keyspace, provider_id, key_id))
.collect()
}
pub(crate) async fn read_admin_provider_pool_cooldown_counts(
runner: &RedisKvRunner,
provider_ids: &[String],
) -> BTreeMap<String, usize> {
if provider_ids.is_empty() {
return BTreeMap::new();
}
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
warn!("gateway admin provider pool: failed to connect redis for cooldown counts");
return BTreeMap::new();
};
let keyspace = runner.keyspace().clone();
let mut pipeline = redis::pipe();
for provider_id in provider_ids {
pipeline
.cmd("SCARD")
.arg(pool_cooldown_index_key(&keyspace, provider_id));
}
match pipeline.query_async::<Vec<u64>>(&mut connection).await {
Ok(counts) => provider_ids
.iter()
.cloned()
.zip(counts.into_iter())
.map(|(provider_id, count)| (provider_id, count as usize))
.collect(),
Err(err) => {
warn!(
"gateway admin provider pool: failed to batch read cooldown counts: {:?}",
err
);
BTreeMap::new()
}
}
}
pub(crate) async fn read_admin_provider_pool_runtime_state(
runner: &RedisKvRunner,
provider_id: &str,
key_ids: &[String],
pool_config: AdminProviderPoolConfig,
) -> AdminProviderPoolRuntimeState {
let mut runtime = AdminProviderPoolRuntimeState::default();
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
warn!("gateway admin provider pool: failed to connect redis for provider {provider_id}");
return runtime;
};
let keyspace = runner.keyspace().clone();
let cooldown_keys = pool_cooldown_keys(&keyspace, provider_id, key_ids);
let cost_keys = pool_cost_keys(&keyspace, provider_id, key_ids);
let sticky_keys = match scan_redis_keys(
&mut connection,
&pool_sticky_pattern(&keyspace, provider_id),
)
.await
{
Ok(keys) => keys,
Err(err) => {
warn!(
"gateway admin provider pool: failed to scan sticky keys for provider {provider_id}: {:?}",
err
);
Vec::new()
}
};
runtime.total_sticky_sessions = sticky_keys.len();
if !sticky_keys.is_empty() {
for chunk in sticky_keys.chunks(ADMIN_PROVIDER_POOL_SCAN_BATCH as usize) {
let values = redis::cmd("MGET")
.arg(chunk)
.query_async::<Vec<Option<String>>>(&mut connection)
.await;
let Ok(values) = values else {
warn!("gateway admin provider pool: failed to read sticky bindings for provider {provider_id}");
break;
};
for bound_key_id in values.into_iter().flatten() {
*runtime
.sticky_sessions_by_key
.entry(bound_key_id)
.or_insert(0) += 1;
}
}
}
if !cooldown_keys.is_empty() {
let cooldown_reasons = redis::cmd("MGET")
.arg(&cooldown_keys)
.query_async::<Vec<Option<String>>>(&mut connection)
.await
.unwrap_or_else(|err| {
warn!(
"gateway admin provider pool: failed to batch read cooldown reasons for provider {provider_id}: {:?}",
err
);
vec![None; cooldown_keys.len()]
});
let mut ttl_pipeline = redis::pipe();
for cooldown_key in &cooldown_keys {
ttl_pipeline.cmd("TTL").arg(cooldown_key);
}
let cooldown_ttls = ttl_pipeline
.query_async::<Vec<i64>>(&mut connection)
.await
.unwrap_or_else(|err| {
warn!(
"gateway admin provider pool: failed to batch read cooldown ttl for provider {provider_id}: {:?}",
err
);
vec![-1; cooldown_keys.len()]
});
for (((key_id, _cooldown_key), reason), ttl) in key_ids
.iter()
.zip(cooldown_keys.iter())
.zip(cooldown_reasons.into_iter())
.zip(cooldown_ttls.into_iter())
{
if let Some(reason) = reason {
runtime
.cooldown_reason_by_key
.insert(key_id.clone(), reason);
if let Ok(ttl_seconds) = u64::try_from(ttl) {
if ttl_seconds > 0 {
runtime
.cooldown_ttl_by_key
.insert(key_id.clone(), ttl_seconds);
}
}
}
}
}
if !cost_keys.is_empty() {
let window_start = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
.saturating_sub(pool_config.cost_window_seconds);
let mut cost_pipeline = redis::pipe();
for cost_key in &cost_keys {
cost_pipeline
.cmd("ZRANGEBYSCORE")
.arg(cost_key)
.arg(window_start)
.arg("+inf");
}
let members_by_key = cost_pipeline
.query_async::<Vec<Vec<String>>>(&mut connection)
.await
.unwrap_or_else(|err| {
warn!(
"gateway admin provider pool: failed to batch read cost windows for provider {provider_id}: {:?}",
err
);
vec![Vec::new(); cost_keys.len()]
});
for (key_id, members) in key_ids.iter().zip(members_by_key.into_iter()) {
let total = members
.iter()
.map(|member| parse_pool_cost_member(member))
.sum::<u64>();
runtime
.cost_window_usage_by_key
.insert(key_id.clone(), total);
}
}
if pool_config.lru_enabled && !key_ids.is_empty() {
let mut command = redis::cmd("ZMSCORE");
command.arg(pool_lru_key(&keyspace, provider_id));
for key_id in key_ids {
command.arg(key_id);
}
if let Ok(scores) = command
.query_async::<Vec<Option<f64>>>(&mut connection)
.await
{
for (key_id, score) in key_ids.iter().zip(scores.into_iter()) {
if let Some(score) = score {
runtime.lru_score_by_key.insert(key_id.clone(), score);
}
}
}
}
runtime
}
pub(crate) async fn read_admin_provider_pool_cooldown_count(
runner: &RedisKvRunner,
provider_id: &str,
) -> usize {
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
warn!("gateway admin provider pool: failed to connect redis for provider {provider_id}");
return 0;
};
let keyspace = runner.keyspace().clone();
redis::cmd("SCARD")
.arg(pool_cooldown_index_key(&keyspace, provider_id))
.query_async::<u64>(&mut connection)
.await
.map(|value| value as usize)
.unwrap_or(0)
}
pub(crate) async fn read_admin_provider_pool_cooldown_key_ids(
runner: &RedisKvRunner,
provider_id: &str,
) -> Vec<String> {
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
warn!("gateway admin provider pool: failed to connect redis for provider {provider_id}");
return Vec::new();
};
let keyspace = runner.keyspace().clone();
redis::cmd("SMEMBERS")
.arg(pool_cooldown_index_key(&keyspace, provider_id))
.query_async::<Vec<String>>(&mut connection)
.await
.unwrap_or_default()
}
pub(crate) async fn build_admin_provider_pool_status_payload(
state: &AppState,
provider_id: &str,
) -> Option<serde_json::Value> {
if !state.has_provider_catalog_data_reader() {
return None;
}
let provider = state
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
.await
.ok()
.and_then(|mut providers| providers.drain(..).next())?;
let Some(pool_config) = admin_provider_pool_config(&provider) else {
return Some(json!({
"provider_id": provider.id,
"provider_name": provider.name,
"pool_enabled": false,
"total_keys": 0,
"total_sticky_sessions": 0,
"keys": [],
}));
};
let keys = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
.await
.ok()
.unwrap_or_default();
let key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>();
let runtime = match state.redis_kv_runner() {
Some(runner) => {
read_admin_provider_pool_runtime_state(&runner, &provider.id, &key_ids, pool_config)
.await
}
None => AdminProviderPoolRuntimeState::default(),
};
let key_payloads = keys
.into_iter()
.map(|key| {
let cooldown_reason = runtime.cooldown_reason_by_key.get(&key.id).cloned();
json!({
"key_id": key.id,
"key_name": key.name,
"is_active": key.is_active,
"cooldown_reason": cooldown_reason,
"cooldown_ttl_seconds": cooldown_reason
.as_ref()
.and_then(|_| runtime.cooldown_ttl_by_key.get(&key.id).copied()),
"cost_window_usage": runtime.cost_window_usage_by_key.get(&key.id).copied().unwrap_or(0),
"cost_limit": pool_config.cost_limit_per_key_tokens,
"sticky_sessions": runtime.sticky_sessions_by_key.get(&key.id).copied().unwrap_or(0),
"lru_score": runtime.lru_score_by_key.get(&key.id).copied(),
})
})
.collect::<Vec<_>>();
Some(json!({
"provider_id": provider.id,
"provider_name": provider.name,
"pool_enabled": true,
"total_keys": key_payloads.len(),
"total_sticky_sessions": runtime.total_sticky_sessions,
"keys": key_payloads,
}))
}
pub(crate) async fn clear_admin_provider_pool_cooldown(
state: &AppState,
provider_id: &str,
key_id: &str,
) {
let Some(runner) = state.redis_kv_runner() else {
return;
};
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
warn!("gateway admin provider pool: failed to connect redis to clear cooldown for key {key_id}");
return;
};
let keyspace = runner.keyspace().clone();
let _: Result<(), _> = redis::pipe()
.cmd("DEL")
.arg(pool_cooldown_key(&keyspace, provider_id, key_id))
.ignore()
.cmd("SREM")
.arg(pool_cooldown_index_key(&keyspace, provider_id))
.arg(key_id)
.ignore()
.query_async(&mut connection)
.await;
}
pub(crate) async fn reset_admin_provider_pool_cost(
state: &AppState,
provider_id: &str,
key_id: &str,
) {
let Some(runner) = state.redis_kv_runner() else {
return;
};
let _ = runner.del(&format!("ap:{provider_id}:cost:{key_id}")).await;
}
@@ -0,0 +1,623 @@
use super::{
admin_pool_provider_id_from_path, build_admin_pool_error_response,
ADMIN_POOL_BANNED_KEY_CLEANUP_EMPTY_MESSAGE,
ADMIN_POOL_PROVIDER_CATALOG_READER_UNAVAILABLE_DETAIL,
ADMIN_POOL_PROVIDER_CATALOG_WRITER_UNAVAILABLE_DETAIL,
};
use super::{payloads as pool_payloads, selection as pool_selection};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::pool::{
clear_admin_provider_pool_cooldown, reset_admin_provider_pool_cost,
};
use crate::handlers::admin::shared::{
attach_admin_audit_response, encrypt_catalog_secret_with_fallbacks,
};
use crate::{AppState, GatewayError, LocalProviderDeleteTaskState};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::collections::BTreeSet;
use std::time::{SystemTime, UNIX_EPOCH};
use uuid::Uuid;
#[derive(Debug, Default, serde::Deserialize)]
struct AdminPoolBatchActionRequest {
#[serde(default)]
key_ids: Vec<String>,
#[serde(default)]
action: String,
#[serde(default)]
payload: Option<serde_json::Value>,
}
#[derive(Debug, Default, serde::Deserialize)]
struct AdminPoolBatchImportRequest {
#[serde(default)]
keys: Vec<AdminPoolBatchImportItem>,
#[serde(default)]
proxy_node_id: Option<String>,
}
#[derive(Debug, Default, serde::Deserialize)]
struct AdminPoolBatchImportItem {
#[serde(default)]
name: String,
#[serde(default)]
api_key: String,
#[serde(default)]
auth_type: String,
}
fn admin_pool_batch_delete_task_parts(request_path: &str) -> Option<(String, String)> {
let raw = request_path.strip_prefix("/api/admin/pool/")?;
let (provider_id, suffix) = raw.split_once("/keys/batch-delete-task/")?;
let provider_id = provider_id.trim();
let task_id = suffix.trim().trim_matches('/');
if provider_id.is_empty()
|| provider_id.contains('/')
|| task_id.is_empty()
|| task_id.contains('/')
{
return None;
}
Some((provider_id.to_string(), task_id.to_string()))
}
fn admin_pool_key_proxy_value(proxy_node_id: Option<&str>) -> Option<serde_json::Value> {
proxy_node_id
.map(str::trim)
.filter(|value| !value.is_empty())
.map(|value| json!({ "node_id": value, "enabled": true }))
}
fn build_admin_pool_batch_delete_task_payload(
task: &LocalProviderDeleteTaskState,
) -> serde_json::Value {
json!({
"task_id": task.task_id,
"provider_id": task.provider_id,
"status": task.status,
"stage": task.stage,
"total_keys": task.total_keys,
"deleted_keys": task.deleted_keys,
"total_endpoints": task.total_endpoints,
"deleted_endpoints": task.deleted_endpoints,
"message": task.message,
})
}
fn attach_admin_pool_batch_delete_task_terminal_audit(
provider_id: &str,
task_id: &str,
task_status: &str,
response: Response<Body>,
) -> Response<Body> {
match task_status {
"completed" => attach_admin_audit_response(
response,
"admin_pool_batch_delete_task_completed_viewed",
"view_pool_batch_delete_task_terminal_state",
"provider_key_batch_delete_task",
&format!("{provider_id}:{task_id}"),
),
"failed" => attach_admin_audit_response(
response,
"admin_pool_batch_delete_task_failed_viewed",
"view_pool_batch_delete_task_terminal_state",
"provider_key_batch_delete_task",
&format!("{provider_id}:{task_id}"),
),
_ => response,
}
}
fn admin_pool_resolved_api_formats(
endpoints: &[StoredProviderCatalogEndpoint],
existing_keys: &[StoredProviderCatalogKey],
) -> Vec<String> {
let mut formats = Vec::new();
let mut seen = BTreeSet::new();
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
let api_format = endpoint.api_format.trim();
if api_format.is_empty() || !seen.insert(api_format.to_string()) {
continue;
}
formats.push(api_format.to_string());
}
if !formats.is_empty() {
return formats;
}
for key in existing_keys {
for api_format in pool_payloads::admin_pool_api_formats(key) {
if !seen.insert(api_format.clone()) {
continue;
}
formats.push(api_format);
}
}
formats
}
async fn build_admin_pool_cleanup_banned_keys_response(
state: &AppState,
provider_id: String,
) -> 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) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_admin_pool_error_response(
http::StatusCode::NOT_FOUND,
format!("Provider {provider_id} 不存在"),
));
};
let banned_keys = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
.await?
.into_iter()
.filter(pool_selection::admin_pool_key_is_known_banned)
.collect::<Vec<_>>();
if banned_keys.is_empty() {
return Ok(Json(json!({
"affected": 0,
"message": ADMIN_POOL_BANNED_KEY_CLEANUP_EMPTY_MESSAGE,
}))
.into_response());
}
let deleted_key_ids = banned_keys
.iter()
.map(|key| key.id.clone())
.collect::<Vec<_>>();
for key in &banned_keys {
clear_admin_provider_pool_cooldown(state, &provider.id, &key.id).await;
reset_admin_provider_pool_cost(state, &provider.id, &key.id).await;
}
let mut affected = 0usize;
for key_id in &deleted_key_ids {
if state.delete_provider_catalog_key(key_id).await? {
affected += 1;
}
}
state
.cleanup_deleted_provider_catalog_refs(&provider.id, &[], &deleted_key_ids)
.await?;
Ok(Json(json!({
"affected": affected,
"message": format!("已清理 {affected} 个异常账号"),
}))
.into_response())
}
async fn build_admin_pool_batch_import_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
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.request_path) else {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"provider_id 无效",
));
};
let payload = match request_body {
Some(body) if !body.is_empty() => {
match serde_json::from_slice::<AdminPoolBatchImportRequest>(body) {
Ok(value) => value,
Err(_) => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"Invalid JSON request body",
));
}
}
}
_ => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"Invalid JSON request body",
));
}
};
if payload.keys.len() > 500 {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"keys length must be less than or equal to 500",
));
}
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_admin_pool_error_response(
http::StatusCode::NOT_FOUND,
format!("Provider {provider_id} 不存在"),
));
};
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let existing_keys = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let api_formats = admin_pool_resolved_api_formats(&endpoints, &existing_keys);
if api_formats.is_empty() {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"Provider 没有可用 endpoint 或现有 key,无法推断 api_formats",
));
}
let proxy = admin_pool_key_proxy_value(payload.proxy_node_id.as_deref());
let mut imported = 0usize;
let skipped = 0usize;
let mut errors = Vec::new();
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
for (index, item) in payload.keys.iter().enumerate() {
let api_key = item.api_key.trim();
if api_key.is_empty() {
errors.push(json!({
"index": index,
"reason": "api_key is empty",
}));
continue;
}
let Some(encrypted_api_key) = encrypt_catalog_secret_with_fallbacks(state, api_key) else {
errors.push(json!({
"index": index,
"reason": "gateway 未配置 provider key 加密密钥",
}));
continue;
};
let auth_type = item.auth_type.trim().to_ascii_lowercase();
let auth_type = if auth_type.is_empty() {
"api_key".to_string()
} else {
auth_type
};
let name = item.name.trim();
let mut record = match StoredProviderCatalogKey::new(
Uuid::new_v4().to_string(),
provider.id.clone(),
if name.is_empty() {
format!("imported-{index}")
} else {
name.to_string()
},
auth_type,
None,
true,
) {
Ok(value) => value,
Err(err) => {
errors.push(json!({
"index": index,
"reason": err.to_string(),
}));
continue;
}
};
record = match record.with_transport_fields(
Some(json!(api_formats)),
encrypted_api_key,
None,
None,
None,
None,
None,
proxy.clone(),
None,
) {
Ok(value) => value,
Err(err) => {
errors.push(json!({
"index": index,
"reason": err.to_string(),
}));
continue;
}
};
record.request_count = Some(0);
record.success_count = Some(0);
record.error_count = Some(0);
record.total_response_time_ms = Some(0);
record.health_by_format = Some(json!({}));
record.circuit_breaker_by_format = Some(json!({}));
record.created_at_unix_secs = Some(now_unix_secs);
record.updated_at_unix_secs = Some(now_unix_secs);
let Some(_) = state.create_provider_catalog_key(&record).await? else {
return Ok(build_admin_pool_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
ADMIN_POOL_PROVIDER_CATALOG_WRITER_UNAVAILABLE_DETAIL,
));
};
imported += 1;
}
Ok(Json(json!({
"imported": imported,
"skipped": skipped,
"errors": errors,
}))
.into_response())
}
async fn build_admin_pool_batch_action_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
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.request_path) else {
return Ok(build_admin_pool_error_response(
http::StatusCode::NOT_FOUND,
"Provider 不存在",
));
};
let payload = match request_body {
Some(body) if !body.is_empty() => {
match serde_json::from_slice::<AdminPoolBatchActionRequest>(body) {
Ok(value) => value,
Err(_) => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"Invalid JSON request body",
));
}
}
}
_ => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"Invalid JSON request body",
));
}
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_admin_pool_error_response(
http::StatusCode::NOT_FOUND,
format!("Provider {provider_id} 不存在"),
));
};
let action = payload.action.trim().to_ascii_lowercase();
let action_label = match action.as_str() {
"enable" => "enabled",
"disable" => "disabled",
"clear_proxy" => "proxy cleared",
"set_proxy" => "proxy set",
"delete" => "deleted",
_ => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
format!(
"Invalid action: {action}. Supported locally: enable, disable, clear_proxy, set_proxy, delete"
),
));
}
};
let key_ids = payload
.key_ids
.into_iter()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.collect::<BTreeSet<_>>()
.into_iter()
.collect::<Vec<_>>();
if key_ids.is_empty() {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"key_ids should not be empty",
));
}
let proxy_payload = if action == "set_proxy" {
match payload.payload {
Some(serde_json::Value::Object(map)) if !map.is_empty() => {
Some(serde_json::Value::Object(map))
}
_ => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"set_proxy action requires a non-empty payload with proxy config",
));
}
}
} else {
None
};
let keys = state
.read_provider_catalog_keys_by_ids(&key_ids)
.await?
.into_iter()
.filter(|key| key.provider_id == provider.id)
.collect::<Vec<_>>();
if action == "delete" {
let deleted_key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>();
for key in &keys {
clear_admin_provider_pool_cooldown(state, &provider.id, &key.id).await;
reset_admin_provider_pool_cost(state, &provider.id, &key.id).await;
}
let mut affected = 0usize;
for key_id in &deleted_key_ids {
if state.delete_provider_catalog_key(key_id).await? {
affected = affected.saturating_add(1);
}
}
state
.cleanup_deleted_provider_catalog_refs(&provider.id, &[], &deleted_key_ids)
.await?;
return Ok(Json(json!({
"affected": affected,
"message": format!("{affected} keys {action_label}"),
}))
.into_response());
}
let mut affected = 0usize;
for mut key in keys {
match action.as_str() {
"enable" => key.is_active = true,
"disable" => key.is_active = false,
"clear_proxy" => key.proxy = None,
"set_proxy" => key.proxy = proxy_payload.clone(),
_ => unreachable!(),
}
if state.update_provider_catalog_key(&key).await?.is_some() {
affected = affected.saturating_add(1);
}
}
Ok(Json(json!({
"affected": affected,
"message": format!("{affected} keys {action_label}"),
}))
.into_response())
}
async fn build_admin_pool_batch_delete_task_status_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
) -> Result<Response<Body>, GatewayError> {
let Some((provider_id, task_id)) =
admin_pool_batch_delete_task_parts(&request_context.request_path)
else {
return Ok(build_admin_pool_error_response(
http::StatusCode::NOT_FOUND,
"批量删除任务不存在",
));
};
let Some(task) = state.get_provider_delete_task(&task_id) else {
return Ok(build_admin_pool_error_response(
http::StatusCode::NOT_FOUND,
"批量删除任务不存在",
));
};
if task.provider_id != provider_id {
return Ok(build_admin_pool_error_response(
http::StatusCode::NOT_FOUND,
"批量删除任务不存在",
));
}
Ok(attach_admin_pool_batch_delete_task_terminal_audit(
&provider_id,
&task_id,
task.status.as_str(),
Json(build_admin_pool_batch_delete_task_payload(&task)).into_response(),
))
}
pub(super) async fn maybe_build_local_admin_pool_batch_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
match decision_route_kind(request_context) {
Some("cleanup_banned_keys") if request_context.request_method == http::Method::POST => {
let Some(provider_id) = admin_pool_provider_id_from_path(&request_context.request_path)
else {
return Ok(Some(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"provider_id 无效",
)));
};
Ok(Some(
build_admin_pool_cleanup_banned_keys_response(state, provider_id).await?,
))
}
Some("batch_import_keys") => Ok(Some(
build_admin_pool_batch_import_response(state, request_context, request_body).await?,
)),
Some("batch_action_keys") => Ok(Some(
build_admin_pool_batch_action_response(state, request_context, request_body).await?,
)),
Some("batch_delete_task_status") => Ok(Some(
build_admin_pool_batch_delete_task_status_response(state, request_context).await?,
)),
_ => Ok(None),
}
}
fn decision_route_kind<'a>(request_context: &'a GatewayPublicRequestContext) -> Option<&'a str> {
request_context
.control_decision
.as_ref()?
.route_kind
.as_deref()
}
@@ -0,0 +1,179 @@
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::shared::query_param_value;
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
const ADMIN_POOL_PROVIDER_CATALOG_READER_UNAVAILABLE_DETAIL: &str =
"Admin pool overview requires provider catalog reader";
const ADMIN_POOL_PROVIDER_CATALOG_WRITER_UNAVAILABLE_DETAIL: &str =
"Admin pool cleanup requires provider catalog writer";
const ADMIN_POOL_BANNED_KEY_CLEANUP_EMPTY_MESSAGE: &str = "未发现可清理的异常账号";
mod batch_routes;
mod payloads;
mod read_routes;
mod selection;
#[derive(Debug, Default, serde::Deserialize)]
struct AdminPoolResolveSelectionRequest {
#[serde(default)]
search: String,
#[serde(default)]
quick_selectors: Vec<String>,
}
fn build_admin_pool_error_response(
status: http::StatusCode,
detail: impl Into<String>,
) -> Response<Body> {
(status, Json(json!({ "detail": detail.into() }))).into_response()
}
fn parse_admin_pool_page(query: Option<&str>) -> Result<usize, String> {
match query_param_value(query, "page") {
None => Ok(1),
Some(value) => {
let parsed = value
.parse::<usize>()
.map_err(|_| "page must be an integer between 1 and 10000".to_string())?;
if (1..=10_000).contains(&parsed) {
Ok(parsed)
} else {
Err("page must be an integer between 1 and 10000".to_string())
}
}
}
}
fn parse_admin_pool_page_size(query: Option<&str>) -> Result<usize, String> {
match query_param_value(query, "page_size") {
None => Ok(50),
Some(value) => {
let parsed = value
.parse::<usize>()
.map_err(|_| "page_size must be an integer between 1 and 200".to_string())?;
if (1..=200).contains(&parsed) {
Ok(parsed)
} else {
Err("page_size must be an integer between 1 and 200".to_string())
}
}
}
}
fn parse_admin_pool_search(query: Option<&str>) -> Option<String> {
query_param_value(query, "search")
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
}
fn parse_admin_pool_status_filter(query: Option<&str>) -> Result<String, String> {
let value = query_param_value(query, "status")
.unwrap_or_else(|| "all".to_string())
.trim()
.to_ascii_lowercase();
match value.as_str() {
"all" | "active" | "inactive" | "cooldown" => Ok(value),
_ => Err("status must be one of: all, active, cooldown, inactive".to_string()),
}
}
fn admin_pool_provider_id_from_path(request_path: &str) -> Option<String> {
let raw = request_path.strip_prefix("/api/admin/pool/")?;
let mut segments = raw.split('/');
let provider_id = segments.next()?.trim();
let keys_segment = segments.next()?.trim();
if provider_id.is_empty() || keys_segment != "keys" {
None
} else {
Some(provider_id.to_string())
}
}
fn is_admin_pool_route(request_context: &GatewayPublicRequestContext) -> bool {
let normalized_path = request_context.request_path.trim_end_matches('/');
let path = if normalized_path.is_empty() {
request_context.request_path.as_str()
} else {
normalized_path
};
(request_context.request_method == http::Method::GET && path == "/api/admin/pool/overview")
|| (request_context.request_method == http::Method::GET
&& path == "/api/admin/pool/scheduling-presets")
|| (request_context.request_method == http::Method::GET
&& path.starts_with("/api/admin/pool/")
&& path.ends_with("/keys")
&& path.matches('/').count() == 5)
|| (request_context.request_method == http::Method::POST
&& path.starts_with("/api/admin/pool/")
&& path.ends_with("/keys/batch-import")
&& path.matches('/').count() == 6)
|| (request_context.request_method == http::Method::POST
&& path.starts_with("/api/admin/pool/")
&& path.ends_with("/keys/batch-action")
&& path.matches('/').count() == 6)
|| (request_context.request_method == http::Method::POST
&& path.starts_with("/api/admin/pool/")
&& path.ends_with("/keys/resolve-selection")
&& path.matches('/').count() == 6)
|| (request_context.request_method == http::Method::GET
&& path.starts_with("/api/admin/pool/")
&& path.contains("/keys/batch-delete-task/")
&& path.matches('/').count() == 7)
|| (request_context.request_method == http::Method::POST
&& path.starts_with("/api/admin/pool/")
&& path.ends_with("/keys/cleanup-banned")
&& path.matches('/').count() == 6)
}
pub(crate) async fn maybe_build_local_admin_pool_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.control_decision.as_ref() else {
return Ok(None);
};
if decision.route_family.as_deref() != Some("pool_manage") {
return Ok(None);
}
if !is_admin_pool_route(request_context) {
return Ok(None);
}
if let Some(response) = batch_routes::maybe_build_local_admin_pool_batch_response(
state,
request_context,
request_body,
)
.await?
{
return Ok(Some(response));
}
if let Some(response) = read_routes::maybe_build_local_admin_pool_read_response(
state,
request_context,
request_body,
)
.await?
{
return Ok(Some(response));
}
Ok(Some(build_admin_pool_error_response(
http::StatusCode::NOT_FOUND,
format!(
"Unsupported admin pool route {} {}",
request_context.request_method, request_context.request_path
),
)))
}
@@ -0,0 +1,221 @@
use crate::handlers::admin::provider::shared::{
AdminProviderPoolConfig, AdminProviderPoolRuntimeState,
};
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use serde_json::json;
pub(super) fn admin_pool_api_formats(key: &StoredProviderCatalogKey) -> Vec<String> {
key.api_formats
.as_ref()
.and_then(serde_json::Value::as_array)
.map(|values| {
values
.iter()
.filter_map(serde_json::Value::as_str)
.map(ToOwned::to_owned)
.collect::<Vec<_>>()
})
.unwrap_or_default()
}
fn admin_pool_string_list(value: Option<&serde_json::Value>) -> Option<Vec<String>> {
let values = value
.and_then(serde_json::Value::as_array)
.map(|items| {
items
.iter()
.filter_map(serde_json::Value::as_str)
.map(ToOwned::to_owned)
.collect::<Vec<_>>()
})
.unwrap_or_default();
if values.is_empty() {
None
} else {
Some(values)
}
}
fn admin_pool_json_object(
value: Option<&serde_json::Value>,
) -> Option<serde_json::Map<String, serde_json::Value>> {
value
.and_then(serde_json::Value::as_object)
.cloned()
.filter(|value| !value.is_empty())
}
fn admin_pool_health_score(key: &StoredProviderCatalogKey) -> f64 {
let scores = key
.health_by_format
.as_ref()
.and_then(serde_json::Value::as_object)
.map(|formats| {
formats
.values()
.filter_map(serde_json::Value::as_object)
.filter_map(|item| item.get("health_score"))
.filter_map(serde_json::Value::as_f64)
.collect::<Vec<_>>()
})
.unwrap_or_default();
if scores.is_empty() {
1.0
} else {
scores.into_iter().fold(1.0, f64::min)
}
}
fn admin_pool_circuit_breaker_open(key: &StoredProviderCatalogKey) -> bool {
key.circuit_breaker_by_format
.as_ref()
.and_then(serde_json::Value::as_object)
.map(|formats| {
formats
.values()
.filter_map(serde_json::Value::as_object)
.any(|item| {
item.get("open")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
})
})
.unwrap_or(false)
}
fn admin_pool_scheduling_payload(
key: &StoredProviderCatalogKey,
cooldown_reason: Option<&str>,
cooldown_ttl_seconds: Option<u64>,
health_score: f64,
circuit_breaker_open: bool,
) -> (String, String, String, Vec<serde_json::Value>) {
if !key.is_active {
return (
"blocked".to_string(),
"inactive".to_string(),
"已禁用".to_string(),
vec![json!({
"code": "inactive",
"label": "已禁用",
"blocking": true,
"source": "manual",
"ttl_seconds": serde_json::Value::Null,
"detail": serde_json::Value::Null,
})],
);
}
if let Some(reason) = cooldown_reason {
return (
"degraded".to_string(),
"cooldown".to_string(),
"冷却中".to_string(),
vec![json!({
"code": "cooldown",
"label": "冷却中",
"blocking": true,
"source": "pool",
"ttl_seconds": cooldown_ttl_seconds,
"detail": reason,
})],
);
}
if circuit_breaker_open {
return (
"degraded".to_string(),
"circuit_breaker".to_string(),
"熔断中".to_string(),
vec![json!({
"code": "circuit_breaker",
"label": "熔断中",
"blocking": true,
"source": "health",
"ttl_seconds": serde_json::Value::Null,
"detail": serde_json::Value::Null,
})],
);
}
if health_score < 0.5 {
return (
"degraded".to_string(),
"health_low".to_string(),
"健康度较低".to_string(),
vec![json!({
"code": "health_low",
"label": "健康度较低",
"blocking": false,
"source": "health",
"ttl_seconds": serde_json::Value::Null,
"detail": serde_json::Value::Null,
})],
);
}
(
"available".to_string(),
"available".to_string(),
"可用".to_string(),
Vec::new(),
)
}
pub(super) fn build_admin_pool_key_payload(
key: &StoredProviderCatalogKey,
runtime: &AdminProviderPoolRuntimeState,
pool_config: Option<AdminProviderPoolConfig>,
) -> serde_json::Value {
let cooldown_reason = runtime.cooldown_reason_by_key.get(&key.id).cloned();
let cooldown_ttl_seconds = cooldown_reason
.as_ref()
.and_then(|_| runtime.cooldown_ttl_by_key.get(&key.id).copied());
let health_score = admin_pool_health_score(key);
let circuit_breaker_open = admin_pool_circuit_breaker_open(key);
let (scheduling_status, scheduling_reason, scheduling_label, scheduling_reasons) =
admin_pool_scheduling_payload(
key,
cooldown_reason.as_deref(),
cooldown_ttl_seconds,
health_score,
circuit_breaker_open,
);
json!({
"key_id": key.id,
"key_name": key.name,
"is_active": key.is_active,
"auth_type": key.auth_type,
"status_snapshot": key.status_snapshot.clone().unwrap_or_else(|| json!({})),
"health_score": health_score,
"circuit_breaker_open": circuit_breaker_open,
"api_formats": admin_pool_api_formats(key),
"rate_multipliers": admin_pool_json_object(key.rate_multipliers.as_ref()),
"internal_priority": key.internal_priority,
"rpm_limit": key.rpm_limit,
"cache_ttl_minutes": key.cache_ttl_minutes,
"max_probe_interval_minutes": key.max_probe_interval_minutes,
"note": key.note,
"allowed_models": admin_pool_string_list(key.allowed_models.as_ref()),
"capabilities": admin_pool_json_object(key.capabilities.as_ref()),
"auto_fetch_models": key.auto_fetch_models,
"locked_models": admin_pool_string_list(key.locked_models.as_ref()),
"model_include_patterns": admin_pool_string_list(key.model_include_patterns.as_ref()),
"model_exclude_patterns": admin_pool_string_list(key.model_exclude_patterns.as_ref()),
"proxy": key.proxy.clone(),
"fingerprint": key.fingerprint.clone(),
"cooldown_reason": cooldown_reason,
"cooldown_ttl_seconds": cooldown_ttl_seconds,
"cost_window_usage": runtime.cost_window_usage_by_key.get(&key.id).copied().unwrap_or(0),
"cost_limit": pool_config.map(|config| config.cost_limit_per_key_tokens),
"request_count": key.request_count.unwrap_or(0),
"total_tokens": 0,
"total_cost_usd": "0.00000000",
"sticky_sessions": runtime.sticky_sessions_by_key.get(&key.id).copied().unwrap_or(0),
"lru_score": runtime.lru_score_by_key.get(&key.id).copied(),
"created_at": key.created_at_unix_secs.and_then(unix_secs_to_rfc3339),
"last_used_at": key.last_used_at_unix_secs.and_then(unix_secs_to_rfc3339),
"scheduling_status": scheduling_status,
"scheduling_reason": scheduling_reason,
"scheduling_label": scheduling_label,
"scheduling_reasons": scheduling_reasons,
})
}
@@ -0,0 +1,506 @@
use super::{
admin_pool_provider_id_from_path, build_admin_pool_error_response, parse_admin_pool_page,
parse_admin_pool_page_size, parse_admin_pool_search, parse_admin_pool_status_filter,
AdminPoolResolveSelectionRequest, ADMIN_POOL_PROVIDER_CATALOG_READER_UNAVAILABLE_DETAIL,
};
use super::{payloads as pool_payloads, selection as pool_selection};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::pool::{
admin_provider_pool_config, read_admin_provider_pool_cooldown_counts,
read_admin_provider_pool_cooldown_key_ids, read_admin_provider_pool_runtime_state,
};
use crate::handlers::admin::provider::shared::AdminProviderPoolRuntimeState;
use crate::{AppState, GatewayError};
use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyListQuery;
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use std::collections::BTreeMap;
async fn build_admin_pool_overview_payload(
state: &AppState,
) -> Result<serde_json::Value, GatewayError> {
let providers = state.list_provider_catalog_providers(false).await?;
let pool_enabled_providers = providers
.into_iter()
.filter_map(|provider| {
admin_provider_pool_config(&provider).map(|config| (provider, config))
})
.collect::<Vec<_>>();
let provider_ids = pool_enabled_providers
.iter()
.map(|(provider, _)| provider.id.clone())
.collect::<Vec<_>>();
let key_stats = if provider_ids.is_empty() {
Vec::new()
} else {
state
.list_provider_catalog_key_stats_by_provider_ids(&provider_ids)
.await?
};
let key_stats_by_provider = key_stats
.into_iter()
.map(|item| (item.provider_id.clone(), item))
.collect::<BTreeMap<_, _>>();
let redis_runner = state.redis_kv_runner();
let cooldown_counts_by_provider = match redis_runner.as_ref() {
Some(runner) if !provider_ids.is_empty() => {
read_admin_provider_pool_cooldown_counts(runner, &provider_ids).await
}
_ => BTreeMap::new(),
};
let mut items = Vec::with_capacity(pool_enabled_providers.len());
for (provider, _pool_config) in pool_enabled_providers {
let stats = key_stats_by_provider.get(&provider.id);
let total_keys = stats.map(|item| item.total_keys as usize).unwrap_or(0);
let active_keys = stats.map(|item| item.active_keys as usize).unwrap_or(0);
let cooldown_count = cooldown_counts_by_provider
.get(&provider.id)
.copied()
.unwrap_or(0);
items.push(json!({
"provider_id": provider.id,
"provider_name": provider.name,
"provider_type": provider.provider_type,
"total_keys": total_keys,
"active_keys": active_keys,
"cooldown_count": cooldown_count,
"pool_enabled": true,
}));
}
Ok(json!({ "items": items }))
}
async fn build_admin_pool_list_keys_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
) -> 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,
));
}
let Some(provider_id) = admin_pool_provider_id_from_path(&request_context.request_path) else {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"provider_id 无效",
));
};
let query = request_context.request_query_string.as_deref();
let page = match parse_admin_pool_page(query) {
Ok(value) => value,
Err(detail) => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
detail,
));
}
};
let page_size = match parse_admin_pool_page_size(query) {
Ok(value) => value,
Err(detail) => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
detail,
));
}
};
let search = parse_admin_pool_search(query).map(|value| value.to_ascii_lowercase());
let status = match parse_admin_pool_status_filter(query) {
Ok(value) => value,
Err(detail) => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
detail,
));
}
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_admin_pool_error_response(
http::StatusCode::NOT_FOUND,
format!("Provider {provider_id} 不存在"),
));
};
let pool_config = admin_provider_pool_config(&provider);
let page_offset = page.saturating_sub(1).saturating_mul(page_size);
let (keys, total) = if status == "cooldown" {
let cooldown_key_ids = if let Some(runner) = state.redis_kv_runner() {
read_admin_provider_pool_cooldown_key_ids(&runner, &provider.id).await
} else {
Vec::new()
};
let mut keys = if cooldown_key_ids.is_empty() {
Vec::new()
} else {
state
.list_provider_catalog_keys_by_ids(&cooldown_key_ids)
.await?
};
if let Some(keyword) = search.as_ref() {
keys.retain(|key| {
key.name.to_ascii_lowercase().contains(keyword)
|| key.id.to_ascii_lowercase().contains(keyword)
});
}
pool_selection::admin_pool_sort_keys(&mut keys);
let total = keys.len();
let keys = keys
.into_iter()
.skip(page_offset)
.take(page_size)
.collect::<Vec<_>>();
(keys, total)
} else {
let key_page = state
.list_provider_catalog_key_page(&ProviderCatalogKeyListQuery {
provider_id: provider.id.clone(),
search: search.clone(),
is_active: match status.as_str() {
"active" => Some(true),
"inactive" => Some(false),
_ => None,
},
offset: page_offset,
limit: page_size,
})
.await?;
(key_page.items, key_page.total)
};
let key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>();
let runtime = match (state.redis_kv_runner(), pool_config) {
(Some(runner), Some(pool_config)) if !key_ids.is_empty() => {
read_admin_provider_pool_runtime_state(&runner, &provider.id, &key_ids, pool_config)
.await
}
_ => AdminProviderPoolRuntimeState::default(),
};
let items = keys
.into_iter()
.map(|key| pool_payloads::build_admin_pool_key_payload(&key, &runtime, pool_config))
.collect::<Vec<_>>();
Ok(Json(json!({
"total": total,
"page": page,
"page_size": page_size,
"keys": items,
}))
.into_response())
}
async fn build_admin_pool_resolve_selection_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
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,
));
}
let Some(provider_id) = admin_pool_provider_id_from_path(&request_context.request_path) else {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"provider_id 无效",
));
};
let payload = match request_body {
None => AdminPoolResolveSelectionRequest::default(),
Some(body) if body.is_empty() => AdminPoolResolveSelectionRequest::default(),
Some(body) => match serde_json::from_slice::<AdminPoolResolveSelectionRequest>(body) {
Ok(value) => value,
Err(_) => {
return Ok(build_admin_pool_error_response(
http::StatusCode::BAD_REQUEST,
"Invalid JSON request body",
));
}
},
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(build_admin_pool_error_response(
http::StatusCode::NOT_FOUND,
format!("Provider {provider_id} 不存在"),
));
};
let provider_type = provider.provider_type.clone();
let search = payload.search.trim();
let mut quick_selectors = payload
.quick_selectors
.into_iter()
.map(pool_selection::admin_pool_normalize_text)
.filter(|value| {
matches!(
value.as_str(),
"banned"
| "no_5h_limit"
| "no_weekly_limit"
| "plan_free"
| "plan_team"
| "oauth_invalid"
| "proxy_unset"
| "proxy_set"
| "disabled"
| "enabled"
)
})
.collect::<Vec<_>>();
quick_selectors.sort();
quick_selectors.dedup();
let mut keys = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
.await?
.into_iter()
.filter(|key| {
pool_selection::admin_pool_matches_search(state, key, &provider_type, Some(search))
})
.filter(|key| {
quick_selectors.is_empty()
|| quick_selectors.iter().all(|selector| {
pool_selection::admin_pool_matches_quick_selector(
state,
key,
&provider_type,
selector,
)
})
})
.collect::<Vec<_>>();
keys.sort_by(|left, right| {
left.internal_priority
.cmp(&right.internal_priority)
.then_with(|| left.name.cmp(&right.name))
});
let items = keys
.iter()
.map(|key| {
json!({
"key_id": key.id,
"key_name": key.name,
"auth_type": key.auth_type,
})
})
.collect::<Vec<_>>();
Ok(Json(json!({
"total": items.len(),
"items": items,
}))
.into_response())
}
fn build_admin_pool_scheduling_presets_payload() -> serde_json::Value {
json!([
{
"name": "lru",
"label": "LRU 轮转",
"description": "最久未使用的 Key 优先",
"providers": [],
"modes": serde_json::Value::Null,
"default_mode": serde_json::Value::Null,
"mutex_group": "distribution_mode",
"evidence_hint": "依据 LRU 时间戳(最近未使用优先)",
},
{
"name": "cache_affinity",
"label": "缓存亲和",
"description": "优先复用最近使用过的 Key,利用 Prompt Caching",
"providers": [],
"modes": serde_json::Value::Null,
"default_mode": serde_json::Value::Null,
"mutex_group": "distribution_mode",
"evidence_hint": "依据 LRU 时间戳(最近使用优先,与 LRU 轮转相反)",
},
{
"name": "cost_first",
"label": "成本优先",
"description": "优先选择窗口消耗更低的账号",
"providers": [],
"modes": serde_json::Value::Null,
"default_mode": serde_json::Value::Null,
"mutex_group": serde_json::Value::Null,
"evidence_hint": "依据窗口成本/Token 用量,缺失时回退配额使用率",
},
{
"name": "free_first",
"label": "Free 优先",
"description": "优先消耗 Free 账号(依赖 plan_type)",
"providers": ["codex", "kiro"],
"modes": serde_json::Value::Null,
"default_mode": serde_json::Value::Null,
"mutex_group": serde_json::Value::Null,
"evidence_hint": "依据 plan_type(Free 账号优先调度)",
},
{
"name": "health_first",
"label": "健康优先",
"description": "优先选择健康分更高、失败更少的账号",
"providers": [],
"modes": serde_json::Value::Null,
"default_mode": serde_json::Value::Null,
"mutex_group": serde_json::Value::Null,
"evidence_hint": "依据 health_by_format 聚合分(含熔断/失败衰减)",
},
{
"name": "latency_first",
"label": "延迟优先",
"description": "优先选择最近延迟更低的账号",
"providers": [],
"modes": serde_json::Value::Null,
"default_mode": serde_json::Value::Null,
"mutex_group": serde_json::Value::Null,
"evidence_hint": "依据号池延迟窗口均值(latency_window_seconds)",
},
{
"name": "load_balance",
"label": "负载均衡",
"description": "随机分散 Key 使用,均匀分摊负载",
"providers": [],
"modes": serde_json::Value::Null,
"default_mode": serde_json::Value::Null,
"mutex_group": "distribution_mode",
"evidence_hint": "每次随机分值,实现完全均匀分散",
},
{
"name": "plus_first",
"label": "Plus 优先",
"description": "优先消耗 Plus/Pro 账号(依赖 plan_type)",
"providers": ["codex", "kiro"],
"modes": serde_json::Value::Null,
"default_mode": serde_json::Value::Null,
"mutex_group": serde_json::Value::Null,
"evidence_hint": "依据 plan_type(Plus/Pro 账号优先调度)",
},
{
"name": "priority_first",
"label": "优先级优先",
"description": "按账号优先级顺序调度(数字越小越优先)",
"providers": [],
"modes": serde_json::Value::Null,
"default_mode": serde_json::Value::Null,
"mutex_group": serde_json::Value::Null,
"evidence_hint": "依据 internal_priority(支持拖拽/手工编辑)",
},
{
"name": "quota_balanced",
"label": "额度平均",
"description": "优先选额度消耗最少的账号",
"providers": [],
"modes": serde_json::Value::Null,
"default_mode": serde_json::Value::Null,
"mutex_group": serde_json::Value::Null,
"evidence_hint": "依据账号配额使用率;无配额时回退到窗口成本使用",
},
{
"name": "recent_refresh",
"label": "额度刷新优先",
"description": "优先选即将刷新额度的账号",
"providers": ["codex", "kiro"],
"modes": serde_json::Value::Null,
"default_mode": serde_json::Value::Null,
"mutex_group": serde_json::Value::Null,
"evidence_hint": "依据账号额度重置倒计时(next_reset / reset_seconds)",
},
{
"name": "single_account",
"label": "单号优先",
"description": "集中使用同一账号(反向 LRU)",
"providers": [],
"modes": serde_json::Value::Null,
"default_mode": serde_json::Value::Null,
"mutex_group": "distribution_mode",
"evidence_hint": "先按账号优先级(internal_priority),同级再按反向 LRU 集中",
},
{
"name": "team_first",
"label": "Team 优先",
"description": "优先消耗 Team 账号(依赖 plan_type)",
"providers": ["codex", "kiro"],
"modes": serde_json::Value::Null,
"default_mode": serde_json::Value::Null,
"mutex_group": serde_json::Value::Null,
"evidence_hint": "依据 plan_type(Team 账号优先调度)",
}
])
}
pub(super) async fn maybe_build_local_admin_pool_read_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&axum::body::Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
match request_context
.control_decision
.as_ref()
.and_then(|decision| decision.route_kind.as_deref())
{
Some("overview")
if request_context.request_method == http::Method::GET
&& matches!(
request_context.request_path.trim_end_matches('/'),
"/api/admin/pool/overview"
) =>
{
if !state.has_provider_catalog_data_reader() {
return Ok(Some(build_admin_pool_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
ADMIN_POOL_PROVIDER_CATALOG_READER_UNAVAILABLE_DETAIL,
)));
}
Ok(Some(
Json(build_admin_pool_overview_payload(state).await?).into_response(),
))
}
Some("scheduling_presets")
if request_context.request_method == http::Method::GET
&& matches!(
request_context.request_path.trim_end_matches('/'),
"/api/admin/pool/scheduling-presets"
) =>
{
Ok(Some(
Json(build_admin_pool_scheduling_presets_payload()).into_response(),
))
}
Some("list_keys") => Ok(Some(
build_admin_pool_list_keys_response(state, request_context).await?,
)),
Some("resolve_selection") => Ok(Some(
build_admin_pool_resolve_selection_response(state, request_context, request_body)
.await?,
)),
_ => Ok(None),
}
}
@@ -0,0 +1,246 @@
use crate::handlers::admin::shared::decrypt_catalog_secret_with_fallbacks;
use crate::AppState;
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
fn admin_pool_reason_indicates_ban(reason: &str) -> bool {
let normalized = reason.trim().to_ascii_lowercase();
!normalized.is_empty()
&& [
"banned",
"forbidden",
"blocked",
"suspend",
"deactivated",
"disabled",
"verification",
"workspace",
"受限",
"封",
"禁",
]
.iter()
.any(|hint| normalized.contains(hint))
}
pub(super) fn admin_pool_normalize_text(value: impl AsRef<str>) -> String {
value.as_ref().trim().to_ascii_lowercase()
}
fn admin_pool_parse_auth_config_json(
state: &AppState,
key: &StoredProviderCatalogKey,
) -> Option<serde_json::Map<String, serde_json::Value>> {
let ciphertext = key.encrypted_auth_config.as_deref()?.trim();
if ciphertext.is_empty() {
return None;
}
let plaintext = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)?;
serde_json::from_str::<serde_json::Value>(&plaintext)
.ok()?
.as_object()
.cloned()
}
fn admin_pool_derive_oauth_plan_type(
state: &AppState,
key: &StoredProviderCatalogKey,
provider_type: &str,
) -> Option<String> {
let normalize = |value: &str| {
let mut text = value.trim().to_string();
if text.is_empty() {
return None;
}
let provider_type = provider_type.trim().to_ascii_lowercase();
if !provider_type.is_empty() && text.to_ascii_lowercase().starts_with(&provider_type) {
text = text[provider_type.len()..]
.trim_matches(|ch: char| [' ', ':', '-', '_'].contains(&ch))
.to_string();
}
if text.is_empty() {
None
} else {
Some(text.to_ascii_lowercase())
}
};
if key.auth_type.trim() != "oauth" {
return None;
}
if let Some(auth_config) = admin_pool_parse_auth_config_json(state, key) {
for plan_key in ["plan_type", "tier", "plan", "subscription_plan"] {
if let Some(value) = auth_config
.get(plan_key)
.and_then(serde_json::Value::as_str)
{
if let Some(normalized) = normalize(value) {
return Some(normalized);
}
}
}
}
let upstream_metadata = key.upstream_metadata.as_ref()?.as_object()?;
let provider_bucket = upstream_metadata
.get(&provider_type.trim().to_ascii_lowercase())
.and_then(serde_json::Value::as_object);
for source in provider_bucket
.into_iter()
.chain(std::iter::once(upstream_metadata))
{
for plan_key in [
"plan_type",
"tier",
"subscription_title",
"subscription_plan",
] {
if let Some(value) = source.get(plan_key).and_then(serde_json::Value::as_str) {
if let Some(normalized) = normalize(value) {
return Some(normalized);
}
}
}
}
None
}
fn admin_pool_has_proxy(key: &StoredProviderCatalogKey) -> bool {
match key.proxy.as_ref() {
Some(serde_json::Value::Object(values)) => !values.is_empty(),
Some(serde_json::Value::String(value)) => !value.trim().is_empty(),
Some(serde_json::Value::Bool(value)) => *value,
Some(serde_json::Value::Number(_)) => true,
Some(serde_json::Value::Array(values)) => !values.is_empty(),
_ => false,
}
}
fn admin_pool_is_oauth_invalid(key: &StoredProviderCatalogKey) -> bool {
if key.auth_type.trim() != "oauth" {
return false;
}
if key
.oauth_invalid_reason
.as_deref()
.is_some_and(|value| !value.trim().is_empty())
{
return true;
}
key.expires_at_unix_secs
.is_some_and(|value| value > 0 && value <= chrono::Utc::now().timestamp().max(0) as u64)
}
pub(super) fn admin_pool_matches_quick_selector(
state: &AppState,
key: &StoredProviderCatalogKey,
provider_type: &str,
selector: &str,
) -> bool {
match selector {
"banned" => admin_pool_key_is_known_banned(key),
"oauth_invalid" => admin_pool_is_oauth_invalid(key),
"proxy_unset" => !admin_pool_has_proxy(key),
"proxy_set" => admin_pool_has_proxy(key),
"disabled" => !key.is_active,
"enabled" => key.is_active,
"plan_free" => admin_pool_derive_oauth_plan_type(state, key, provider_type)
.is_some_and(|value| value.contains("free")),
"plan_team" => admin_pool_derive_oauth_plan_type(state, key, provider_type)
.is_some_and(|value| value.contains("team")),
"no_5h_limit" | "no_weekly_limit" => false,
_ => false,
}
}
pub(super) fn admin_pool_matches_search(
state: &AppState,
key: &StoredProviderCatalogKey,
provider_type: &str,
search: Option<&str>,
) -> bool {
let Some(search) = search else {
return true;
};
let search = admin_pool_normalize_text(search);
if search.is_empty() {
return true;
}
let oauth_plan_type = admin_pool_derive_oauth_plan_type(state, key, provider_type);
let mut search_fields = vec![
key.id.clone(),
key.name.clone(),
key.auth_type.clone(),
if key.is_active {
"已启用".to_string()
} else {
"已禁用".to_string()
},
if admin_pool_has_proxy(key) {
"独立代理".to_string()
} else {
"未配置代理".to_string()
},
];
if let Some(reason) = key.oauth_invalid_reason.as_ref() {
search_fields.push(reason.clone());
}
if let Some(note) = key.note.as_ref() {
search_fields.push(note.clone());
}
if let Some(plan_type) = oauth_plan_type {
search_fields.push(plan_type);
}
search_fields
.into_iter()
.any(|value| admin_pool_normalize_text(&value).contains(&search))
}
pub(super) fn admin_pool_key_is_known_banned(key: &StoredProviderCatalogKey) -> bool {
if key
.oauth_invalid_reason
.as_deref()
.is_some_and(admin_pool_reason_indicates_ban)
{
return true;
}
let Some(account) = key
.status_snapshot
.as_ref()
.and_then(serde_json::Value::as_object)
.and_then(|snapshot| snapshot.get("account"))
.and_then(serde_json::Value::as_object)
else {
return false;
};
if !account
.get("blocked")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
{
return false;
}
account
.get("code")
.and_then(serde_json::Value::as_str)
.is_some_and(admin_pool_reason_indicates_ban)
|| account
.get("reason")
.and_then(serde_json::Value::as_str)
.is_some_and(admin_pool_reason_indicates_ban)
}
pub(super) fn admin_pool_sort_keys(keys: &mut [StoredProviderCatalogKey]) {
keys.sort_by(|left, right| {
left.internal_priority
.cmp(&right.internal_priority)
.then(left.name.cmp(&right.name))
.then(left.id.cmp(&right.id))
});
}
@@ -0,0 +1,17 @@
use crate::control::GatewayPublicRequestContext;
use crate::{AppState, GatewayError};
use axum::body::{Body, Bytes};
use axum::http::Response;
mod models;
mod routes;
mod shared;
pub(crate) async fn maybe_build_local_admin_provider_query_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
routes::maybe_build_local_admin_provider_query_response(state, request_context, request_body)
.await
}
@@ -0,0 +1,228 @@
use super::shared::{
build_admin_provider_query_bad_request_response, build_admin_provider_query_not_found_response,
provider_query_extract_api_key_id, provider_query_extract_provider_id,
ADMIN_PROVIDER_QUERY_API_KEY_NOT_FOUND_DETAIL, ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
ADMIN_PROVIDER_QUERY_NO_LOCAL_MODELS_DETAIL, ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL,
ADMIN_PROVIDER_QUERY_PROVIDER_NOT_FOUND_DETAIL,
};
use crate::{AppState, GatewayError};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
};
use axum::{body::Body, http::Response, response::IntoResponse, Json};
use serde_json::json;
use std::collections::{BTreeMap, BTreeSet};
pub(super) const ADMIN_PROVIDER_QUERY_LOCAL_TEST_MODEL_MESSAGE: &str =
"Rust local provider-query model test is not configured";
pub(super) const ADMIN_PROVIDER_QUERY_LOCAL_TEST_MODEL_FAILOVER_MESSAGE: &str =
"Rust local provider-query failover simulation is not configured";
fn provider_query_string_list(value: Option<&serde_json::Value>) -> Vec<String> {
value
.and_then(serde_json::Value::as_array)
.map(|items| {
items
.iter()
.filter_map(serde_json::Value::as_str)
.map(str::trim)
.filter(|item| !item.is_empty())
.map(ToOwned::to_owned)
.collect::<Vec<_>>()
})
.unwrap_or_default()
}
fn provider_query_resolved_api_formats(
endpoints: &[StoredProviderCatalogEndpoint],
selected_key: Option<&StoredProviderCatalogKey>,
) -> Vec<String> {
let mut seen = BTreeSet::new();
let key_formats = selected_key
.map(|key| provider_query_string_list(key.api_formats.as_ref()))
.unwrap_or_default();
let mut formats = Vec::new();
for endpoint in endpoints.iter().filter(|endpoint| endpoint.is_active) {
let api_format = endpoint.api_format.trim();
if api_format.is_empty() {
continue;
}
if !key_formats.is_empty() && !key_formats.iter().any(|value| value == api_format) {
continue;
}
if seen.insert(api_format.to_string()) {
formats.push(api_format.to_string());
}
}
if formats.is_empty() {
for api_format in key_formats {
if seen.insert(api_format.clone()) {
formats.push(api_format);
}
}
}
formats
}
pub(super) async fn build_admin_provider_query_models_response(
state: &AppState,
payload: &serde_json::Value,
) -> Result<Response<Body>, GatewayError> {
let Some(provider_id) = provider_query_extract_provider_id(payload) else {
return Ok(build_admin_provider_query_bad_request_response(
ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL,
));
};
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.find(|item| item.id == provider_id)
else {
return Ok(build_admin_provider_query_not_found_response(
ADMIN_PROVIDER_QUERY_PROVIDER_NOT_FOUND_DETAIL,
));
};
let provider_ids = vec![provider.id.clone()];
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(&provider_ids)
.await?;
let keys = state
.list_provider_catalog_keys_by_provider_ids(&provider_ids)
.await?;
let selected_key = if let Some(api_key_id) = provider_query_extract_api_key_id(payload) {
let Some(key) = keys.iter().find(|key| key.id == api_key_id) else {
return Ok(build_admin_provider_query_not_found_response(
ADMIN_PROVIDER_QUERY_API_KEY_NOT_FOUND_DETAIL,
));
};
Some(key)
} else {
None
};
let active_keys = keys.iter().filter(|key| key.is_active).count();
if selected_key.is_none() && active_keys == 0 {
return Ok(build_admin_provider_query_bad_request_response(
ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL,
));
}
let resolved_api_formats = provider_query_resolved_api_formats(&endpoints, selected_key);
let provider_models = state
.list_admin_provider_available_source_models(&provider.id)
.await?;
let mut grouped: BTreeMap<
String,
(
aether_data_contracts::repository::global_models::StoredAdminProviderModel,
BTreeSet<String>,
),
> = BTreeMap::new();
for model in provider_models {
let entry = grouped
.entry(model.provider_model_name.clone())
.or_insert_with(|| (model.clone(), BTreeSet::new()));
for api_format in &resolved_api_formats {
entry.1.insert(api_format.clone());
}
}
let models: Vec<_> = grouped
.into_iter()
.map(|(model_id, (model, api_formats))| {
let display_name = model
.global_model_display_name
.clone()
.or(model.global_model_name.clone())
.unwrap_or_else(|| model_id.clone());
let api_formats: Vec<_> = api_formats.into_iter().collect();
json!({
"id": model_id,
"object": "model",
"created": model.created_at_unix_secs,
"owned_by": provider.name,
"display_name": display_name,
"api_format": api_formats.first().cloned(),
"api_formats": api_formats,
"provider_model_name": model.provider_model_name,
"global_model_id": model.global_model_id,
"global_model_name": model.global_model_name,
"supports_streaming": model.supports_streaming,
"supports_function_calling": model.supports_function_calling,
"supports_vision": model.supports_vision,
"supports_extended_thinking": model.supports_extended_thinking,
"supports_image_generation": model.supports_image_generation,
"is_available": model.is_available,
})
})
.collect();
let success = !models.is_empty();
let error = if success {
None
} else {
Some(ADMIN_PROVIDER_QUERY_NO_LOCAL_MODELS_DETAIL)
};
Ok(Json(json!({
"success": success,
"data": {
"models": models,
"error": error,
"from_cache": true,
"keys_total": active_keys,
"keys_cached": 0,
"keys_fetched": 0,
},
"provider": {
"id": provider.id,
"name": provider.name,
"display_name": provider.name,
},
}))
.into_response())
}
pub(super) fn build_admin_provider_query_test_model_response(
provider_id: String,
model: String,
) -> Response<Body> {
Json(json!({
"success": false,
"tested": false,
"provider_id": provider_id,
"model": model,
"attempts": [],
"total_candidates": 0,
"total_attempts": 0,
"error": ADMIN_PROVIDER_QUERY_LOCAL_TEST_MODEL_MESSAGE,
"source": "local",
"message": ADMIN_PROVIDER_QUERY_LOCAL_TEST_MODEL_MESSAGE,
}))
.into_response()
}
pub(super) fn build_admin_provider_query_test_model_failover_response(
provider_id: String,
failover_models: Vec<String>,
) -> Response<Body> {
Json(json!({
"success": false,
"tested": false,
"provider_id": provider_id,
"model": failover_models.first().cloned(),
"failover_models": failover_models,
"attempts": [],
"total_candidates": 0,
"total_attempts": 0,
"error": ADMIN_PROVIDER_QUERY_LOCAL_TEST_MODEL_FAILOVER_MESSAGE,
"source": "local",
"message": ADMIN_PROVIDER_QUERY_LOCAL_TEST_MODEL_FAILOVER_MESSAGE,
}))
.into_response()
}
@@ -0,0 +1,140 @@
use super::models::{
build_admin_provider_query_models_response,
build_admin_provider_query_test_model_failover_response,
build_admin_provider_query_test_model_response,
};
use super::shared::{
build_admin_provider_query_bad_request_response, parse_admin_provider_query_body,
provider_query_extract_failover_models, provider_query_extract_model,
provider_query_extract_provider_id, provider_query_extract_request_id,
provider_query_payload_keys, ADMIN_PROVIDER_QUERY_FAILOVER_MODELS_REQUIRED_DETAIL,
ADMIN_PROVIDER_QUERY_MODEL_REQUIRED_DETAIL, ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL,
};
use crate::control::GatewayPublicRequestContext;
use crate::log_ids::short_request_id;
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
http::Response,
};
use tracing::warn;
fn log_admin_provider_query_validation_failure(
request_context: &GatewayPublicRequestContext,
route_kind: &str,
detail: &'static str,
payload: &serde_json::Value,
) {
let provider_id =
provider_query_extract_provider_id(payload).unwrap_or_else(|| "-".to_string());
let model = provider_query_extract_model(payload).unwrap_or_else(|| "-".to_string());
let request_id = provider_query_extract_request_id(payload).unwrap_or_else(|| "-".to_string());
let request_id_for_log = short_request_id(request_id.as_str());
let payload_keys = provider_query_payload_keys(payload);
warn!(
event_name = "admin_provider_query_request_rejected",
log_type = "validation",
route_kind,
path = %request_context.request_path,
request_id = %request_id_for_log,
provider_id = %provider_id,
model = %model,
payload_keys = ?payload_keys,
detail,
"admin provider query request rejected"
);
}
pub(super) async fn maybe_build_local_admin_provider_query_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.control_decision.as_ref() else {
return Ok(None);
};
if decision.route_family.as_deref() != Some("provider_query_manage") {
return Ok(None);
}
if request_context.request_method != http::Method::POST {
return Ok(None);
}
let payload = match parse_admin_provider_query_body(request_body) {
Ok(value) => value,
Err(response) => return Ok(Some(response)),
};
let route_kind = decision.route_kind.as_deref().unwrap_or("query_models");
match route_kind {
"query_models" => Ok(Some(
build_admin_provider_query_models_response(state, &payload).await?,
)),
"test_model" => {
let Some(provider_id) = provider_query_extract_provider_id(&payload) else {
log_admin_provider_query_validation_failure(
request_context,
route_kind,
ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL,
&payload,
);
return Ok(Some(build_admin_provider_query_bad_request_response(
ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL,
)));
};
let Some(model) = provider_query_extract_model(&payload) else {
log_admin_provider_query_validation_failure(
request_context,
route_kind,
ADMIN_PROVIDER_QUERY_MODEL_REQUIRED_DETAIL,
&payload,
);
return Ok(Some(build_admin_provider_query_bad_request_response(
ADMIN_PROVIDER_QUERY_MODEL_REQUIRED_DETAIL,
)));
};
Ok(Some(build_admin_provider_query_test_model_response(
provider_id,
model,
)))
}
"test_model_failover" => {
let Some(provider_id) = provider_query_extract_provider_id(&payload) else {
log_admin_provider_query_validation_failure(
request_context,
route_kind,
ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL,
&payload,
);
return Ok(Some(build_admin_provider_query_bad_request_response(
ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL,
)));
};
let failover_models = provider_query_extract_failover_models(&payload);
if failover_models.is_empty() {
log_admin_provider_query_validation_failure(
request_context,
route_kind,
ADMIN_PROVIDER_QUERY_FAILOVER_MODELS_REQUIRED_DETAIL,
&payload,
);
return Ok(Some(build_admin_provider_query_bad_request_response(
ADMIN_PROVIDER_QUERY_FAILOVER_MODELS_REQUIRED_DETAIL,
)));
}
Ok(Some(
build_admin_provider_query_test_model_failover_response(
provider_id,
failover_models,
),
))
}
_ => Ok(Some(
build_admin_provider_query_models_response(state, &payload).await?,
)),
}
}
@@ -0,0 +1,119 @@
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) const ADMIN_PROVIDER_QUERY_INVALID_JSON_DETAIL: &str = "Invalid JSON request body";
pub(super) const ADMIN_PROVIDER_QUERY_PROVIDER_ID_REQUIRED_DETAIL: &str = "provider_id is required";
pub(super) const ADMIN_PROVIDER_QUERY_MODEL_REQUIRED_DETAIL: &str = "model is required";
pub(super) const ADMIN_PROVIDER_QUERY_FAILOVER_MODELS_REQUIRED_DETAIL: &str =
"failover_models should not be empty";
pub(super) const ADMIN_PROVIDER_QUERY_PROVIDER_NOT_FOUND_DETAIL: &str = "Provider not found";
pub(super) const ADMIN_PROVIDER_QUERY_API_KEY_NOT_FOUND_DETAIL: &str = "API Key not found";
pub(super) const ADMIN_PROVIDER_QUERY_NO_ACTIVE_API_KEY_DETAIL: &str =
"No active API Key found for this provider";
pub(super) const ADMIN_PROVIDER_QUERY_NO_LOCAL_MODELS_DETAIL: &str =
"No models available from local provider catalog";
pub(super) fn build_admin_provider_query_bad_request_response(
detail: &'static str,
) -> Response<Body> {
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response()
}
pub(super) fn build_admin_provider_query_not_found_response(
detail: &'static str,
) -> Response<Body> {
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": detail })),
)
.into_response()
}
pub(super) fn parse_admin_provider_query_body(
request_body: Option<&Bytes>,
) -> Result<serde_json::Value, Response<Body>> {
let Some(raw_body) = request_body else {
return Ok(json!({}));
};
if raw_body.is_empty() {
return Ok(json!({}));
}
serde_json::from_slice::<serde_json::Value>(raw_body).map_err(|_| {
build_admin_provider_query_bad_request_response(ADMIN_PROVIDER_QUERY_INVALID_JSON_DETAIL)
})
}
pub(super) fn provider_query_extract_provider_id(payload: &serde_json::Value) -> Option<String> {
payload
.get("provider_id")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
pub(super) fn provider_query_extract_api_key_id(payload: &serde_json::Value) -> Option<String> {
payload
.get("api_key_id")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
pub(super) fn provider_query_extract_model(payload: &serde_json::Value) -> Option<String> {
payload
.get("model")
.or_else(|| payload.get("model_name"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
pub(super) fn provider_query_extract_failover_models(payload: &serde_json::Value) -> Vec<String> {
if let Some(items) = payload
.get("failover_models")
.or_else(|| payload.get("models"))
.and_then(serde_json::Value::as_array)
{
return items
.iter()
.filter_map(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.collect::<Vec<_>>();
}
provider_query_extract_model(payload)
.into_iter()
.collect::<Vec<_>>()
}
pub(super) fn provider_query_extract_request_id(payload: &serde_json::Value) -> Option<String> {
payload
.get("request_id")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
pub(super) fn provider_query_payload_keys(payload: &serde_json::Value) -> Vec<String> {
let Some(object) = payload.as_object() else {
return Vec::new();
};
let mut keys = object.keys().cloned().collect::<Vec<_>>();
keys.sort();
keys
}
@@ -0,0 +1,7 @@
mod paths;
mod payloads;
mod support;
pub(crate) use self::paths::*;
pub(crate) use self::payloads::*;
pub(crate) use self::support::*;
@@ -0,0 +1,418 @@
pub(crate) fn admin_provider_id_for_keys(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/endpoints/providers/")?
.strip_suffix("/keys")
.map(ToOwned::to_owned)
}
pub(crate) fn is_admin_provider_ops_architectures_root(request_path: &str) -> bool {
matches!(
request_path,
"/api/admin/provider-ops/architectures" | "/api/admin/provider-ops/architectures/"
)
}
pub(crate) fn admin_provider_ops_architecture_id_from_path(request_path: &str) -> Option<String> {
let raw = request_path.strip_prefix("/api/admin/provider-ops/architectures/")?;
let normalized = raw.trim().trim_matches('/');
if normalized.is_empty() || normalized.contains('/') {
None
} else {
Some(normalized.to_string())
}
}
pub(crate) fn admin_provider_id_for_provider_ops_status(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-ops/providers/")?
.strip_suffix("/status")
.map(|value| value.trim().trim_matches('/').to_string())
.filter(|value| !value.is_empty() && !value.contains('/'))
}
pub(crate) fn admin_provider_id_for_provider_ops_config(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-ops/providers/")?
.strip_suffix("/config")
.map(|value| value.trim().trim_matches('/').to_string())
.filter(|value| !value.is_empty() && !value.contains('/'))
}
pub(crate) fn admin_provider_id_for_provider_ops_disconnect(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-ops/providers/")?
.strip_suffix("/disconnect")
.map(|value| value.trim().trim_matches('/').to_string())
.filter(|value| !value.is_empty() && !value.contains('/'))
}
pub(crate) fn admin_provider_id_for_provider_ops_connect(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-ops/providers/")?
.strip_suffix("/connect")
.map(|value| value.trim().trim_matches('/').to_string())
.filter(|value| !value.is_empty() && !value.contains('/'))
}
pub(crate) fn admin_provider_id_for_provider_ops_verify(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-ops/providers/")?
.strip_suffix("/verify")
.map(|value| value.trim().trim_matches('/').to_string())
.filter(|value| !value.is_empty() && !value.contains('/'))
}
pub(crate) fn admin_provider_id_for_provider_ops_balance(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-ops/providers/")?
.strip_suffix("/balance")
.map(|value| value.trim().trim_matches('/').to_string())
.filter(|value| !value.is_empty() && !value.contains('/'))
}
pub(crate) fn admin_provider_id_for_provider_ops_checkin(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-ops/providers/")?
.strip_suffix("/checkin")
.map(|value| value.trim().trim_matches('/').to_string())
.filter(|value| !value.is_empty() && !value.contains('/'))
}
pub(crate) fn is_admin_provider_strategy_strategies_root(request_path: &str) -> bool {
matches!(
request_path,
"/api/admin/provider-strategy/strategies" | "/api/admin/provider-strategy/strategies/"
)
}
pub(crate) fn admin_provider_id_for_provider_strategy_billing(
request_path: &str,
) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-strategy/providers/")?
.strip_suffix("/billing")
.map(|value| value.trim().trim_matches('/').to_string())
.filter(|value| !value.is_empty() && !value.contains('/'))
}
pub(crate) fn admin_provider_id_for_provider_strategy_stats(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-strategy/providers/")?
.strip_suffix("/stats")
.map(|value| value.trim().trim_matches('/').to_string())
.filter(|value| !value.is_empty() && !value.contains('/'))
}
pub(crate) fn admin_provider_id_for_provider_strategy_quota(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-strategy/providers/")?
.strip_suffix("/quota")
.map(|value| value.trim().trim_matches('/').to_string())
.filter(|value| !value.is_empty() && !value.contains('/'))
}
pub(crate) fn admin_provider_ops_action_route_parts(
request_path: &str,
) -> Option<(String, String)> {
let raw = request_path.strip_prefix("/api/admin/provider-ops/providers/")?;
let (provider_id, action_type) = raw.split_once("/actions/")?;
let provider_id = provider_id.trim().trim_matches('/');
let action_type = action_type.trim().trim_matches('/');
if provider_id.is_empty()
|| action_type.is_empty()
|| provider_id.contains('/')
|| action_type.contains('/')
{
None
} else {
Some((provider_id.to_string(), action_type.to_string()))
}
}
pub(crate) fn is_admin_provider_ops_batch_balance_root(request_path: &str) -> bool {
matches!(
request_path,
"/api/admin/provider-ops/batch/balance" | "/api/admin/provider-ops/batch/balance/"
)
}
pub(crate) fn admin_provider_id_for_refresh_quota(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/endpoints/providers/")?
.strip_suffix("/refresh-quota")
.map(ToOwned::to_owned)
}
pub(crate) fn admin_provider_id_for_health_monitor(request_path: &str) -> Option<String> {
let raw = request_path.strip_prefix("/api/admin/providers/")?;
let raw = raw.strip_suffix("/health-monitor")?;
let normalized = raw.trim().trim_matches('/');
if normalized.is_empty() || normalized.contains('/') {
None
} else {
Some(normalized.to_string())
}
}
pub(crate) fn admin_provider_id_for_mapping_preview(request_path: &str) -> Option<String> {
let raw = request_path.strip_prefix("/api/admin/providers/")?;
let raw = raw.strip_suffix("/mapping-preview")?;
let normalized = raw.trim().trim_matches('/');
if normalized.is_empty() || normalized.contains('/') {
None
} else {
Some(normalized.to_string())
}
}
pub(crate) fn admin_provider_id_for_pool_status(request_path: &str) -> Option<String> {
admin_provider_id_for_suffix(request_path, "/pool-status")
}
pub(crate) fn admin_provider_pool_key_route_parts(
request_path: &str,
marker: &str,
) -> Option<(String, String)> {
let raw = request_path.strip_prefix("/api/admin/providers/")?;
let (provider_id, key_id) = raw.split_once(marker)?;
let provider_id = provider_id.trim().trim_matches('/');
let key_id = key_id.trim().trim_matches('/');
if provider_id.is_empty()
|| key_id.is_empty()
|| provider_id.contains('/')
|| key_id.contains('/')
{
None
} else {
Some((provider_id.to_string(), key_id.to_string()))
}
}
pub(crate) fn admin_provider_clear_pool_cooldown_parts(
request_path: &str,
) -> Option<(String, String)> {
admin_provider_pool_key_route_parts(request_path, "/pool/clear-cooldown/")
}
pub(crate) fn admin_provider_reset_pool_cost_parts(request_path: &str) -> Option<(String, String)> {
admin_provider_pool_key_route_parts(request_path, "/pool/reset-cost/")
}
pub(crate) fn admin_provider_delete_task_parts(request_path: &str) -> Option<(String, String)> {
let raw = request_path.strip_prefix("/api/admin/providers/")?;
let (provider_id, task_id) = raw.split_once("/delete-task/")?;
let provider_id = provider_id.trim().trim_matches('/');
let task_id = task_id.trim().trim_matches('/');
if provider_id.is_empty()
|| task_id.is_empty()
|| provider_id.contains('/')
|| task_id.contains('/')
{
None
} else {
Some((provider_id.to_string(), task_id.to_string()))
}
}
pub(crate) fn admin_provider_id_for_summary(request_path: &str) -> Option<String> {
let raw = request_path.strip_prefix("/api/admin/providers/")?;
let raw = raw.strip_suffix("/summary")?;
let normalized = raw.trim().trim_matches('/');
if normalized.is_empty() || normalized.contains('/') {
None
} else {
Some(normalized.to_string())
}
}
pub(crate) fn admin_provider_id_for_manage_path(request_path: &str) -> Option<String> {
let raw = request_path.strip_prefix("/api/admin/providers/")?;
let normalized = raw.trim().trim_matches('/');
if normalized.is_empty() || normalized.contains('/') {
None
} else {
Some(normalized.to_string())
}
}
pub(crate) fn is_admin_providers_root(request_path: &str) -> bool {
matches!(
request_path,
"/api/admin/providers" | "/api/admin/providers/"
)
}
pub(crate) fn admin_provider_id_for_models_list(request_path: &str) -> Option<String> {
admin_provider_id_for_suffix(request_path, "/models")
}
pub(crate) fn admin_provider_id_for_suffix(request_path: &str, suffix: &str) -> Option<String> {
let raw = request_path.strip_prefix("/api/admin/providers/")?;
let raw = raw.strip_suffix(suffix)?;
let normalized = raw.trim().trim_matches('/');
if normalized.is_empty() || normalized.contains('/') {
None
} else {
Some(normalized.to_string())
}
}
pub(crate) fn admin_provider_model_route_parts(request_path: &str) -> Option<(String, String)> {
let raw = request_path.strip_prefix("/api/admin/providers/")?;
let (provider_id, model_id) = raw.split_once("/models/")?;
let provider_id = provider_id.trim().trim_matches('/');
let model_id = model_id.trim().trim_matches('/');
if provider_id.is_empty()
|| model_id.is_empty()
|| provider_id.contains('/')
|| model_id.contains('/')
{
None
} else {
Some((provider_id.to_string(), model_id.to_string()))
}
}
pub(crate) fn admin_provider_models_batch_path(request_path: &str) -> Option<String> {
admin_provider_id_for_suffix(request_path, "/models/batch")
}
pub(crate) fn admin_provider_available_source_models_path(request_path: &str) -> Option<String> {
admin_provider_id_for_suffix(request_path, "/available-source-models")
}
pub(crate) fn admin_provider_assign_global_models_path(request_path: &str) -> Option<String> {
admin_provider_id_for_suffix(request_path, "/assign-global-models")
}
pub(crate) fn admin_provider_import_models_path(request_path: &str) -> Option<String> {
admin_provider_id_for_suffix(request_path, "/import-from-upstream")
}
pub(crate) fn admin_reveal_key_id(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/endpoints/keys/")?
.strip_suffix("/reveal")
.map(ToOwned::to_owned)
}
pub(crate) fn admin_export_key_id(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/endpoints/keys/")?
.strip_suffix("/export")
.map(ToOwned::to_owned)
}
pub(crate) fn admin_clear_oauth_invalid_key_id(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/endpoints/keys/")?
.strip_suffix("/clear-oauth-invalid")
.map(ToOwned::to_owned)
}
pub(crate) fn admin_update_key_id(request_path: &str) -> Option<String> {
let key_id = request_path.strip_prefix("/api/admin/endpoints/keys/")?;
(!key_id.is_empty() && !key_id.contains('/')).then_some(key_id.to_string())
}
pub(crate) fn admin_provider_oauth_start_key_id(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-oauth/keys/")?
.strip_suffix("/start")
.filter(|key_id| !key_id.is_empty() && !key_id.contains('/'))
.map(ToOwned::to_owned)
}
pub(crate) fn admin_provider_oauth_start_provider_id(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-oauth/providers/")?
.strip_suffix("/start")
.filter(|provider_id| !provider_id.is_empty() && !provider_id.contains('/'))
.map(ToOwned::to_owned)
}
pub(crate) fn admin_provider_oauth_complete_key_id(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-oauth/keys/")?
.strip_suffix("/complete")
.filter(|key_id| !key_id.is_empty() && !key_id.contains('/'))
.map(ToOwned::to_owned)
}
pub(crate) fn admin_provider_oauth_refresh_key_id(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-oauth/keys/")?
.strip_suffix("/refresh")
.filter(|key_id| !key_id.is_empty() && !key_id.contains('/'))
.map(ToOwned::to_owned)
}
pub(crate) fn admin_provider_oauth_complete_provider_id(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-oauth/providers/")?
.strip_suffix("/complete")
.filter(|provider_id| !provider_id.is_empty() && !provider_id.contains('/'))
.map(ToOwned::to_owned)
}
pub(crate) fn admin_provider_oauth_import_provider_id(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-oauth/providers/")?
.strip_suffix("/import-refresh-token")
.filter(|provider_id| !provider_id.is_empty() && !provider_id.contains('/'))
.map(ToOwned::to_owned)
}
pub(crate) fn admin_provider_oauth_batch_import_provider_id(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-oauth/providers/")?
.strip_suffix("/batch-import")
.filter(|provider_id| !provider_id.is_empty() && !provider_id.contains('/'))
.map(ToOwned::to_owned)
}
pub(crate) fn admin_provider_oauth_batch_import_task_provider_id(
request_path: &str,
) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-oauth/providers/")?
.strip_suffix("/batch-import/tasks")
.filter(|provider_id| !provider_id.is_empty() && !provider_id.contains('/'))
.map(ToOwned::to_owned)
}
pub(crate) fn admin_provider_oauth_batch_import_task_path(
request_path: &str,
) -> Option<(String, String)> {
let suffix = request_path
.strip_prefix("/api/admin/provider-oauth/providers/")?
.strip_suffix("/")
.unwrap_or(request_path.strip_prefix("/api/admin/provider-oauth/providers/")?);
let (provider_id, task_path) = suffix.split_once("/batch-import/tasks/")?;
if provider_id.is_empty()
|| provider_id.contains('/')
|| task_path.is_empty()
|| task_path.contains('/')
{
return None;
}
Some((provider_id.to_string(), task_path.to_string()))
}
pub(crate) fn admin_provider_oauth_device_authorize_provider_id(
request_path: &str,
) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-oauth/providers/")?
.strip_suffix("/device-authorize")
.filter(|provider_id| !provider_id.is_empty() && !provider_id.contains('/'))
.map(ToOwned::to_owned)
}
pub(crate) fn admin_provider_oauth_device_poll_provider_id(request_path: &str) -> Option<String> {
request_path
.strip_prefix("/api/admin/provider-oauth/providers/")?
.strip_suffix("/device-poll")
.filter(|provider_id| !provider_id.is_empty() && !provider_id.contains('/'))
.map(ToOwned::to_owned)
}
@@ -0,0 +1,263 @@
use serde::Deserialize;
#[derive(Debug, Deserialize)]
pub(crate) struct AdminProviderKeyCreateRequest {
#[serde(default)]
pub(crate) api_formats: Option<Vec<String>>,
#[serde(default)]
pub(crate) api_key: Option<String>,
#[serde(default)]
pub(crate) auth_type: Option<String>,
#[serde(default)]
pub(crate) auth_config: Option<serde_json::Value>,
pub(crate) name: String,
#[serde(default)]
pub(crate) rate_multipliers: Option<serde_json::Value>,
#[serde(default)]
pub(crate) internal_priority: Option<i32>,
#[serde(default)]
pub(crate) rpm_limit: Option<u32>,
#[serde(default)]
pub(crate) allowed_models: Option<Vec<String>>,
#[serde(default)]
pub(crate) capabilities: Option<serde_json::Value>,
#[serde(default)]
pub(crate) cache_ttl_minutes: Option<i32>,
#[serde(default)]
pub(crate) max_probe_interval_minutes: Option<i32>,
#[serde(default)]
pub(crate) note: Option<String>,
#[serde(default)]
pub(crate) auto_fetch_models: Option<bool>,
#[serde(default)]
pub(crate) locked_models: Option<Vec<String>>,
#[serde(default)]
pub(crate) model_include_patterns: Option<Vec<String>>,
#[serde(default)]
pub(crate) model_exclude_patterns: Option<Vec<String>>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct AdminProviderKeyUpdateRequest {
#[serde(default)]
pub(crate) api_formats: Option<Vec<String>>,
#[serde(default)]
pub(crate) api_key: Option<String>,
#[serde(default)]
pub(crate) auth_type: Option<String>,
#[serde(default)]
pub(crate) auth_config: Option<serde_json::Value>,
#[serde(default)]
pub(crate) name: Option<String>,
#[serde(default)]
pub(crate) rate_multipliers: Option<serde_json::Value>,
#[serde(default)]
pub(crate) internal_priority: Option<i32>,
#[serde(default)]
pub(crate) global_priority_by_format: Option<serde_json::Value>,
#[serde(default)]
pub(crate) rpm_limit: Option<u32>,
#[serde(default)]
pub(crate) allowed_models: Option<Vec<String>>,
#[serde(default)]
pub(crate) capabilities: Option<serde_json::Value>,
#[serde(default)]
pub(crate) cache_ttl_minutes: Option<i32>,
#[serde(default)]
pub(crate) max_probe_interval_minutes: Option<i32>,
#[serde(default)]
pub(crate) is_active: Option<bool>,
#[serde(default)]
pub(crate) note: Option<String>,
#[serde(default)]
pub(crate) auto_fetch_models: Option<bool>,
#[serde(default)]
pub(crate) locked_models: Option<Vec<String>>,
#[serde(default)]
pub(crate) model_include_patterns: Option<Vec<String>>,
#[serde(default)]
pub(crate) model_exclude_patterns: Option<Vec<String>>,
#[serde(default)]
pub(crate) proxy: Option<serde_json::Value>,
#[serde(default)]
pub(crate) fingerprint: Option<serde_json::Value>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct AdminProviderKeyBatchDeleteRequest {
pub(crate) ids: Vec<String>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct AdminProviderQuotaRefreshRequest {
#[serde(default)]
pub(crate) key_ids: Option<Vec<String>>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct AdminProviderCreateRequest {
pub(crate) name: String,
#[serde(default)]
pub(crate) provider_type: Option<String>,
#[serde(default)]
pub(crate) description: Option<String>,
#[serde(default)]
pub(crate) website: Option<String>,
#[serde(default)]
pub(crate) billing_type: Option<String>,
#[serde(default)]
pub(crate) monthly_quota_usd: Option<f64>,
#[serde(default)]
pub(crate) quota_reset_day: Option<u64>,
#[serde(default)]
pub(crate) quota_last_reset_at: Option<String>,
#[serde(default)]
pub(crate) quota_expires_at: Option<String>,
#[serde(default)]
pub(crate) provider_priority: Option<i32>,
#[serde(default)]
pub(crate) keep_priority_on_conversion: Option<bool>,
#[serde(default)]
pub(crate) is_active: Option<bool>,
#[serde(default)]
pub(crate) concurrent_limit: Option<i32>,
#[serde(default)]
pub(crate) max_retries: Option<i32>,
#[serde(default)]
pub(crate) proxy: Option<serde_json::Value>,
#[serde(default)]
pub(crate) stream_first_byte_timeout: Option<f64>,
#[serde(default)]
pub(crate) request_timeout: Option<f64>,
#[serde(default)]
pub(crate) pool_advanced: Option<serde_json::Value>,
#[serde(default)]
pub(crate) claude_code_advanced: Option<serde_json::Value>,
#[serde(default)]
pub(crate) failover_rules: Option<serde_json::Value>,
#[serde(default)]
pub(crate) config: Option<serde_json::Value>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct AdminProviderUpdateRequest {
#[serde(default)]
pub(crate) name: Option<String>,
#[serde(default)]
pub(crate) provider_type: Option<String>,
#[serde(default)]
pub(crate) description: Option<String>,
#[serde(default)]
pub(crate) website: Option<String>,
#[serde(default)]
pub(crate) billing_type: Option<String>,
#[serde(default)]
pub(crate) monthly_quota_usd: Option<f64>,
#[serde(default)]
pub(crate) quota_reset_day: Option<u64>,
#[serde(default)]
pub(crate) quota_last_reset_at: Option<String>,
#[serde(default)]
pub(crate) quota_expires_at: Option<String>,
#[serde(default)]
pub(crate) provider_priority: Option<i32>,
#[serde(default)]
pub(crate) keep_priority_on_conversion: Option<bool>,
#[serde(default)]
pub(crate) is_active: Option<bool>,
#[serde(default)]
pub(crate) concurrent_limit: Option<i32>,
#[serde(default)]
pub(crate) max_retries: Option<i32>,
#[serde(default)]
pub(crate) proxy: Option<serde_json::Value>,
#[serde(default)]
pub(crate) stream_first_byte_timeout: Option<f64>,
#[serde(default)]
pub(crate) request_timeout: Option<f64>,
#[serde(default)]
pub(crate) pool_advanced: Option<serde_json::Value>,
#[serde(default)]
pub(crate) claude_code_advanced: Option<serde_json::Value>,
#[serde(default)]
pub(crate) failover_rules: Option<serde_json::Value>,
#[serde(default)]
pub(crate) enable_format_conversion: Option<bool>,
#[serde(default)]
pub(crate) config: Option<serde_json::Value>,
}
pub(crate) const CODEX_WHAM_USAGE_URL: &str = "https://chatgpt.com/backend-api/wham/usage";
pub(crate) const KIRO_USAGE_LIMITS_PATH: &str = "/getUsageLimits";
pub(crate) const KIRO_USAGE_SDK_VERSION: &str = "1.0.0";
pub(crate) const ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH: &str = "/v1internal:fetchAvailableModels";
pub(crate) const OAUTH_ACCOUNT_BLOCK_PREFIX: &str = "[ACCOUNT_BLOCK] ";
pub(crate) const OAUTH_REFRESH_FAILED_PREFIX: &str = "[REFRESH_FAILED] ";
pub(crate) const OAUTH_EXPIRED_PREFIX: &str = "[OAUTH_EXPIRED] ";
pub(crate) const OAUTH_REQUEST_FAILED_PREFIX: &str = "[REQUEST_FAILED] ";
#[derive(Debug, Deserialize)]
pub(crate) struct AdminProviderModelCreateRequest {
pub(crate) provider_model_name: String,
#[serde(default)]
pub(crate) provider_model_mappings: Option<serde_json::Value>,
pub(crate) global_model_id: String,
#[serde(default)]
pub(crate) price_per_request: Option<f64>,
#[serde(default)]
pub(crate) tiered_pricing: Option<serde_json::Value>,
#[serde(default)]
pub(crate) supports_vision: Option<bool>,
#[serde(default)]
pub(crate) supports_function_calling: Option<bool>,
#[serde(default)]
pub(crate) supports_streaming: Option<bool>,
#[serde(default)]
pub(crate) supports_extended_thinking: Option<bool>,
#[serde(default)]
pub(crate) is_active: Option<bool>,
#[serde(default)]
pub(crate) config: Option<serde_json::Value>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct AdminProviderModelUpdateRequest {
#[serde(default)]
pub(crate) provider_model_name: Option<String>,
#[serde(default)]
pub(crate) provider_model_mappings: Option<serde_json::Value>,
#[serde(default)]
pub(crate) global_model_id: Option<String>,
#[serde(default)]
pub(crate) price_per_request: Option<f64>,
#[serde(default)]
pub(crate) tiered_pricing: Option<serde_json::Value>,
#[serde(default)]
pub(crate) supports_vision: Option<bool>,
#[serde(default)]
pub(crate) supports_function_calling: Option<bool>,
#[serde(default)]
pub(crate) supports_streaming: Option<bool>,
#[serde(default)]
pub(crate) supports_extended_thinking: Option<bool>,
#[serde(default)]
pub(crate) is_active: Option<bool>,
#[serde(default)]
pub(crate) is_available: Option<bool>,
#[serde(default)]
pub(crate) config: Option<serde_json::Value>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct AdminBatchAssignGlobalModelsRequest {
pub(crate) global_model_ids: Vec<String>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct AdminImportProviderModelsRequest {
pub(crate) model_ids: Vec<String>,
#[serde(default)]
pub(crate) tiered_pricing: Option<serde_json::Value>,
#[serde(default)]
pub(crate) price_per_request: Option<f64>,
}
@@ -0,0 +1,50 @@
use crate::{AppState, LocalProviderDeleteTaskState};
use serde_json::json;
use std::collections::BTreeMap;
pub(crate) const ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_KEYS: usize = 200;
pub(crate) const ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_MODELS: usize = 500;
pub(crate) const ADMIN_PROVIDER_MAPPING_PREVIEW_FETCH_LIMIT: usize = 10_000;
pub(crate) const ADMIN_PROVIDER_POOL_SCAN_BATCH: u64 = 200;
pub(crate) const ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL: &str =
"Admin provider OAuth data unavailable";
#[derive(Debug, Clone, Copy)]
pub(crate) struct AdminProviderPoolConfig {
pub(crate) lru_enabled: bool,
pub(crate) cost_window_seconds: u64,
pub(crate) cost_limit_per_key_tokens: Option<u64>,
}
#[derive(Debug, Default)]
pub(crate) struct AdminProviderPoolRuntimeState {
pub(crate) total_sticky_sessions: usize,
pub(crate) sticky_sessions_by_key: BTreeMap<String, usize>,
pub(crate) cooldown_reason_by_key: BTreeMap<String, String>,
pub(crate) cooldown_ttl_by_key: BTreeMap<String, u64>,
pub(crate) cost_window_usage_by_key: BTreeMap<String, u64>,
pub(crate) lru_score_by_key: BTreeMap<String, f64>,
}
pub(crate) fn build_admin_provider_delete_task_payload(
task: &LocalProviderDeleteTaskState,
) -> serde_json::Value {
json!({
"task_id": task.task_id,
"provider_id": task.provider_id,
"status": task.status,
"stage": task.stage,
"total_keys": task.total_keys,
"deleted_keys": task.deleted_keys,
"total_endpoints": task.total_endpoints,
"deleted_endpoints": task.deleted_endpoints,
"message": task.message,
})
}
pub(crate) fn put_admin_provider_delete_task(
state: &AppState,
task: &LocalProviderDeleteTaskState,
) {
state.put_provider_delete_task(task.clone());
}
@@ -0,0 +1,259 @@
use super::super::write::{normalize_provider_billing_type, parse_optional_rfc3339_unix_secs};
use super::shared::admin_provider_strategy_provider_not_found_response;
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
use crate::{AppState, GatewayError};
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde::Deserialize;
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
#[derive(Debug, Deserialize)]
pub(super) struct AdminProviderStrategyBillingRequest {
pub(super) billing_type: String,
#[serde(default)]
pub(super) monthly_quota_usd: Option<f64>,
#[serde(default = "default_provider_strategy_quota_reset_day")]
pub(super) quota_reset_day: u64,
#[serde(default)]
pub(super) quota_last_reset_at: Option<String>,
#[serde(default)]
pub(super) quota_expires_at: Option<String>,
#[serde(default)]
pub(super) rpm_limit: Option<i32>,
#[serde(default = "default_provider_strategy_provider_priority")]
pub(super) provider_priority: i32,
}
fn default_provider_strategy_quota_reset_day() -> u64 {
30
}
fn default_provider_strategy_provider_priority() -> i32 {
100
}
pub(super) fn build_provider_strategy_list_response() -> Response<Body> {
Json(json!({
"strategies": [{
"name": "sticky_priority",
"priority": 110,
"version": "1.0.0",
"description": "粘性优先级负载均衡策略,正常时始终使用同一提供商",
"author": "System",
}],
"total": 1,
}))
.into_response()
}
pub(super) async fn build_provider_strategy_update_billing_response(
state: &AppState,
provider_id: String,
payload: AdminProviderStrategyBillingRequest,
) -> Result<Response<Body>, GatewayError> {
let Some(existing) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(admin_provider_strategy_provider_not_found_response());
};
let billing_type = match normalize_provider_billing_type(&payload.billing_type) {
Ok(value) => value,
Err(message) => {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": message })),
)
.into_response());
}
};
if payload
.monthly_quota_usd
.is_some_and(|value| !value.is_finite() || value < 0.0)
{
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "monthly_quota_usd 必须是非负数" })),
)
.into_response());
}
if !(1..=365).contains(&payload.quota_reset_day) {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "quota_reset_day 必须是 1 到 365 之间的整数" })),
)
.into_response());
}
if !(0..=10_000).contains(&payload.provider_priority) {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "provider_priority 必须在 0 到 10000 之间" })),
)
.into_response());
}
let quota_last_reset_at_unix_secs = match payload.quota_last_reset_at.as_deref() {
Some(value) => match parse_optional_rfc3339_unix_secs(value, "quota_last_reset_at") {
Ok(value) => Some(value),
Err(message) => {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": message })),
)
.into_response());
}
},
None => existing.quota_last_reset_at_unix_secs,
};
let quota_expires_at_unix_secs = match payload.quota_expires_at.as_deref() {
Some(value) => match parse_optional_rfc3339_unix_secs(value, "quota_expires_at") {
Ok(value) => Some(value),
Err(message) => {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": message })),
)
.into_response());
}
},
None => existing.quota_expires_at_unix_secs,
};
let synced_monthly_used_usd = match quota_last_reset_at_unix_secs {
Some(quota_last_reset_at_unix_secs) if state.has_usage_data_reader() => Some(
state
.summarize_provider_usage_since(&provider_id, quota_last_reset_at_unix_secs)
.await?
.total_cost_usd,
),
_ => existing.monthly_used_usd,
};
let _ignored_rpm_limit = payload.rpm_limit;
let updated = existing
.clone()
.with_billing_fields(
Some(billing_type.clone()),
payload.monthly_quota_usd,
synced_monthly_used_usd,
Some(payload.quota_reset_day),
quota_last_reset_at_unix_secs,
quota_expires_at_unix_secs,
)
.with_routing_fields(payload.provider_priority);
let Some(updated) = state.update_provider_catalog_provider(&updated).await? else {
return Ok(admin_provider_strategy_provider_not_found_response());
};
Ok(Json(json!({
"message": "Provider billing config updated successfully",
"provider": {
"id": updated.id,
"name": updated.name,
"billing_type": billing_type,
"provider_priority": updated.provider_priority,
},
}))
.into_response())
}
pub(super) async fn build_provider_strategy_stats_response(
state: &AppState,
provider_id: String,
hours: u64,
) -> Result<Response<Body>, GatewayError> {
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let since_unix_secs = now_unix_secs.saturating_sub(hours.saturating_mul(3600));
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(admin_provider_strategy_provider_not_found_response());
};
let summary = state
.summarize_provider_usage_since(&provider_id, since_unix_secs)
.await?;
let monthly_used_usd = provider.monthly_used_usd.unwrap_or(0.0);
let quota_remaining_usd = provider
.monthly_quota_usd
.map(|value| value - monthly_used_usd);
let success_rate = if summary.total_requests > 0 {
summary.successful_requests as f64 / summary.total_requests as f64
} else {
0.0
};
Ok(Json(json!({
"provider_id": provider_id,
"provider_name": provider.name,
"period_hours": hours,
"billing_info": {
"billing_type": provider.billing_type,
"monthly_quota_usd": provider.monthly_quota_usd,
"monthly_used_usd": monthly_used_usd,
"quota_remaining_usd": quota_remaining_usd,
"quota_expires_at": provider.quota_expires_at_unix_secs.and_then(unix_secs_to_rfc3339),
},
"usage_stats": {
"total_requests": summary.total_requests,
"successful_requests": summary.successful_requests,
"failed_requests": summary.failed_requests,
"success_rate": success_rate,
"avg_response_time_ms": (summary.avg_response_time_ms * 100.0).round() / 100.0,
"total_cost_usd": (summary.total_cost_usd * 10_000.0).round() / 10_000.0,
},
}))
.into_response())
}
pub(super) async fn build_provider_strategy_reset_quota_response(
state: &AppState,
provider_id: String,
) -> Result<Response<Body>, GatewayError> {
let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await?
.into_iter()
.next()
else {
return Ok(admin_provider_strategy_provider_not_found_response());
};
if provider.billing_type.as_deref() != Some("monthly_quota") {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "Only monthly quota providers can be reset" })),
)
.into_response());
}
let previous_used = provider.monthly_used_usd.unwrap_or(0.0);
let mut updated = provider.clone();
updated.monthly_used_usd = Some(0.0);
let Some(updated) = state.update_provider_catalog_provider(&updated).await? else {
return Ok(admin_provider_strategy_provider_not_found_response());
};
Ok(Json(json!({
"message": "Provider quota reset successfully",
"provider_name": updated.name,
"previous_used": previous_used,
"current_used": 0.0,
}))
.into_response())
}
@@ -0,0 +1,17 @@
use crate::control::GatewayPublicRequestContext;
use crate::{AppState, GatewayError};
use axum::body::{Body, Bytes};
use axum::http::Response;
mod builders;
mod routes;
mod shared;
pub(crate) async fn maybe_build_local_admin_provider_strategy_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
routes::maybe_build_local_admin_provider_strategy_response(state, request_context, request_body)
.await
}
@@ -0,0 +1,136 @@
use super::builders::{
build_provider_strategy_list_response, build_provider_strategy_reset_quota_response,
build_provider_strategy_stats_response, build_provider_strategy_update_billing_response,
AdminProviderStrategyBillingRequest,
};
use super::shared::{
admin_provider_strategy_data_unavailable_response,
admin_provider_strategy_dispatcher_not_found_response,
admin_provider_strategy_provider_not_found_response,
ADMIN_PROVIDER_STRATEGY_DATA_UNAVAILABLE_DETAIL,
ADMIN_PROVIDER_STRATEGY_STATS_DATA_UNAVAILABLE_DETAIL,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::admin::provider::shared::{
admin_provider_id_for_provider_strategy_billing, admin_provider_id_for_provider_strategy_quota,
admin_provider_id_for_provider_strategy_stats, is_admin_provider_strategy_strategies_root,
};
use crate::handlers::admin::shared::query_param_value;
use crate::{AppState, GatewayError};
use axum::{
body::{Body, Bytes},
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) async fn maybe_build_local_admin_provider_strategy_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
request_body: Option<&Bytes>,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(decision) = request_context.control_decision.as_ref() else {
return Ok(None);
};
if decision.route_family.as_deref() != Some("provider_strategy_manage") {
return Ok(None);
}
if decision.route_kind.as_deref() == Some("list_strategies")
&& request_context.request_method == http::Method::GET
&& is_admin_provider_strategy_strategies_root(&request_context.request_path)
{
return Ok(Some(build_provider_strategy_list_response()));
}
if decision.route_kind.as_deref() == Some("update_provider_billing")
&& request_context.request_method == http::Method::PUT
{
if !state.has_provider_catalog_data_reader() || !state.has_provider_catalog_data_writer() {
return Ok(Some(admin_provider_strategy_data_unavailable_response(
ADMIN_PROVIDER_STRATEGY_DATA_UNAVAILABLE_DETAIL,
)));
}
let Some(provider_id) =
admin_provider_id_for_provider_strategy_billing(&request_context.request_path)
else {
return Ok(Some(admin_provider_strategy_provider_not_found_response()));
};
let Some(request_body) = request_body else {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求体不能为空" })),
)
.into_response(),
));
};
let payload =
match serde_json::from_slice::<AdminProviderStrategyBillingRequest>(request_body) {
Ok(payload) => payload,
Err(_) => {
return Ok(Some(
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "请求数据验证失败" })),
)
.into_response(),
))
}
};
return Ok(Some(
build_provider_strategy_update_billing_response(state, provider_id, payload).await?,
));
}
if decision.route_kind.as_deref() == Some("get_provider_stats")
&& request_context.request_method == http::Method::GET
{
if !state.has_provider_catalog_data_reader() || !state.has_usage_data_reader() {
return Ok(Some(admin_provider_strategy_data_unavailable_response(
ADMIN_PROVIDER_STRATEGY_STATS_DATA_UNAVAILABLE_DETAIL,
)));
}
let Some(provider_id) =
admin_provider_id_for_provider_strategy_stats(&request_context.request_path)
else {
return Ok(Some(admin_provider_strategy_provider_not_found_response()));
};
let hours = query_param_value(request_context.request_query_string.as_deref(), "hours")
.and_then(|value| value.parse::<u64>().ok())
.filter(|value| *value > 0)
.unwrap_or(24);
return Ok(Some(
build_provider_strategy_stats_response(state, provider_id, hours).await?,
));
}
if decision.route_kind.as_deref() == Some("reset_provider_quota")
&& request_context.request_method == http::Method::DELETE
{
if !state.has_provider_catalog_data_reader() || !state.has_provider_catalog_data_writer() {
return Ok(Some(admin_provider_strategy_data_unavailable_response(
ADMIN_PROVIDER_STRATEGY_DATA_UNAVAILABLE_DETAIL,
)));
}
let Some(provider_id) =
admin_provider_id_for_provider_strategy_quota(&request_context.request_path)
else {
return Ok(Some(admin_provider_strategy_provider_not_found_response()));
};
return Ok(Some(
build_provider_strategy_reset_quota_response(state, provider_id).await?,
));
}
Ok(Some(admin_provider_strategy_dispatcher_not_found_response()))
}
@@ -0,0 +1,38 @@
use super::super::build_proxy_error_response;
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
pub(super) const ADMIN_PROVIDER_STRATEGY_DATA_UNAVAILABLE_DETAIL: &str =
"Admin provider strategy data unavailable";
pub(super) const ADMIN_PROVIDER_STRATEGY_STATS_DATA_UNAVAILABLE_DETAIL: &str =
"Admin provider strategy stats data unavailable";
pub(super) fn admin_provider_strategy_data_unavailable_response(detail: &str) -> Response<Body> {
build_proxy_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"data_unavailable",
detail,
Some(json!({ "error": detail })),
)
}
pub(super) fn admin_provider_strategy_provider_not_found_response() -> Response<Body> {
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider not found" })),
)
.into_response()
}
pub(super) fn admin_provider_strategy_dispatcher_not_found_response() -> Response<Body> {
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "Provider strategy route not found" })),
)
.into_response()
}
@@ -0,0 +1,687 @@
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
use crate::handlers::public::{
provider_key_api_formats, request_candidate_event_unix_secs, request_candidate_status_label,
};
use crate::AppState;
use aether_data_contracts::repository::candidates::{
RequestCandidateStatus, StoredRequestCandidate,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use serde_json::json;
use std::collections::{BTreeMap, BTreeSet};
use std::time::{SystemTime, UNIX_EPOCH};
pub(crate) async fn build_admin_providers_payload(
state: &AppState,
skip: usize,
limit: usize,
is_active: Option<bool>,
) -> Option<serde_json::Value> {
if !state.has_provider_catalog_data_reader() {
return None;
}
let active_only = is_active.unwrap_or(false);
let mut providers = state
.list_provider_catalog_providers(active_only)
.await
.ok()
.unwrap_or_default();
if matches!(is_active, Some(false)) {
providers.retain(|provider| !provider.is_active);
}
providers.sort_by(|left, right| {
left.provider_priority
.cmp(&right.provider_priority)
.then_with(|| left.name.cmp(&right.name))
});
let providers = providers
.into_iter()
.skip(skip)
.take(limit)
.collect::<Vec<_>>();
let provider_ids = providers
.iter()
.map(|provider| provider.id.clone())
.collect::<Vec<_>>();
let endpoints = if provider_ids.is_empty() {
Vec::new()
} else {
state
.list_provider_catalog_endpoints_by_provider_ids(&provider_ids)
.await
.ok()
.unwrap_or_default()
};
let keys = if provider_ids.is_empty() {
Vec::new()
} else {
state
.list_provider_catalog_keys_by_provider_ids(&provider_ids)
.await
.ok()
.unwrap_or_default()
};
let first_endpoint_by_provider = endpoints
.into_iter()
.filter(|endpoint| endpoint.is_active)
.fold(
BTreeMap::<String, StoredProviderCatalogEndpoint>::new(),
|mut acc, endpoint| {
acc.entry(endpoint.provider_id.clone()).or_insert(endpoint);
acc
},
);
let has_any_key_by_provider =
keys.into_iter()
.fold(BTreeSet::<String>::new(), |mut acc, key| {
acc.insert(key.provider_id);
acc
});
Some(serde_json::Value::Array(
providers
.into_iter()
.map(|provider| {
let provider_id = provider.id.clone();
let endpoint = first_endpoint_by_provider.get(&provider_id);
json!({
"id": provider_id.clone(),
"name": provider.name,
"api_format": endpoint.map(|item| item.api_format.clone()),
"base_url": endpoint.map(|item| item.base_url.clone()),
"api_key": has_any_key_by_provider.contains(&provider_id).then_some("***"),
"priority": provider.provider_priority,
"is_active": provider.is_active,
"created_at": provider.created_at_unix_secs.and_then(unix_secs_to_rfc3339),
"updated_at": provider.updated_at_unix_secs.and_then(unix_secs_to_rfc3339),
})
})
.collect(),
))
}
fn json_truthy(value: &serde_json::Value) -> bool {
match value {
serde_json::Value::Null => false,
serde_json::Value::Bool(value) => *value,
serde_json::Value::Number(value) => value.as_f64().is_some_and(|value| value != 0.0),
serde_json::Value::String(value) => !value.trim().is_empty(),
serde_json::Value::Array(value) => !value.is_empty(),
serde_json::Value::Object(value) => !value.is_empty(),
}
}
fn endpoint_timestamp_or_now(value: Option<u64>, now_unix_secs: u64) -> serde_json::Value {
unix_secs_to_rfc3339(value.unwrap_or(now_unix_secs))
.map(serde_json::Value::String)
.unwrap_or(serde_json::Value::Null)
}
pub(crate) fn build_admin_provider_summary_value(
provider: &StoredProviderCatalogProvider,
endpoints: &[StoredProviderCatalogEndpoint],
keys: &[StoredProviderCatalogKey],
quota_snapshot: Option<&aether_data_contracts::repository::quota::StoredProviderQuotaSnapshot>,
model_stats: Option<
&aether_data_contracts::repository::global_models::StoredProviderModelStats,
>,
active_global_model_ids: Vec<String>,
now_unix_secs: u64,
) -> serde_json::Value {
let total_endpoints = endpoints.len();
let active_endpoints = endpoints
.iter()
.filter(|endpoint| endpoint.is_active)
.count();
let total_keys = keys.len();
let active_keys = keys.iter().filter(|key| key.is_active).count();
let total_models = model_stats
.map(|stats| stats.total_models as usize)
.unwrap_or(0);
let active_models = model_stats
.map(|stats| stats.active_models as usize)
.unwrap_or(0);
let api_formats = endpoints
.iter()
.map(|endpoint| endpoint.api_format.clone())
.collect::<Vec<_>>();
let format_to_endpoint_id = endpoints
.iter()
.map(|endpoint| (endpoint.api_format.clone(), endpoint.id.clone()))
.collect::<BTreeMap<_, _>>();
let mut keys_by_endpoint = BTreeMap::<String, Vec<&StoredProviderCatalogKey>>::new();
for endpoint in endpoints {
keys_by_endpoint.entry(endpoint.id.clone()).or_default();
}
for key in keys {
for api_format in provider_key_api_formats(key) {
if let Some(endpoint_id) = format_to_endpoint_id.get(&api_format) {
keys_by_endpoint
.entry(endpoint_id.clone())
.or_default()
.push(key);
}
}
}
let mut endpoint_health_scores = Vec::with_capacity(endpoints.len());
let endpoint_health_details = endpoints
.iter()
.map(|endpoint| {
let endpoint_keys = keys_by_endpoint
.get(&endpoint.id)
.cloned()
.unwrap_or_default();
let health_score = if endpoint_keys.is_empty() {
1.0
} else {
let mut scores = Vec::new();
for key in &endpoint_keys {
let score = key
.health_by_format
.as_ref()
.and_then(|value| value.get(&endpoint.api_format))
.and_then(|value| value.get("health_score"))
.and_then(serde_json::Value::as_f64)
.unwrap_or(1.0);
scores.push(score);
}
scores.iter().sum::<f64>() / scores.len() as f64
};
endpoint_health_scores.push(health_score);
json!({
"api_format": endpoint.api_format,
"health_score": health_score,
"is_active": endpoint.is_active,
"total_keys": endpoint_keys.len(),
"active_keys": endpoint_keys.iter().filter(|key| key.is_active).count(),
})
})
.collect::<Vec<_>>();
let avg_health_score = if endpoint_health_scores.is_empty() {
1.0
} else {
endpoint_health_scores.iter().sum::<f64>() / endpoint_health_scores.len() as f64
};
let unhealthy_endpoints = endpoint_health_scores
.iter()
.filter(|score| **score < 0.5)
.count();
let provider_config = provider.config.clone();
let config = provider_config
.as_ref()
.and_then(serde_json::Value::as_object);
let provider_ops_config = config.and_then(|cfg| cfg.get("provider_ops"));
let ops_configured = provider_ops_config.is_some_and(json_truthy);
let ops_architecture_id = provider_ops_config
.and_then(serde_json::Value::as_object)
.and_then(|cfg| cfg.get("architecture_id"))
.and_then(serde_json::Value::as_str)
.map(ToOwned::to_owned);
json!({
"id": provider.id.clone(),
"name": provider.name.clone(),
"provider_type": provider.provider_type.clone(),
"description": provider.description.clone(),
"website": provider.website.clone(),
"provider_priority": provider.provider_priority,
"keep_priority_on_conversion": provider.keep_priority_on_conversion,
"enable_format_conversion": provider.enable_format_conversion,
"is_active": provider.is_active,
"billing_type": quota_snapshot.map(|quota| quota.billing_type.clone()),
"monthly_quota_usd": quota_snapshot.and_then(|quota| quota.monthly_quota_usd),
"monthly_used_usd": quota_snapshot.map(|quota| quota.monthly_used_usd),
"quota_reset_day": quota_snapshot.and_then(|quota| quota.quota_reset_day),
"quota_last_reset_at": quota_snapshot
.and_then(|quota| quota.quota_last_reset_at_unix_secs)
.and_then(unix_secs_to_rfc3339),
"quota_expires_at": quota_snapshot
.and_then(|quota| quota.quota_expires_at_unix_secs)
.and_then(unix_secs_to_rfc3339),
"max_retries": provider.max_retries,
"proxy": provider.proxy.clone(),
"stream_first_byte_timeout": provider.stream_first_byte_timeout_secs,
"request_timeout": provider.request_timeout_secs,
"claude_code_advanced": config.and_then(|cfg| cfg.get("claude_code_advanced")).cloned(),
"pool_advanced": config.and_then(|cfg| cfg.get("pool_advanced")).cloned(),
"failover_rules": config.and_then(|cfg| cfg.get("failover_rules")).cloned(),
"total_endpoints": total_endpoints,
"active_endpoints": active_endpoints,
"total_keys": total_keys,
"active_keys": active_keys,
"total_models": total_models,
"active_models": active_models,
"global_model_ids": active_global_model_ids,
"avg_health_score": avg_health_score,
"unhealthy_endpoints": unhealthy_endpoints,
"api_formats": api_formats,
"endpoint_health_details": endpoint_health_details,
"ops_configured": ops_configured,
"ops_architecture_id": ops_architecture_id,
"created_at": endpoint_timestamp_or_now(provider.created_at_unix_secs, now_unix_secs),
"updated_at": endpoint_timestamp_or_now(provider.updated_at_unix_secs, now_unix_secs),
})
}
pub(crate) async fn build_admin_provider_summary_payload(
state: &AppState,
provider_id: &str,
) -> Option<serde_json::Value> {
if !state.has_provider_catalog_data_reader() {
return None;
}
let provider_ids = vec![provider_id.to_string()];
let provider = state
.read_provider_catalog_providers_by_ids(&provider_ids)
.await
.ok()?
.into_iter()
.next()?;
let (
endpoints_result,
keys_result,
quota_snapshot_result,
model_stats_result,
active_global_model_ids_result,
) = tokio::join!(
state.list_provider_catalog_endpoints_by_provider_ids(&provider_ids),
state.list_provider_catalog_keys_by_provider_ids(&provider_ids),
state.read_provider_quota_snapshot(provider_id),
state.list_provider_model_stats(&provider_ids),
state.list_active_global_model_ids_by_provider_ids(&provider_ids),
);
let endpoints = endpoints_result.ok().unwrap_or_default();
let keys = keys_result.ok().unwrap_or_default();
let quota_snapshot = quota_snapshot_result.ok().flatten();
let model_stats = model_stats_result
.ok()
.unwrap_or_default()
.into_iter()
.find(|stats| stats.provider_id == provider_id);
let active_global_model_ids = active_global_model_ids_result
.ok()
.unwrap_or_default()
.into_iter()
.filter(|row| row.provider_id == provider_id)
.map(|row| row.global_model_id)
.collect::<BTreeSet<_>>()
.into_iter()
.collect::<Vec<_>>();
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
Some(build_admin_provider_summary_value(
&provider,
&endpoints,
&keys,
quota_snapshot.as_ref(),
model_stats.as_ref(),
active_global_model_ids,
now_unix_secs,
))
}
pub(crate) async fn build_admin_providers_summary_payload(
state: &AppState,
page: usize,
page_size: usize,
search: &str,
status: &str,
api_format: &str,
model_id: &str,
) -> Option<serde_json::Value> {
if !state.has_provider_catalog_data_reader() {
return None;
}
let normalized_search = search.trim().to_ascii_lowercase();
let search_keywords = normalized_search
.split_whitespace()
.filter(|value| !value.is_empty())
.collect::<Vec<_>>();
let normalized_status = status.trim().to_ascii_lowercase();
let normalized_api_format = api_format.trim();
let normalized_model_id = model_id.trim();
let mut providers = state
.list_provider_catalog_providers(false)
.await
.ok()
.unwrap_or_default();
let all_provider_ids = providers
.iter()
.map(|provider| provider.id.clone())
.collect::<Vec<_>>();
let all_endpoints = if all_provider_ids.is_empty() {
Vec::new()
} else {
state
.list_provider_catalog_endpoints_by_provider_ids(&all_provider_ids)
.await
.ok()
.unwrap_or_default()
};
let active_global_model_refs = if all_provider_ids.is_empty() {
Vec::new()
} else {
state
.list_active_global_model_ids_by_provider_ids(&all_provider_ids)
.await
.ok()
.unwrap_or_default()
};
let mut api_formats_by_provider = BTreeMap::<String, BTreeSet<String>>::new();
for endpoint in &all_endpoints {
api_formats_by_provider
.entry(endpoint.provider_id.clone())
.or_default()
.insert(endpoint.api_format.clone());
}
let mut active_global_model_ids_by_provider = BTreeMap::<String, BTreeSet<String>>::new();
for row in active_global_model_refs {
active_global_model_ids_by_provider
.entry(row.provider_id)
.or_default()
.insert(row.global_model_id);
}
providers.retain(|provider| {
if !search_keywords.is_empty() {
let provider_name = provider.name.to_ascii_lowercase();
if !search_keywords
.iter()
.all(|keyword| provider_name.contains(keyword))
{
return false;
}
}
match normalized_status.as_str() {
"active" if !provider.is_active => return false,
"inactive" if provider.is_active => return false,
_ => {}
}
if normalized_api_format != "all"
&& !normalized_api_format.is_empty()
&& !api_formats_by_provider
.get(&provider.id)
.is_some_and(|items| items.contains(normalized_api_format))
{
return false;
}
if normalized_model_id != "all"
&& !normalized_model_id.is_empty()
&& !active_global_model_ids_by_provider
.get(&provider.id)
.is_some_and(|items| items.contains(normalized_model_id))
{
return false;
}
true
});
providers.sort_by(|left, right| {
right
.is_active
.cmp(&left.is_active)
.then_with(|| left.provider_priority.cmp(&right.provider_priority))
.then_with(|| left.created_at_unix_secs.cmp(&right.created_at_unix_secs))
});
let total = providers.len();
let offset = page.saturating_sub(1).saturating_mul(page_size);
let providers = providers
.into_iter()
.skip(offset)
.take(page_size)
.collect::<Vec<_>>();
let provider_ids = providers
.iter()
.map(|provider| provider.id.clone())
.collect::<Vec<_>>();
let endpoints = if provider_ids.is_empty() {
Vec::new()
} else {
state
.list_provider_catalog_endpoints_by_provider_ids(&provider_ids)
.await
.ok()
.unwrap_or_default()
};
let keys = if provider_ids.is_empty() {
Vec::new()
} else {
state
.list_provider_catalog_keys_by_provider_ids(&provider_ids)
.await
.ok()
.unwrap_or_default()
};
let model_stats = if provider_ids.is_empty() {
Vec::new()
} else {
state
.list_provider_model_stats(&provider_ids)
.await
.ok()
.unwrap_or_default()
};
let mut endpoints_by_provider = BTreeMap::<String, Vec<StoredProviderCatalogEndpoint>>::new();
for endpoint in endpoints {
endpoints_by_provider
.entry(endpoint.provider_id.clone())
.or_default()
.push(endpoint);
}
let mut keys_by_provider = BTreeMap::<String, Vec<StoredProviderCatalogKey>>::new();
for key in keys {
keys_by_provider
.entry(key.provider_id.clone())
.or_default()
.push(key);
}
let model_stats_by_provider = model_stats
.into_iter()
.map(|stats| (stats.provider_id.clone(), stats))
.collect::<BTreeMap<_, _>>();
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
let mut items = Vec::with_capacity(providers.len());
for provider in providers {
let quota_snapshot = state
.read_provider_quota_snapshot(&provider.id)
.await
.ok()
.flatten();
let active_global_model_ids = active_global_model_ids_by_provider
.get(&provider.id)
.cloned()
.unwrap_or_default()
.into_iter()
.collect::<Vec<_>>();
items.push(build_admin_provider_summary_value(
&provider,
endpoints_by_provider
.get(&provider.id)
.map(Vec::as_slice)
.unwrap_or(&[]),
keys_by_provider
.get(&provider.id)
.map(Vec::as_slice)
.unwrap_or(&[]),
quota_snapshot.as_ref(),
model_stats_by_provider.get(&provider.id),
active_global_model_ids,
now_unix_secs,
));
}
Some(json!({
"total": total,
"page": page,
"page_size": page_size,
"items": items,
}))
}
pub(crate) async fn build_admin_provider_health_monitor_payload(
state: &AppState,
provider_id: &str,
lookback_hours: u64,
per_endpoint_limit: usize,
) -> Option<serde_json::Value> {
if !state.has_provider_catalog_data_reader() || !state.has_request_candidate_data_reader() {
return None;
}
let provider = state
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
.await
.ok()
.and_then(|mut providers| providers.drain(..).next())?;
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let since_unix_secs = now_unix_secs.saturating_sub(lookback_hours * 3600);
let mut endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await
.ok()
.unwrap_or_default();
endpoints.sort_by(|left, right| {
left.api_format
.cmp(&right.api_format)
.then_with(|| left.id.cmp(&right.id))
});
if endpoints.is_empty() {
return Some(json!({
"provider_id": provider.id,
"provider_name": provider.name,
"generated_at": unix_secs_to_rfc3339(now_unix_secs),
"endpoints": [],
}));
}
let endpoint_ids = endpoints
.iter()
.map(|endpoint| endpoint.id.clone())
.collect::<Vec<_>>();
let fetch_limit = per_endpoint_limit
.saturating_mul(endpoint_ids.len())
.max(per_endpoint_limit);
let attempts = state
.list_finalized_request_candidates_by_endpoint_ids_since(
&endpoint_ids,
since_unix_secs,
fetch_limit,
)
.await
.ok()
.unwrap_or_default();
let mut attempts_by_endpoint = BTreeMap::<String, Vec<StoredRequestCandidate>>::new();
for candidate in attempts {
let Some(endpoint_id) = candidate.endpoint_id.clone() else {
continue;
};
attempts_by_endpoint
.entry(endpoint_id)
.or_default()
.push(candidate);
}
for candidates in attempts_by_endpoint.values_mut() {
candidates.sort_by(|left, right| {
right
.created_at_unix_secs
.cmp(&left.created_at_unix_secs)
.then_with(|| right.id.cmp(&left.id))
});
candidates.truncate(per_endpoint_limit);
candidates.sort_by(|left, right| {
request_candidate_event_unix_secs(left)
.cmp(&request_candidate_event_unix_secs(right))
.then_with(|| left.id.cmp(&right.id))
});
}
let endpoints = endpoints
.into_iter()
.map(|endpoint| {
let candidates = attempts_by_endpoint.remove(&endpoint.id).unwrap_or_default();
let success_count = candidates
.iter()
.filter(|candidate| candidate.status == RequestCandidateStatus::Success)
.count();
let failed_count = candidates
.iter()
.filter(|candidate| candidate.status == RequestCandidateStatus::Failed)
.count();
let skipped_count = candidates
.iter()
.filter(|candidate| candidate.status == RequestCandidateStatus::Skipped)
.count();
let total_attempts = candidates.len();
let success_rate = if total_attempts > 0 {
success_count as f64 / total_attempts as f64
} else {
1.0
};
let last_event_at = candidates
.last()
.and_then(|candidate| unix_secs_to_rfc3339(request_candidate_event_unix_secs(candidate)));
let events = candidates
.into_iter()
.filter_map(|candidate| {
Some(json!({
"timestamp": unix_secs_to_rfc3339(request_candidate_event_unix_secs(&candidate))?,
"status": request_candidate_status_label(candidate.status),
"status_code": candidate.status_code,
"latency_ms": candidate.latency_ms,
"error_type": candidate.error_type,
"error_message": candidate.error_message,
}))
})
.collect::<Vec<_>>();
json!({
"endpoint_id": endpoint.id,
"api_format": endpoint.api_format,
"is_active": endpoint.is_active,
"total_attempts": total_attempts,
"success_count": success_count,
"failed_count": failed_count,
"skipped_count": skipped_count,
"success_rate": success_rate,
"last_event_at": last_event_at,
"events": events,
})
})
.collect::<Vec<_>>();
Some(json!({
"provider_id": provider.id,
"provider_name": provider.name,
"generated_at": unix_secs_to_rfc3339(now_unix_secs),
"endpoints": endpoints,
}))
}
@@ -0,0 +1,495 @@
use super::{
normalize_auth_type, normalize_json_object, normalize_string_list, validate_vertex_api_formats,
};
use crate::handlers::admin::provider::shared::{
AdminProviderKeyCreateRequest, AdminProviderKeyUpdateRequest,
};
use crate::handlers::admin::shared::{
build_admin_provider_key_response, decrypt_catalog_secret_with_fallbacks,
encrypt_catalog_secret_with_fallbacks, json_string_list, parse_catalog_auth_config_json,
};
use crate::AppState;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
use uuid::Uuid;
pub(crate) async fn build_admin_create_provider_key_record(
state: &AppState,
provider: &StoredProviderCatalogProvider,
payload: AdminProviderKeyCreateRequest,
) -> Result<StoredProviderCatalogKey, String> {
let name = payload.name.trim();
if name.is_empty() {
return Err("name 为必填字段".to_string());
}
let api_formats = normalize_string_list(payload.api_formats)
.ok_or_else(|| "api_formats 为必填字段".to_string())?;
let auth_type = normalize_auth_type(payload.auth_type.as_deref())?;
validate_vertex_api_formats(&provider.provider_type, &auth_type, &api_formats)?;
let api_key = payload.api_key.unwrap_or_default().trim().to_string();
let auth_config = normalize_json_object(payload.auth_config, "auth_config")?;
let auth_config_object = auth_config
.as_ref()
.and_then(serde_json::Value::as_object)
.cloned();
match auth_type.as_str() {
"api_key" => {
if api_key.is_empty() {
return Err("API Key 认证模式下 api_key 为必填字段".to_string());
}
}
"service_account" => {
if auth_config_object.is_none() {
return Err("Service Account 认证模式下 auth_config 为必填字段".to_string());
}
}
"oauth" => {
if !api_key.is_empty() {
return Err("OAuth 认证模式下不允许直接填写 api_key".to_string());
}
}
_ => {}
}
let existing_keys = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
.await
.map_err(|err| format!("{err:?}"))?;
if auth_type == "api_key" {
for existing in existing_keys
.iter()
.filter(|existing| existing.auth_type.trim().eq_ignore_ascii_case("api_key"))
{
let Some(decrypted) = decrypt_catalog_secret_with_fallbacks(
state.encryption_key(),
&existing.encrypted_api_key,
) else {
continue;
};
if decrypted != "__placeholder__" && decrypted == api_key {
return Err(format!(
"该 API Key 已存在于当前 Provider 中(名称: {})",
existing.name
));
}
}
}
if auth_type == "service_account" {
let new_client_email = auth_config_object
.as_ref()
.and_then(|config| config.get("client_email"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
if let Some(new_client_email) = new_client_email {
for existing in existing_keys.iter().filter(|existing| {
matches!(
existing.auth_type.trim().to_ascii_lowercase().as_str(),
"service_account" | "vertex_ai"
)
}) {
let Some(existing_config) = parse_catalog_auth_config_json(state, existing) else {
continue;
};
let Some(existing_email) = existing_config
.get("client_email")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
continue;
};
if existing_email == new_client_email {
return Err(format!(
"该 Service Account ({new_client_email}) 已存在于当前 Provider 中(名称: {})",
existing.name
));
}
}
}
}
let encrypted_api_key = match auth_type.as_str() {
"api_key" => encrypt_catalog_secret_with_fallbacks(state, &api_key),
_ => encrypt_catalog_secret_with_fallbacks(state, "__placeholder__"),
}
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?;
let encrypted_auth_config = auth_config
.as_ref()
.map(serde_json::to_string)
.transpose()
.map_err(|err| err.to_string())?
.and_then(|plaintext| encrypt_catalog_secret_with_fallbacks(state, &plaintext));
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let mut key = StoredProviderCatalogKey::new(
Uuid::new_v4().to_string(),
provider.id.clone(),
name.to_string(),
auth_type,
normalize_json_object(payload.capabilities, "capabilities")?,
true,
)
.map_err(|err| err.to_string())?
.with_transport_fields(
Some(json!(api_formats)),
encrypted_api_key,
encrypted_auth_config,
normalize_json_object(payload.rate_multipliers, "rate_multipliers")?,
None,
normalize_string_list(payload.allowed_models).map(|value| json!(value)),
None,
None,
None,
)
.map_err(|err| err.to_string())?;
key.note = payload
.note
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
key.internal_priority = payload.internal_priority.unwrap_or(50);
key.rpm_limit = payload.rpm_limit;
key.cache_ttl_minutes = payload.cache_ttl_minutes.unwrap_or(5);
key.max_probe_interval_minutes = payload.max_probe_interval_minutes.unwrap_or(32);
key.request_count = Some(0);
key.success_count = Some(0);
key.error_count = Some(0);
key.total_response_time_ms = Some(0);
key.auto_fetch_models = payload.auto_fetch_models.unwrap_or(false);
key.locked_models = normalize_string_list(payload.locked_models).map(|value| json!(value));
key.model_include_patterns =
normalize_string_list(payload.model_include_patterns).map(|value| json!(value));
key.model_exclude_patterns =
normalize_string_list(payload.model_exclude_patterns).map(|value| json!(value));
key.health_by_format = Some(json!({}));
key.circuit_breaker_by_format = Some(json!({}));
key.created_at_unix_secs = Some(now_unix_secs);
key.updated_at_unix_secs = Some(now_unix_secs);
Ok(key)
}
pub(crate) async fn build_admin_update_provider_key_record(
state: &AppState,
provider: &StoredProviderCatalogProvider,
existing: &StoredProviderCatalogKey,
raw_payload: &serde_json::Map<String, serde_json::Value>,
payload: AdminProviderKeyUpdateRequest,
) -> Result<StoredProviderCatalogKey, String> {
let mut updated = existing.clone();
let current_auth_type = normalize_auth_type(Some(&existing.auth_type))?;
let target_auth_type = payload
.auth_type
.as_deref()
.map(|value| normalize_auth_type(Some(value)))
.transpose()?
.unwrap_or_else(|| current_auth_type.clone());
let auth_type_switch = payload
.auth_type
.as_deref()
.is_some_and(|_| target_auth_type != current_auth_type);
let api_key_present = raw_payload.contains_key("api_key");
let api_key_value = payload
.api_key
.as_deref()
.map(str::trim)
.map(ToOwned::to_owned);
if api_key_present && api_key_value.as_deref() == Some("") {
return Err("api_key 不能为空".to_string());
}
let auth_config_present = raw_payload.contains_key("auth_config");
let auth_config = normalize_json_object(payload.auth_config, "auth_config")?;
let auth_config_object = auth_config
.as_ref()
.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" => {
if auth_type_switch
&& matches!(
api_key_value.as_deref(),
None | Some("") | Some("__placeholder__")
)
{
return Err("切换到 API Key 认证模式时,必须提供新的 API Key".to_string());
}
if api_key_present
&& matches!(
api_key_value.as_deref(),
None | Some("") | Some("__placeholder__")
)
{
return Err("API Key 认证模式下 api_key 不能为空".to_string());
}
if let Some(api_key) = api_key_value.as_deref() {
for existing_key in existing_keys.iter().filter(|key| {
key.id != existing.id && key.auth_type.trim().eq_ignore_ascii_case("api_key")
}) {
let Some(decrypted) = decrypt_catalog_secret_with_fallbacks(
state.encryption_key(),
&existing_key.encrypted_api_key,
) else {
continue;
};
if decrypted != "__placeholder__" && decrypted == api_key {
return Err(format!(
"该 API Key 已存在于当前 Provider 中(名称: {})",
existing_key.name
));
}
}
updated.encrypted_api_key =
encrypt_catalog_secret_with_fallbacks(state, api_key)
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?;
}
updated.encrypted_auth_config = None;
}
"service_account" => {
if auth_type_switch && auth_config_object.is_none() {
return Err(
"切换到 Service Account 认证模式时,必须提供 Service Account JSON".to_string(),
);
}
if api_key_present
&& !matches!(
api_key_value.as_deref(),
None | Some("") | Some("__placeholder__")
)
{
return Err("Service Account 认证模式下不允许直接填写 api_key".to_string());
}
if auth_type_switch || api_key_present {
updated.encrypted_api_key =
encrypt_catalog_secret_with_fallbacks(state, "__placeholder__")
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?;
}
if let Some(client_email) = auth_config_object
.as_ref()
.and_then(|config| config.get("client_email"))
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
{
for existing_key in existing_keys.iter().filter(|key| {
key.id != existing.id
&& matches!(
key.auth_type.trim().to_ascii_lowercase().as_str(),
"service_account" | "vertex_ai"
)
}) {
let Some(existing_config) = parse_catalog_auth_config_json(state, existing_key)
else {
continue;
};
let Some(existing_email) = existing_config
.get("client_email")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
else {
continue;
};
if existing_email == client_email {
return Err(format!(
"该 Service Account ({client_email}) 已存在于当前 Provider 中(名称: {})",
existing_key.name
));
}
}
}
if auth_config_present {
updated.encrypted_auth_config = auth_config
.as_ref()
.map(serde_json::to_string)
.transpose()
.map_err(|err| err.to_string())?
.map(|plaintext| {
encrypt_catalog_secret_with_fallbacks(state, &plaintext)
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())
})
.transpose()?;
}
}
"oauth" => {
if api_key_present
&& !matches!(
api_key_value.as_deref(),
None | Some("") | Some("__placeholder__")
)
{
return Err("OAuth 认证模式下不允许直接填写 api_key".to_string());
}
if auth_type_switch {
updated.encrypted_api_key =
encrypt_catalog_secret_with_fallbacks(state, "__placeholder__")
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?;
updated.encrypted_auth_config = None;
}
}
_ => {}
}
if raw_payload.contains_key("api_formats") {
let api_formats = normalize_string_list(payload.api_formats)
.ok_or_else(|| "api_formats 为必填字段".to_string())?;
validate_vertex_api_formats(&provider.provider_type, &target_auth_type, &api_formats)?;
updated.api_formats = Some(json!(api_formats));
} else if payload.auth_type.is_some() {
let api_formats = json_string_list(existing.api_formats.as_ref());
validate_vertex_api_formats(&provider.provider_type, &target_auth_type, &api_formats)?;
}
updated.auth_type = target_auth_type;
if let Some(name) = payload.name {
let trimmed = name.trim();
if trimmed.is_empty() {
return Err("name 为必填字段".to_string());
}
updated.name = trimmed.to_string();
}
if raw_payload.contains_key("rate_multipliers") {
updated.rate_multipliers =
normalize_json_object(payload.rate_multipliers, "rate_multipliers")?;
}
if let Some(internal_priority) = payload.internal_priority {
updated.internal_priority = internal_priority;
}
if raw_payload.contains_key("global_priority_by_format") {
updated.global_priority_by_format = normalize_json_object(
payload.global_priority_by_format,
"global_priority_by_format",
)?;
}
if raw_payload.contains_key("rpm_limit") {
updated.rpm_limit = payload.rpm_limit;
if payload.rpm_limit.is_none() {
updated.learned_rpm_limit = None;
}
}
if raw_payload.contains_key("allowed_models") {
updated.allowed_models =
normalize_string_list(payload.allowed_models).map(|value| json!(value));
}
if raw_payload.contains_key("capabilities") {
updated.capabilities = normalize_json_object(payload.capabilities, "capabilities")?;
}
if let Some(cache_ttl_minutes) = payload.cache_ttl_minutes {
updated.cache_ttl_minutes = cache_ttl_minutes;
}
if let Some(max_probe_interval_minutes) = payload.max_probe_interval_minutes {
updated.max_probe_interval_minutes = max_probe_interval_minutes;
}
if let Some(is_active) = payload.is_active {
updated.is_active = is_active;
}
if raw_payload.contains_key("note") {
updated.note = payload
.note
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
}
if let Some(auto_fetch_models) = payload.auto_fetch_models {
updated.auto_fetch_models = auto_fetch_models;
}
if raw_payload.contains_key("locked_models") {
updated.locked_models =
normalize_string_list(payload.locked_models).map(|value| json!(value));
}
if raw_payload.contains_key("model_include_patterns") {
updated.model_include_patterns =
normalize_string_list(payload.model_include_patterns).map(|value| json!(value));
}
if raw_payload.contains_key("model_exclude_patterns") {
updated.model_exclude_patterns =
normalize_string_list(payload.model_exclude_patterns).map(|value| json!(value));
}
if raw_payload.contains_key("proxy") {
updated.proxy = normalize_json_object(payload.proxy, "proxy")?;
}
if raw_payload.contains_key("fingerprint") {
updated.fingerprint = normalize_json_object(payload.fingerprint, "fingerprint")?;
}
if auth_config_present && !auth_type_switch && updated.auth_type != "api_key" {
updated.encrypted_auth_config = auth_config
.as_ref()
.map(serde_json::to_string)
.transpose()
.map_err(|err| err.to_string())?
.map(|plaintext| {
encrypt_catalog_secret_with_fallbacks(state, &plaintext)
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())
})
.transpose()?;
}
updated.updated_at_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs());
Ok(updated)
}
pub(crate) async fn build_admin_provider_keys_payload(
state: &AppState,
provider_id: &str,
skip: usize,
limit: usize,
) -> Option<serde_json::Value> {
if !state.has_provider_catalog_data_reader() {
return None;
}
let provider = state
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
.await
.ok()
.and_then(|mut providers| providers.drain(..).next())?;
let mut keys = state
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
.await
.ok()
.unwrap_or_default();
keys.sort_by(|left, right| {
left.internal_priority
.cmp(&right.internal_priority)
.then_with(|| {
left.created_at_unix_secs
.unwrap_or_default()
.cmp(&right.created_at_unix_secs.unwrap_or_default())
})
.then_with(|| left.id.cmp(&right.id))
});
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
Some(serde_json::Value::Array(
keys.into_iter()
.skip(skip)
.take(limit)
.map(|key| build_admin_provider_key_response(state, &key, now_unix_secs))
.collect(),
))
}
@@ -0,0 +1,85 @@
use crate::handlers::admin::shared::{normalize_json_object, normalize_string_list};
pub(crate) fn normalize_provider_type_input(value: &str) -> Result<String, String> {
let normalized = value.trim().to_ascii_lowercase();
match normalized.as_str() {
"custom" | "claude_code" | "kiro" | "codex" | "gemini_cli" | "antigravity"
| "vertex_ai" => Ok(normalized),
_ => Err(
"provider_type 仅支持 custom / claude_code / kiro / codex / gemini_cli / antigravity / vertex_ai"
.to_string(),
),
}
}
pub(crate) fn normalize_provider_billing_type(value: &str) -> Result<String, String> {
let normalized = value.trim().to_ascii_lowercase();
match normalized.as_str() {
"monthly_quota" | "pay_as_you_go" | "free_tier" => Ok(normalized),
_ => Err("billing_type 仅支持 monthly_quota / pay_as_you_go / free_tier".to_string()),
}
}
pub(crate) fn parse_optional_rfc3339_unix_secs(
value: &str,
field_name: &str,
) -> Result<u64, String> {
let trimmed = value.trim();
if trimmed.is_empty() {
return Err(format!("{field_name} 不能为空"));
}
let parsed = chrono::DateTime::parse_from_rfc3339(trimmed)
.map_err(|_| format!("{field_name} 必须是合法的 RFC3339 时间"))?;
u64::try_from(parsed.timestamp()).map_err(|_| format!("{field_name} 超出有效时间范围"))
}
pub(crate) fn normalize_auth_type(value: Option<&str>) -> Result<String, String> {
let auth_type = value.unwrap_or("api_key").trim().to_ascii_lowercase();
match auth_type.as_str() {
"api_key" | "service_account" | "oauth" => Ok(auth_type),
_ => Err("auth_type 仅支持 api_key / service_account / oauth".to_string()),
}
}
pub(crate) fn validate_vertex_api_formats(
provider_type: &str,
auth_type: &str,
api_formats: &[String],
) -> Result<(), String> {
if !provider_type.trim().eq_ignore_ascii_case("vertex_ai") {
return Ok(());
}
let allowed = match auth_type {
"api_key" => &["gemini:chat"][..],
"service_account" | "vertex_ai" => &["claude:chat", "gemini:chat"][..],
_ => return Ok(()),
};
let invalid = api_formats
.iter()
.filter(|value| !allowed.contains(&value.as_str()))
.cloned()
.collect::<Vec<_>>();
if invalid.is_empty() {
return Ok(());
}
Err(format!(
"Vertex {auth_type} 不支持以下 API 格式: {};允许: {}",
invalid.join(", "),
allowed.join(", ")
))
}
mod keys;
mod provider;
mod reveal;
pub(crate) use self::keys::{
build_admin_create_provider_key_record, build_admin_provider_keys_payload,
build_admin_update_provider_key_record,
};
pub(crate) use self::provider::{
build_admin_create_provider_record, build_admin_fixed_provider_endpoint_record,
build_admin_update_provider_record,
};
pub(crate) use self::reveal::{build_admin_export_key_payload, build_admin_reveal_key_payload};
@@ -0,0 +1,510 @@
use super::{
normalize_json_object, normalize_provider_billing_type, normalize_provider_type_input,
parse_optional_rfc3339_unix_secs,
};
use crate::api::ai::{admin_default_body_rules_for_signature, admin_endpoint_signature_parts};
use crate::handlers::admin::provider::shared::{
AdminProviderCreateRequest, AdminProviderUpdateRequest,
};
use crate::handlers::public::normalize_admin_base_url;
use crate::provider_transport::provider_types::provider_type_enables_format_conversion_by_default;
use crate::AppState;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
};
use serde_json::json;
use std::time::{SystemTime, UNIX_EPOCH};
use uuid::Uuid;
pub(crate) async fn build_admin_update_provider_record(
state: &AppState,
existing: &StoredProviderCatalogProvider,
raw_payload: &serde_json::Map<String, serde_json::Value>,
payload: AdminProviderUpdateRequest,
) -> Result<StoredProviderCatalogProvider, String> {
let mut updated = existing.clone();
if let Some(value) = raw_payload.get("name") {
let Some(name) = payload.name.as_deref() else {
return Err(if value.is_null() {
"name 不能为空".to_string()
} else {
"name 必须是字符串".to_string()
});
};
let trimmed = name.trim();
if trimmed.is_empty() {
return Err("name 不能为空".to_string());
}
let duplicate = state
.list_provider_catalog_providers(false)
.await
.map_err(|err| format!("{err:?}"))?
.into_iter()
.any(|provider| provider.id != existing.id && provider.name == trimmed);
if duplicate {
return Err(format!("提供商名称 '{trimmed}' 已存在"));
}
updated.name = trimmed.to_string();
}
let target_provider_type = if let Some(value) = raw_payload.get("provider_type") {
let Some(provider_type) = payload.provider_type.as_deref() else {
return Err(if value.is_null() {
"provider_type 不能为空".to_string()
} else {
"provider_type 必须是字符串".to_string()
});
};
let normalized = normalize_provider_type_input(provider_type)?;
updated.provider_type = normalized.clone();
normalized
} else {
updated.provider_type.clone()
};
if raw_payload.contains_key("description") {
updated.description = payload
.description
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
}
if let Some(value) = raw_payload.get("website") {
updated.website = match payload.website {
None => {
if value.is_null() {
None
} else {
return Err("website 必须是字符串".to_string());
}
}
Some(website) => {
let trimmed = website.trim();
if trimmed.is_empty() {
None
} else if !trimmed.starts_with("http://") && !trimmed.starts_with("https://") {
return Err("website 必须以 http:// 或 https:// 开头".to_string());
} else {
Some(trimmed.to_string())
}
}
};
}
if let Some(value) = raw_payload.get("billing_type") {
let Some(billing_type) = payload.billing_type.as_deref() else {
return Err(if value.is_null() {
"billing_type 不能为空".to_string()
} else {
"billing_type 必须是字符串".to_string()
});
};
updated.billing_type = Some(normalize_provider_billing_type(billing_type)?);
}
if let Some(value) = raw_payload.get("monthly_quota_usd") {
if value.is_null() {
updated.monthly_quota_usd = None;
} else {
let Some(monthly_quota_usd) = payload.monthly_quota_usd else {
return Err("monthly_quota_usd 必须是非负数".to_string());
};
if !monthly_quota_usd.is_finite() || monthly_quota_usd < 0.0 {
return Err("monthly_quota_usd 必须是非负数".to_string());
}
updated.monthly_quota_usd = Some(monthly_quota_usd);
}
}
if let Some(value) = raw_payload.get("quota_reset_day") {
if value.is_null() {
updated.quota_reset_day = None;
} else {
let Some(quota_reset_day) = payload.quota_reset_day else {
return Err("quota_reset_day 必须是 1 到 365 之间的整数".to_string());
};
if !(1..=365).contains(&quota_reset_day) {
return Err("quota_reset_day 必须是 1 到 365 之间的整数".to_string());
}
updated.quota_reset_day = Some(quota_reset_day);
}
}
if let Some(value) = raw_payload.get("quota_last_reset_at") {
if value.is_null() {
updated.quota_last_reset_at_unix_secs = None;
} else {
let Some(raw) = payload.quota_last_reset_at.as_deref() else {
return Err("quota_last_reset_at 必须是字符串".to_string());
};
updated.quota_last_reset_at_unix_secs = Some(parse_optional_rfc3339_unix_secs(
raw,
"quota_last_reset_at",
)?);
}
}
if let Some(value) = raw_payload.get("quota_expires_at") {
if value.is_null() {
updated.quota_expires_at_unix_secs = None;
} else {
let Some(raw) = payload.quota_expires_at.as_deref() else {
return Err("quota_expires_at 必须是字符串".to_string());
};
updated.quota_expires_at_unix_secs =
Some(parse_optional_rfc3339_unix_secs(raw, "quota_expires_at")?);
}
}
if let Some(value) = raw_payload.get("provider_priority") {
let Some(provider_priority) = payload.provider_priority else {
return Err(if value.is_null() {
"provider_priority 不能为空".to_string()
} else {
"provider_priority 必须是整数".to_string()
});
};
if !(0..=10_000).contains(&provider_priority) {
return Err("provider_priority 必须在 0 到 10000 之间".to_string());
}
updated.provider_priority = provider_priority;
}
if let Some(_value) = raw_payload.get("keep_priority_on_conversion") {
let Some(keep_priority_on_conversion) = payload.keep_priority_on_conversion else {
return Err("keep_priority_on_conversion 必须是布尔值".to_string());
};
updated.keep_priority_on_conversion = keep_priority_on_conversion;
}
if let Some(_value) = raw_payload.get("is_active") {
let Some(is_active) = payload.is_active else {
return Err("is_active 必须是布尔值".to_string());
};
updated.is_active = is_active;
}
if raw_payload.contains_key("concurrent_limit") {
updated.concurrent_limit = match payload.concurrent_limit {
Some(value) if value >= 0 => Some(value),
Some(_) => return Err("concurrent_limit 必须是非负整数".to_string()),
None => None,
};
}
if raw_payload.contains_key("max_retries") {
updated.max_retries = match payload.max_retries {
Some(value) if (0..=999).contains(&value) => Some(value),
Some(_) => return Err("max_retries 必须是 0 到 999 之间的整数".to_string()),
None => None,
};
}
if raw_payload.contains_key("proxy") {
updated.proxy = normalize_json_object(payload.proxy, "proxy")?;
}
if raw_payload.contains_key("stream_first_byte_timeout") {
updated.stream_first_byte_timeout_secs = match payload.stream_first_byte_timeout {
Some(value) if (1.0..=300.0).contains(&value) => Some(value),
Some(_) => {
return Err("stream_first_byte_timeout 必须是 1 到 300 之间的数字".to_string())
}
None => None,
};
}
if raw_payload.contains_key("request_timeout") {
updated.request_timeout_secs = match payload.request_timeout {
Some(value) if (1.0..=600.0).contains(&value) => Some(value),
Some(_) => return Err("request_timeout 必须是 1 到 600 之间的数字".to_string()),
None => None,
};
}
if let Some(_value) = raw_payload.get("enable_format_conversion") {
let Some(enable_format_conversion) = payload.enable_format_conversion else {
return Err("enable_format_conversion 必须是布尔值".to_string());
};
updated.enable_format_conversion = enable_format_conversion;
}
let config_seed = if raw_payload.contains_key("config") {
normalize_json_object(payload.config, "config")?
} else {
updated.config.clone()
};
let mut config_map = config_seed
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
if raw_payload.contains_key("claude_code_advanced") {
if raw_payload
.get("claude_code_advanced")
.is_some_and(serde_json::Value::is_null)
{
config_map.remove("claude_code_advanced");
} else {
if target_provider_type != "claude_code" {
return Err("claude_code_advanced 仅适用于 provider_type=claude_code".to_string());
}
let value =
normalize_json_object(payload.claude_code_advanced, "claude_code_advanced")?
.ok_or_else(|| "claude_code_advanced 必须是 JSON 对象".to_string())?;
config_map.insert("claude_code_advanced".to_string(), value);
}
} else if target_provider_type != "claude_code" {
config_map.remove("claude_code_advanced");
}
if raw_payload.contains_key("pool_advanced") {
if raw_payload
.get("pool_advanced")
.is_some_and(serde_json::Value::is_null)
{
config_map.remove("pool_advanced");
} else {
let value = normalize_json_object(payload.pool_advanced, "pool_advanced")?
.ok_or_else(|| "pool_advanced 必须是 JSON 对象".to_string())?;
config_map.insert("pool_advanced".to_string(), value);
}
}
if raw_payload.contains_key("failover_rules") {
if raw_payload
.get("failover_rules")
.is_some_and(serde_json::Value::is_null)
{
config_map.remove("failover_rules");
} else {
let value = normalize_json_object(payload.failover_rules, "failover_rules")?
.ok_or_else(|| "failover_rules 必须是 JSON 对象".to_string())?;
config_map.insert("failover_rules".to_string(), value);
}
}
updated.config = (!config_map.is_empty()).then_some(serde_json::Value::Object(config_map));
updated.updated_at_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs());
Ok(updated)
}
pub(crate) async fn build_admin_create_provider_record(
state: &AppState,
payload: AdminProviderCreateRequest,
) -> Result<(StoredProviderCatalogProvider, Option<i32>), String> {
let name = payload.name.trim();
if name.is_empty() {
return Err("name 为必填字段".to_string());
}
let existing_providers = state
.list_provider_catalog_providers(false)
.await
.map_err(|err| format!("{err:?}"))?;
if existing_providers
.iter()
.any(|provider| provider.name == name)
{
return Err(format!("提供商名称 '{name}' 已存在"));
}
let provider_type =
normalize_provider_type_input(payload.provider_type.as_deref().unwrap_or("custom"))?;
let billing_type = normalize_provider_billing_type(
payload.billing_type.as_deref().unwrap_or("pay_as_you_go"),
)?;
let website = payload.website.and_then(|value| {
let trimmed = value.trim().to_string();
(!trimmed.is_empty()).then_some(trimmed)
});
let website = website.map(|value| {
if value.starts_with("http://") || value.starts_with("https://") {
value
} else {
format!("https://{value}")
}
});
let monthly_quota_usd = match payload.monthly_quota_usd {
Some(value) if value.is_finite() && value >= 0.0 => Some(value),
Some(_) => return Err("monthly_quota_usd 必须是非负数".to_string()),
None => None,
};
let quota_reset_day = match payload.quota_reset_day {
Some(value) if (1..=365).contains(&value) => Some(value),
Some(_) => return Err("quota_reset_day 必须是 1 到 365 之间的整数".to_string()),
None => Some(30),
};
let quota_last_reset_at_unix_secs = payload
.quota_last_reset_at
.as_deref()
.map(|value| parse_optional_rfc3339_unix_secs(value, "quota_last_reset_at"))
.transpose()?;
let quota_expires_at_unix_secs = payload
.quota_expires_at
.as_deref()
.map(|value| parse_optional_rfc3339_unix_secs(value, "quota_expires_at"))
.transpose()?;
let provider_priority = match payload.provider_priority {
Some(value) if (0..=10_000).contains(&value) => value,
Some(_) => return Err("provider_priority 必须在 0 到 10000 之间".to_string()),
None => {
let current_min_priority = existing_providers
.iter()
.map(|provider| provider.provider_priority)
.min();
match current_min_priority {
Some(value) if value <= 0 => 0,
Some(value) => value - 1,
None => 100,
}
}
};
let shift_existing_priorities_from = match payload.provider_priority {
Some(_) => Some(provider_priority),
None => existing_providers
.iter()
.map(|provider| provider.provider_priority)
.min()
.filter(|value| *value <= 0)
.map(|_| 0),
};
let is_active = payload.is_active.unwrap_or(true);
let concurrent_limit = match payload.concurrent_limit {
Some(value) if value >= 0 => Some(value),
Some(_) => return Err("concurrent_limit 必须是非负整数".to_string()),
None => None,
};
let max_retries = match payload.max_retries {
Some(value) if (0..=999).contains(&value) => Some(value),
Some(_) => return Err("max_retries 必须是 0 到 999 之间的整数".to_string()),
None => Some(2),
};
let proxy = normalize_json_object(payload.proxy, "proxy")?;
let stream_first_byte_timeout_secs = match payload.stream_first_byte_timeout {
Some(value) if (1.0..=300.0).contains(&value) => Some(value),
Some(_) => return Err("stream_first_byte_timeout 必须是 1 到 300 之间的数字".to_string()),
None => None,
};
let request_timeout_secs = match payload.request_timeout {
Some(value) if (1.0..=600.0).contains(&value) => Some(value),
Some(_) => return Err("request_timeout 必须是 1 到 600 之间的数字".to_string()),
None => None,
};
let mut config_map = normalize_json_object(payload.config, "config")?
.and_then(|value| value.as_object().cloned())
.unwrap_or_default();
if let Some(value) = normalize_json_object(payload.pool_advanced, "pool_advanced")? {
config_map.insert("pool_advanced".to_string(), value);
}
if let Some(value) = normalize_json_object(payload.failover_rules, "failover_rules")? {
config_map.insert("failover_rules".to_string(), value);
}
if let Some(value) =
normalize_json_object(payload.claude_code_advanced, "claude_code_advanced")?
{
if provider_type != "claude_code" {
return Err("claude_code_advanced 仅适用于 provider_type=claude_code".to_string());
}
config_map.insert("claude_code_advanced".to_string(), value);
}
let config = (!config_map.is_empty()).then_some(serde_json::Value::Object(config_map));
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
let record = StoredProviderCatalogProvider::new(
Uuid::new_v4().to_string(),
name.to_string(),
website,
provider_type.clone(),
)
.map_err(|err| err.to_string())?
.with_description(
payload
.description
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty()),
)
.with_billing_fields(
Some(billing_type),
monthly_quota_usd,
None,
quota_reset_day,
quota_last_reset_at_unix_secs,
quota_expires_at_unix_secs,
)
.with_routing_fields(provider_priority)
.with_transport_fields(
is_active,
payload.keep_priority_on_conversion.unwrap_or(false),
provider_type_enables_format_conversion_by_default(&provider_type),
concurrent_limit,
max_retries,
proxy,
request_timeout_secs,
stream_first_byte_timeout_secs,
config,
)
.with_timestamps(Some(now_unix_secs), Some(now_unix_secs));
Ok((record, shift_existing_priorities_from))
}
pub(crate) fn build_admin_fixed_provider_endpoint_record(
provider: &StoredProviderCatalogProvider,
api_format: &str,
base_url: &str,
) -> Result<StoredProviderCatalogEndpoint, String> {
let (normalized_api_format, api_family, endpoint_kind) =
admin_endpoint_signature_parts(api_format)
.ok_or_else(|| format!("无效的 api_format: {api_format}"))?;
let body_rules = admin_default_body_rules_for_signature(
normalized_api_format,
Some(provider.provider_type.as_str()),
)
.and_then(|(_, rules)| (!rules.is_empty()).then_some(serde_json::Value::Array(rules)));
let endpoint_config =
if provider.provider_type == "codex" && normalized_api_format == "openai:cli" {
Some(json!({ "upstream_stream_policy": "force_stream" }))
} else {
None
};
let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
.unwrap_or(0);
StoredProviderCatalogEndpoint::new(
Uuid::new_v4().to_string(),
provider.id.clone(),
normalized_api_format.to_string(),
Some(api_family.to_string()),
Some(endpoint_kind.to_string()),
true,
)
.map_err(|err| err.to_string())?
.with_timestamps(Some(now_unix_secs), Some(now_unix_secs))
.with_transport_fields(
normalize_admin_base_url(base_url)?,
None,
body_rules,
Some(provider.max_retries.unwrap_or(2)),
None,
endpoint_config,
None,
None,
)
.map_err(|err| err.to_string())
}
@@ -0,0 +1,151 @@
fn normalize_reveal_auth_type(value: &str) -> &str {
match value.trim().to_ascii_lowercase().as_str() {
"service_account" | "vertex_ai" => "service_account",
"oauth" => "oauth",
_ => "api_key",
}
}
use crate::handlers::admin::shared::{
decrypt_catalog_secret_with_fallbacks, parse_catalog_auth_config_json,
};
use crate::AppState;
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use chrono::{SecondsFormat, Utc};
use serde_json::json;
pub(crate) fn build_admin_reveal_key_payload(
state: &AppState,
key: &StoredProviderCatalogKey,
) -> Result<serde_json::Value, String> {
let auth_type = normalize_reveal_auth_type(&key.auth_type);
if matches!(auth_type, "service_account") {
if let Some(auth_config) = parse_catalog_auth_config_json(state, key) {
return Ok(json!({
"auth_type": auth_type,
"auth_config": auth_config,
}));
}
let decrypted =
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), &key.encrypted_api_key)
.ok_or_else(|| {
"无法解密认证配置,可能是加密密钥已更改。请重新添加该密钥。".to_string()
})?;
if decrypted == "__placeholder__" {
return Err("认证配置丢失,请重新添加该密钥。".to_string());
}
return Ok(json!({
"auth_type": auth_type,
"auth_config": decrypted,
}));
}
let decrypted =
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), &key.encrypted_api_key)
.ok_or_else(|| {
"无法解密 API Key,可能是加密密钥已更改。请重新添加该密钥。".to_string()
})?;
Ok(json!({
"auth_type": auth_type,
"api_key": decrypted,
}))
}
fn provider_oauth_export_payload(
provider_type: &str,
auth_config: &serde_json::Map<String, serde_json::Value>,
upstream_metadata: Option<&serde_json::Value>,
) -> serde_json::Map<String, serde_json::Value> {
let normalized_provider_type = provider_type.trim().to_ascii_lowercase();
let skip_keys: &[&str] = match normalized_provider_type.as_str() {
"kiro" => &["access_token", "expires_at", "updated_at"],
_ => &[
"access_token",
"expires_at",
"updated_at",
"token_type",
"scope",
],
};
let mut payload = serde_json::Map::new();
for (key, value) in auth_config {
if skip_keys.contains(&key.as_str()) {
continue;
}
if value.is_null() || value.as_str().is_some_and(str::is_empty) {
continue;
}
payload.insert(key.clone(), value.clone());
}
if normalized_provider_type == "kiro" && !payload.contains_key("email") {
if let Some(email) = upstream_metadata
.and_then(serde_json::Value::as_object)
.and_then(|meta| meta.get("kiro"))
.and_then(serde_json::Value::as_object)
.and_then(|meta| meta.get("email"))
.and_then(serde_json::Value::as_str)
.filter(|value| !value.trim().is_empty())
{
payload.insert("email".to_string(), json!(email));
}
}
payload
}
pub(crate) async fn build_admin_export_key_payload(
state: &AppState,
key: &StoredProviderCatalogKey,
) -> Result<serde_json::Value, String> {
let auth_type = normalize_reveal_auth_type(&key.auth_type);
if auth_type != "oauth" {
return Err("仅 OAuth 类型的 Key 支持导出".to_string());
}
let ciphertext = key
.encrypted_auth_config
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.ok_or_else(|| "缺少认证配置,无法导出".to_string())?;
let plaintext = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)
.ok_or_else(|| "无法解密认证配置".to_string())?;
let auth_config = serde_json::from_str::<serde_json::Value>(&plaintext)
.ok()
.and_then(|value| value.as_object().cloned())
.ok_or_else(|| "无法解密认证配置".to_string())?;
if !auth_config
.get("refresh_token")
.and_then(serde_json::Value::as_str)
.is_some_and(|value| !value.trim().is_empty())
{
return Err("缺少 refresh_token,无法导出".to_string());
}
let provider_type_from_config = auth_config
.get("provider_type")
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let provider_type = if let Some(provider_type) = provider_type_from_config {
provider_type
} else {
state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&key.provider_id))
.await
.map_err(|err| format!("{err:?}"))?
.into_iter()
.next()
.map(|provider| provider.provider_type)
.unwrap_or_default()
};
let mut payload =
provider_oauth_export_payload(&provider_type, &auth_config, key.upstream_metadata.as_ref());
payload.insert("name".to_string(), json!(key.name));
payload.insert(
"exported_at".to_string(),
json!(Utc::now().to_rfc3339_opts(SecondsFormat::Secs, true)),
);
Ok(serde_json::Value::Object(payload))
}