mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 18:59:50 +08:00
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:
@@ -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, ®ion, &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,
|
||||
®ion,
|
||||
&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(¤t_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("a_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))
|
||||
}
|
||||
Reference in New Issue
Block a user