mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 18:59:50 +08:00
refactor: 大规模模块拆分与重组,新增 aether-admin crate
- 新建独立 aether-admin crate 承载 admin 相关共享契约与纯辅助函数 - 拆分 ai_pipeline 下 kiro/private_envelope/conversion/planner 等大文件为子模块目录 - 重组 admin handlers 各业务域(billing/oauth/provider/system/users 等)为目录结构,移除 shared.rs/builders.rs 等反模式 - 移除 ai_pipeline runtime adapters 旧实现(claude/openai/gemini/kiro/vertex/antigravity 等),改由 provider transport 统一承载 - 移除 control_facade/execution_facade/auth_snapshot_facade 等冗余 facade 层 - 拆分 query/billing 与 query/monitoring 模块、state/runtime/payments 与 security 模块 - 扩展架构测试覆盖 admin_billing/admin_model/admin_users 等新模块 - 删除 docs/architecture/refactor-execution-plan.md 已完成的执行计划文档
This commit is contained in:
@@ -0,0 +1,146 @@
|
||||
use crate::handlers::admin::provider::delete_task::run_admin_provider_delete_task;
|
||||
use crate::handlers::admin::provider::shared::paths::{
|
||||
admin_provider_delete_task_parts, admin_provider_id_for_manage_path,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::support::build_admin_provider_delete_task_payload;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::attach_admin_audit_response;
|
||||
use crate::{GatewayError, LocalProviderDeleteTaskState};
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use tracing::warn;
|
||||
use uuid::Uuid;
|
||||
|
||||
fn build_admin_provider_not_found_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
|
||||
fn build_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_provider_delete_task_response(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
route_kind: Option<&str>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
if route_kind == Some("delete_provider") && request_context.method() == http::Method::DELETE {
|
||||
let Some(provider_id) = admin_provider_id_for_manage_path(request_context.path()) else {
|
||||
return Ok(Some(build_admin_provider_not_found_response(
|
||||
"Provider 不存在",
|
||||
)));
|
||||
};
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(Some(build_admin_provider_not_found_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(),
|
||||
};
|
||||
state.put_provider_delete_task(pending_task.clone());
|
||||
if let Err(err) = state
|
||||
.run_admin_provider_delete_task(&provider.id, &task_id)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
"gateway admin provider delete task failed for provider {}: {:?}",
|
||||
provider.id, err
|
||||
);
|
||||
state.put_provider_delete_task(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 route_kind == Some("delete_provider_task") && request_context.method() == http::Method::GET {
|
||||
let Some((provider_id, task_id)) = admin_provider_delete_task_parts(request_context.path())
|
||||
else {
|
||||
return Ok(Some(build_admin_provider_not_found_response(
|
||||
"Task not found",
|
||||
)));
|
||||
};
|
||||
let Some(task) = state.get_provider_delete_task(&task_id) else {
|
||||
return Ok(Some(build_admin_provider_not_found_response(
|
||||
"Task not found",
|
||||
)));
|
||||
};
|
||||
if task.provider_id != provider_id {
|
||||
return Ok(Some(build_admin_provider_not_found_response(
|
||||
"Task not found",
|
||||
)));
|
||||
}
|
||||
return Ok(Some(build_admin_provider_delete_task_terminal_audit(
|
||||
&provider_id,
|
||||
&task_id,
|
||||
task.status.as_str(),
|
||||
Json(build_admin_provider_delete_task_payload(&task)).into_response(),
|
||||
)));
|
||||
}
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
@@ -1,4 +1,8 @@
|
||||
pub(crate) mod delete_task;
|
||||
pub(crate) mod pool;
|
||||
pub(crate) mod reads;
|
||||
mod responses;
|
||||
mod routes;
|
||||
pub(crate) mod writes;
|
||||
|
||||
pub(crate) use routes::maybe_build_local_admin_providers_response;
|
||||
pub(crate) use self::routes::maybe_build_local_admin_providers_response;
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
use crate::handlers::admin::provider::shared::paths::{
|
||||
admin_provider_clear_pool_cooldown_parts, admin_provider_reset_pool_cost_parts,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::attach_admin_audit_response;
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
enum AdminProviderPoolKeyLookup {
|
||||
ProviderMissing,
|
||||
KeyMissing,
|
||||
KeyName(String),
|
||||
}
|
||||
|
||||
fn build_admin_provider_not_found_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
|
||||
async fn lookup_admin_provider_pool_key_name(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
key_id: &str,
|
||||
) -> Result<AdminProviderPoolKeyLookup, GatewayError> {
|
||||
let provider_id_owned = provider_id.to_string();
|
||||
let provider_exists = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id_owned))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
.is_some();
|
||||
if !provider_exists {
|
||||
return Ok(AdminProviderPoolKeyLookup::ProviderMissing);
|
||||
}
|
||||
|
||||
let key = state
|
||||
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider_id_owned))
|
||||
.await?
|
||||
.into_iter()
|
||||
.find(|key| key.id == key_id)
|
||||
.map(|key| key.name);
|
||||
|
||||
Ok(match key {
|
||||
Some(name) => AdminProviderPoolKeyLookup::KeyName(name),
|
||||
None => AdminProviderPoolKeyLookup::KeyMissing,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_build_local_admin_provider_pool_response(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
route_kind: Option<&str>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
if route_kind == Some("clear_pool_cooldown") && request_context.method() == http::Method::POST {
|
||||
let Some((provider_id, key_id)) =
|
||||
admin_provider_clear_pool_cooldown_parts(request_context.path())
|
||||
else {
|
||||
return Ok(Some(build_admin_provider_not_found_response("Key 不存在")));
|
||||
};
|
||||
match lookup_admin_provider_pool_key_name(state, &provider_id, &key_id).await? {
|
||||
AdminProviderPoolKeyLookup::ProviderMissing => {
|
||||
return Ok(Some(build_admin_provider_not_found_response(format!(
|
||||
"Provider {provider_id} 不存在"
|
||||
))));
|
||||
}
|
||||
AdminProviderPoolKeyLookup::KeyMissing => {
|
||||
return Ok(Some(build_admin_provider_not_found_response(format!(
|
||||
"Key {key_id} 不存在"
|
||||
))));
|
||||
}
|
||||
AdminProviderPoolKeyLookup::KeyName(key_name) => {
|
||||
state
|
||||
.clear_admin_provider_pool_cooldown(&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 route_kind == Some("reset_pool_cost") && request_context.method() == http::Method::POST {
|
||||
let Some((provider_id, key_id)) =
|
||||
admin_provider_reset_pool_cost_parts(request_context.path())
|
||||
else {
|
||||
return Ok(Some(build_admin_provider_not_found_response("Key 不存在")));
|
||||
};
|
||||
match lookup_admin_provider_pool_key_name(state, &provider_id, &key_id).await? {
|
||||
AdminProviderPoolKeyLookup::ProviderMissing => {
|
||||
return Ok(Some(build_admin_provider_not_found_response(format!(
|
||||
"Provider {provider_id} 不存在"
|
||||
))));
|
||||
}
|
||||
AdminProviderPoolKeyLookup::KeyMissing => {
|
||||
return Ok(Some(build_admin_provider_not_found_response(format!(
|
||||
"Key {key_id} 不存在"
|
||||
))));
|
||||
}
|
||||
AdminProviderPoolKeyLookup::KeyName(key_name) => {
|
||||
state
|
||||
.reset_admin_provider_pool_cost(&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,195 @@
|
||||
use super::responses::build_admin_providers_data_unavailable_response;
|
||||
use crate::handlers::admin::provider::shared::paths::{
|
||||
admin_provider_id_for_health_monitor, admin_provider_id_for_mapping_preview,
|
||||
admin_provider_id_for_pool_status, admin_provider_id_for_summary, is_admin_providers_root,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::{query_param_optional_bool, query_param_value};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
fn build_admin_provider_not_found_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_build_local_admin_provider_reads_response(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
route_kind: Option<&str>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
if route_kind == Some("list_providers") && is_admin_providers_root(request_context.path()) {
|
||||
if !state.has_provider_catalog_data_reader() {
|
||||
return Ok(Some(build_admin_providers_data_unavailable_response()));
|
||||
}
|
||||
let skip = query_param_value(request_context.query_string(), "skip")
|
||||
.and_then(|value| value.parse::<usize>().ok())
|
||||
.unwrap_or(0);
|
||||
let limit = query_param_value(request_context.query_string(), "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.query_string(), "is_active");
|
||||
let Some(payload) = state
|
||||
.build_admin_providers_payload(skip, limit, is_active)
|
||||
.await
|
||||
else {
|
||||
return Ok(Some(build_admin_providers_data_unavailable_response()));
|
||||
};
|
||||
return Ok(Some(Json(payload).into_response()));
|
||||
}
|
||||
|
||||
if route_kind == Some("summary_list")
|
||||
&& request_context.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.query_string(), "page")
|
||||
.and_then(|value| value.parse::<usize>().ok())
|
||||
.filter(|value| *value > 0)
|
||||
.unwrap_or(1);
|
||||
let page_size = query_param_value(request_context.query_string(), "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.query_string(), "search").unwrap_or_default();
|
||||
let status = query_param_value(request_context.query_string(), "status")
|
||||
.unwrap_or_else(|| "all".to_string());
|
||||
let api_format = query_param_value(request_context.query_string(), "api_format")
|
||||
.unwrap_or_else(|| "all".to_string());
|
||||
let model_id = query_param_value(request_context.query_string(), "model_id")
|
||||
.unwrap_or_else(|| "all".to_string());
|
||||
let Some(payload) = state
|
||||
.build_admin_providers_summary_payload(
|
||||
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 route_kind == Some("provider_summary")
|
||||
&& request_context.path().starts_with("/api/admin/providers/")
|
||||
&& request_context.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.path()) else {
|
||||
return Ok(Some(build_admin_provider_not_found_response(
|
||||
"Provider 不存在",
|
||||
)));
|
||||
};
|
||||
return Ok(Some(
|
||||
match state
|
||||
.build_admin_provider_summary_payload(&provider_id)
|
||||
.await
|
||||
{
|
||||
Some(payload) => Json(payload).into_response(),
|
||||
None => build_admin_provider_not_found_response(format!(
|
||||
"Provider {provider_id} 不存在"
|
||||
)),
|
||||
},
|
||||
));
|
||||
}
|
||||
|
||||
if route_kind == Some("health_monitor")
|
||||
&& request_context.path().starts_with("/api/admin/providers/")
|
||||
&& request_context.path().ends_with("/health-monitor")
|
||||
{
|
||||
let Some(provider_id) = admin_provider_id_for_health_monitor(request_context.path()) else {
|
||||
return Ok(Some(build_admin_provider_not_found_response(
|
||||
"Provider 不存在",
|
||||
)));
|
||||
};
|
||||
let lookback_hours = query_param_value(request_context.query_string(), "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.query_string(), "per_endpoint_limit")
|
||||
.and_then(|value| value.parse::<usize>().ok())
|
||||
.filter(|value| (10..=200).contains(value))
|
||||
.unwrap_or(48);
|
||||
return Ok(Some(
|
||||
match state
|
||||
.build_admin_provider_health_monitor_payload(
|
||||
&provider_id,
|
||||
lookback_hours,
|
||||
per_endpoint_limit,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Some(payload) => Json(payload).into_response(),
|
||||
None => build_admin_provider_not_found_response(format!(
|
||||
"Provider {provider_id} 不存在"
|
||||
)),
|
||||
},
|
||||
));
|
||||
}
|
||||
|
||||
if route_kind == Some("mapping_preview")
|
||||
&& request_context.path().starts_with("/api/admin/providers/")
|
||||
&& request_context.path().ends_with("/mapping-preview")
|
||||
{
|
||||
let Some(provider_id) = admin_provider_id_for_mapping_preview(request_context.path())
|
||||
else {
|
||||
return Ok(Some(build_admin_provider_not_found_response(
|
||||
"Provider 不存在",
|
||||
)));
|
||||
};
|
||||
return Ok(Some(
|
||||
match state
|
||||
.build_admin_provider_mapping_preview_payload(&provider_id)
|
||||
.await
|
||||
{
|
||||
Some(payload) => Json(payload).into_response(),
|
||||
None => build_admin_provider_not_found_response(format!(
|
||||
"Provider {provider_id} 不存在"
|
||||
)),
|
||||
},
|
||||
));
|
||||
}
|
||||
|
||||
if route_kind == Some("pool_status")
|
||||
&& request_context.method() == http::Method::GET
|
||||
&& request_context.path().ends_with("/pool-status")
|
||||
{
|
||||
let Some(provider_id) = admin_provider_id_for_pool_status(request_context.path()) else {
|
||||
return Ok(Some(build_admin_provider_not_found_response(
|
||||
"Provider 不存在",
|
||||
)));
|
||||
};
|
||||
return Ok(Some(
|
||||
match state
|
||||
.build_admin_provider_pool_status_payload(&provider_id)
|
||||
.await
|
||||
{
|
||||
Some(payload) => Json(payload).into_response(),
|
||||
None => build_admin_provider_not_found_response(format!(
|
||||
"Provider {provider_id} 不存在"
|
||||
)),
|
||||
},
|
||||
));
|
||||
}
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
@@ -1,727 +1,16 @@
|
||||
use super::responses::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::runtime::{
|
||||
build_admin_provider_pool_status_payload, clear_admin_provider_pool_cooldown,
|
||||
reset_admin_provider_pool_cost,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::{
|
||||
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, is_admin_providers_root,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
AdminProviderCreateRequest, AdminProviderUpdateRequest,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::support::{
|
||||
build_admin_provider_delete_task_payload, put_admin_provider_delete_task,
|
||||
};
|
||||
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::provider::{
|
||||
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 crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
response::Response,
|
||||
};
|
||||
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,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
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,
|
||||
)
|
||||
state
|
||||
.maybe_build_admin_provider_crud_route_response(request_context, request_body)
|
||||
.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,191 @@
|
||||
use super::responses::build_admin_providers_data_unavailable_response;
|
||||
use crate::handlers::admin::provider::shared::paths::{
|
||||
admin_provider_id_for_manage_path, is_admin_providers_root,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
AdminProviderCreateRequest, AdminProviderUpdateRequest,
|
||||
};
|
||||
use crate::handlers::admin::provider::write::provider::build_admin_fixed_provider_endpoint_record;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::attach_admin_audit_response;
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
fn build_admin_provider_bad_request_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
|
||||
fn build_admin_provider_not_found_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_build_local_admin_provider_writes_response(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
route_kind: Option<&str>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
if route_kind == Some("create_provider")
|
||||
&& request_context.method() == http::Method::POST
|
||||
&& is_admin_providers_root(request_context.path())
|
||||
{
|
||||
let Some(request_body) = request_body else {
|
||||
return Ok(Some(build_admin_provider_bad_request_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(build_admin_provider_bad_request_response(
|
||||
"请求体必须是合法的 JSON 对象",
|
||||
)));
|
||||
}
|
||||
};
|
||||
let (record, shift_existing_priorities_from) =
|
||||
match state.build_admin_create_provider_record(payload).await {
|
||||
Ok(record) => record,
|
||||
Err(message) => {
|
||||
return Ok(Some(build_admin_provider_bad_request_response(message)));
|
||||
}
|
||||
};
|
||||
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)) =
|
||||
state.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(build_admin_provider_bad_request_response(message)));
|
||||
}
|
||||
};
|
||||
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 route_kind == Some("update_provider")
|
||||
&& request_context.method() == http::Method::PATCH
|
||||
&& request_context.path().starts_with("/api/admin/providers/")
|
||||
{
|
||||
let Some(provider_id) = admin_provider_id_for_manage_path(request_context.path()) else {
|
||||
return Ok(Some(build_admin_provider_not_found_response(
|
||||
"Provider 不存在",
|
||||
)));
|
||||
};
|
||||
let Some(request_body) = request_body else {
|
||||
return Ok(Some(build_admin_provider_bad_request_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(build_admin_provider_bad_request_response(
|
||||
"请求体必须是合法的 JSON 对象",
|
||||
)));
|
||||
}
|
||||
};
|
||||
let Some(raw_payload) = raw_value.as_object().cloned() else {
|
||||
return Ok(Some(build_admin_provider_bad_request_response(
|
||||
"请求体必须是合法的 JSON 对象",
|
||||
)));
|
||||
};
|
||||
let payload = match serde_json::from_value::<AdminProviderUpdateRequest>(raw_value) {
|
||||
Ok(payload) => payload,
|
||||
Err(_) => {
|
||||
return Ok(Some(build_admin_provider_bad_request_response(
|
||||
"请求体必须是合法的 JSON 对象",
|
||||
)));
|
||||
}
|
||||
};
|
||||
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(build_admin_provider_not_found_response(format!(
|
||||
"Provider {provider_id} 不存在"
|
||||
))));
|
||||
};
|
||||
let updated_record = match state
|
||||
.build_admin_update_provider_record(&existing_provider, &raw_payload, payload)
|
||||
.await
|
||||
{
|
||||
Ok(record) => record,
|
||||
Err(detail) => return Ok(Some(build_admin_provider_bad_request_response(detail))),
|
||||
};
|
||||
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 state
|
||||
.build_admin_provider_summary_payload(&provider_id)
|
||||
.await
|
||||
{
|
||||
Some(payload) => attach_admin_audit_response(
|
||||
Json(payload).into_response(),
|
||||
"admin_provider_updated",
|
||||
"update_provider",
|
||||
"provider",
|
||||
&provider_id,
|
||||
),
|
||||
None => build_admin_provider_not_found_response(format!(
|
||||
"Provider {provider_id} 不存在"
|
||||
)),
|
||||
},
|
||||
));
|
||||
}
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
@@ -2,9 +2,10 @@ use crate::handlers::admin::provider::shared::support::{
|
||||
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::request::AdminAppState;
|
||||
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 crate::{GatewayError, LocalProviderDeleteTaskState};
|
||||
use aether_data_contracts::repository::global_models::{
|
||||
AdminProviderModelListQuery, PublicGlobalModelQuery, StoredPublicGlobalModel,
|
||||
};
|
||||
@@ -13,10 +14,11 @@ use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(crate) async fn run_admin_provider_delete_task(
|
||||
state: &AppState,
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
task_id: &str,
|
||||
) -> Result<LocalProviderDeleteTaskState, GatewayError> {
|
||||
let app = state.as_ref();
|
||||
let Some(mut provider) = state
|
||||
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
|
||||
.await?
|
||||
@@ -70,7 +72,7 @@ pub(crate) async fn run_admin_provider_delete_task(
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0),
|
||||
);
|
||||
let _ = state.update_provider_catalog_provider(&provider).await?;
|
||||
let _ = app.update_provider_catalog_provider(&provider).await?;
|
||||
}
|
||||
task.stage = "disabling".to_string();
|
||||
task.message = "provider disabled; starting cleanup".to_string();
|
||||
@@ -81,8 +83,7 @@ pub(crate) async fn run_admin_provider_delete_task(
|
||||
.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)
|
||||
app.cleanup_deleted_provider_catalog_refs(&provider.id, &endpoint_ids, &key_ids)
|
||||
.await?;
|
||||
|
||||
task.stage = "deleting_models".to_string();
|
||||
@@ -122,7 +123,7 @@ pub(crate) async fn run_admin_provider_delete_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? {
|
||||
if !app.delete_provider_catalog_provider(&provider.id).await? {
|
||||
task.status = "failed".to_string();
|
||||
task.stage = "failed".to_string();
|
||||
task.message = "provider delete failed".to_string();
|
||||
@@ -155,7 +156,7 @@ pub(crate) fn public_global_model_mapping_patterns(model: &StoredPublicGlobalMod
|
||||
}
|
||||
|
||||
pub(crate) fn mapping_preview_masked_catalog_api_key(
|
||||
state: &AppState,
|
||||
state: &AdminAppState<'_>,
|
||||
key: &StoredProviderCatalogKey,
|
||||
) -> String {
|
||||
let ciphertext = key.encrypted_api_key.trim();
|
||||
@@ -181,20 +182,21 @@ pub(crate) fn mapping_preview_masked_catalog_api_key(
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_provider_mapping_preview_payload(
|
||||
state: &AppState,
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
) -> Option<serde_json::Value> {
|
||||
if !state.has_provider_catalog_data_reader() || !state.has_global_model_data_reader() {
|
||||
let app = state.as_ref();
|
||||
if !app.has_provider_catalog_data_reader() || !app.has_global_model_data_reader() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let provider = state
|
||||
let provider = app
|
||||
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
|
||||
.await
|
||||
.ok()
|
||||
.and_then(|mut providers| providers.drain(..).next())?;
|
||||
|
||||
let mut keys = state
|
||||
let mut keys = app
|
||||
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
|
||||
.await
|
||||
.ok()
|
||||
@@ -209,7 +211,7 @@ pub(crate) async fn build_admin_provider_mapping_preview_payload(
|
||||
keys.truncate(ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_KEYS);
|
||||
}
|
||||
|
||||
let public_models = state
|
||||
let public_models = app
|
||||
.list_public_global_models(&PublicGlobalModelQuery {
|
||||
offset: 0,
|
||||
limit: ADMIN_PROVIDER_MAPPING_PREVIEW_FETCH_LIMIT,
|
||||
|
||||
@@ -2,16 +2,16 @@ mod mutations;
|
||||
mod quota;
|
||||
mod reads;
|
||||
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
response::Response,
|
||||
};
|
||||
|
||||
pub(crate) async fn maybe_build_local_admin_endpoints_keys_response(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
if let Some(response) = reads::maybe_handle(state, request_context, request_body).await? {
|
||||
|
||||
@@ -1,426 +0,0 @@
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::{
|
||||
admin_clear_oauth_invalid_key_id, admin_provider_id_for_keys, admin_update_key_id,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
AdminProviderKeyBatchDeleteRequest, AdminProviderKeyCreateRequest,
|
||||
AdminProviderKeyUpdateRequest,
|
||||
};
|
||||
use crate::handlers::admin::shared::build_admin_provider_key_response;
|
||||
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};
|
||||
|
||||
use super::super::write::keys::{
|
||||
build_admin_create_provider_key_record, build_admin_update_provider_key_record,
|
||||
};
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
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("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("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(),
|
||||
));
|
||||
}
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
@@ -0,0 +1,89 @@
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyBatchDeleteRequest;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
};
|
||||
if decision.route_family.as_deref() != Some("endpoints_manage")
|
||||
|| decision.route_kind.as_deref() != Some("batch_delete_keys")
|
||||
|| request_context.method() != http::Method::POST
|
||||
|| request_context.path() != "/api/admin/endpoints/keys/batch-delete"
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let Some(request_body) = request_body else {
|
||||
return Ok(Some(bad_request_response("请求体不能为空")));
|
||||
};
|
||||
let payload = match serde_json::from_slice::<AdminProviderKeyBatchDeleteRequest>(request_body) {
|
||||
Ok(payload) => payload,
|
||||
Err(_) => return Ok(Some(bad_request_response("请求体必须是合法的 JSON 对象"))),
|
||||
};
|
||||
if payload.ids.len() > 100 {
|
||||
return Ok(Some(bad_request_response("ids 最多 100 个")));
|
||||
}
|
||||
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" }));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Some(
|
||||
Json(json!({
|
||||
"success_count": success_count,
|
||||
"failed_count": failed.len(),
|
||||
"failed": failed,
|
||||
}))
|
||||
.into_response(),
|
||||
))
|
||||
}
|
||||
|
||||
fn bad_request_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_id_for_keys;
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyCreateRequest;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::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: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
};
|
||||
if decision.route_family.as_deref() != Some("endpoints_manage")
|
||||
|| decision.route_kind.as_deref() != Some("create_provider_key")
|
||||
|| request_context.method() != http::Method::POST
|
||||
|| !request_context
|
||||
.path()
|
||||
.starts_with("/api/admin/endpoints/providers/")
|
||||
|| !request_context.path().ends_with("/keys")
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let Some(provider_id) = admin_provider_id_for_keys(request_context.path()) else {
|
||||
return Ok(Some(not_found_response("Provider 不存在")));
|
||||
};
|
||||
let Some(request_body) = request_body else {
|
||||
return Ok(Some(bad_request_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(bad_request_response("请求体必须是合法的 JSON 对象"))),
|
||||
};
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(Some(not_found_response(format!(
|
||||
"Provider {provider_id} 不存在"
|
||||
))));
|
||||
};
|
||||
let record = match state
|
||||
.build_admin_create_provider_key_record(&provider, payload)
|
||||
.await
|
||||
{
|
||||
Ok(record) => record,
|
||||
Err(detail) => return Ok(Some(bad_request_response(detail))),
|
||||
};
|
||||
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);
|
||||
|
||||
Ok(Some(
|
||||
Json(state.build_admin_provider_key_response(&created, now_unix_secs)).into_response(),
|
||||
))
|
||||
}
|
||||
|
||||
fn bad_request_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
|
||||
fn not_found_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
use crate::handlers::admin::provider::shared::paths::admin_update_key_id;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
_request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
};
|
||||
if decision.route_family.as_deref() != Some("endpoints_manage")
|
||||
|| decision.route_kind.as_deref() != Some("delete_key")
|
||||
|| request_context.method() != http::Method::DELETE
|
||||
|| !request_context
|
||||
.path()
|
||||
.starts_with("/api/admin/endpoints/keys/")
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let Some(key_id) = admin_update_key_id(request_context.path()) else {
|
||||
return Ok(Some(not_found_response("Key 不存在")));
|
||||
};
|
||||
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(not_found_response(format!("Key {key_id} 不存在"))));
|
||||
};
|
||||
if !state.delete_provider_catalog_key(&key_id).await? {
|
||||
return Ok(Some(not_found_response(format!("Key {key_id} 不存在"))));
|
||||
}
|
||||
|
||||
Ok(Some(
|
||||
Json(json!({
|
||||
"message": format!("Key {key_id} 已删除")
|
||||
}))
|
||||
.into_response(),
|
||||
))
|
||||
}
|
||||
|
||||
fn not_found_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
mod batch;
|
||||
mod create;
|
||||
mod delete;
|
||||
mod oauth_invalid;
|
||||
mod update;
|
||||
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
response::Response,
|
||||
};
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
if let Some(response) = update::maybe_handle(state, request_context, request_body).await? {
|
||||
return Ok(Some(response));
|
||||
}
|
||||
if let Some(response) = delete::maybe_handle(state, request_context, request_body).await? {
|
||||
return Ok(Some(response));
|
||||
}
|
||||
if let Some(response) = batch::maybe_handle(state, request_context, request_body).await? {
|
||||
return Ok(Some(response));
|
||||
}
|
||||
if let Some(response) =
|
||||
oauth_invalid::maybe_handle(state, request_context, request_body).await?
|
||||
{
|
||||
return Ok(Some(response));
|
||||
}
|
||||
if let Some(response) = create::maybe_handle(state, request_context, request_body).await? {
|
||||
return Ok(Some(response));
|
||||
}
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
+67
@@ -0,0 +1,67 @@
|
||||
use crate::handlers::admin::provider::shared::paths::admin_clear_oauth_invalid_key_id;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
_request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
};
|
||||
if decision.route_family.as_deref() != Some("endpoints_manage")
|
||||
|| decision.route_kind.as_deref() != Some("clear_oauth_invalid")
|
||||
|| request_context.method() != http::Method::POST
|
||||
|| !request_context
|
||||
.path()
|
||||
.starts_with("/api/admin/endpoints/keys/")
|
||||
|| !request_context.path().ends_with("/clear-oauth-invalid")
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let Some(key_id) = admin_clear_oauth_invalid_key_id(request_context.path()) else {
|
||||
return Ok(Some(not_found_response("Key 不存在")));
|
||||
};
|
||||
let Some(key) = state
|
||||
.read_provider_catalog_keys_by_ids(std::slice::from_ref(&key_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(Some(not_found_response(format!("Key {key_id} 不存在"))));
|
||||
};
|
||||
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?;
|
||||
Ok(Some(
|
||||
Json(json!({
|
||||
"message": "已清除 OAuth 失效标记"
|
||||
}))
|
||||
.into_response(),
|
||||
))
|
||||
}
|
||||
|
||||
fn not_found_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
use crate::handlers::admin::provider::shared::paths::admin_update_key_id;
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminProviderKeyUpdateRequest;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::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: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
};
|
||||
if decision.route_family.as_deref() != Some("endpoints_manage")
|
||||
|| decision.route_kind.as_deref() != Some("update_key")
|
||||
|| request_context.method() != http::Method::PUT
|
||||
|| !request_context
|
||||
.path()
|
||||
.starts_with("/api/admin/endpoints/keys/")
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let Some(key_id) = admin_update_key_id(request_context.path()) else {
|
||||
return Ok(Some(not_found_response("Key 不存在")));
|
||||
};
|
||||
let Some(request_body) = request_body else {
|
||||
return Ok(Some(bad_request_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(bad_request_response("请求体必须是合法的 JSON 对象"))),
|
||||
};
|
||||
let Some(raw_payload) = raw_value.as_object().cloned() else {
|
||||
return Ok(Some(bad_request_response("请求体必须是合法的 JSON 对象")));
|
||||
};
|
||||
let payload = match serde_json::from_value::<AdminProviderKeyUpdateRequest>(raw_value) {
|
||||
Ok(payload) => payload,
|
||||
Err(_) => return Ok(Some(bad_request_response("请求体必须是合法的 JSON 对象"))),
|
||||
};
|
||||
|
||||
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(not_found_response(format!("Key {key_id} 不存在"))));
|
||||
};
|
||||
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(not_found_response(format!(
|
||||
"Provider {} 不存在",
|
||||
existing_key.provider_id
|
||||
))));
|
||||
};
|
||||
|
||||
let updated_record = match state
|
||||
.build_admin_update_provider_key_record(&provider, &existing_key, &raw_payload, payload)
|
||||
.await
|
||||
{
|
||||
Ok(record) => record,
|
||||
Err(detail) => return Ok(Some(bad_request_response(detail))),
|
||||
};
|
||||
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);
|
||||
|
||||
Ok(Some(
|
||||
Json(state.build_admin_provider_key_response(&updated, now_unix_secs)).into_response(),
|
||||
))
|
||||
}
|
||||
|
||||
fn bad_request_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
|
||||
fn not_found_response(detail: impl Into<String>) -> Response<Body> {
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": detail.into() })),
|
||||
)
|
||||
.into_response()
|
||||
}
|
||||
@@ -1,9 +1,9 @@
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_id_for_refresh_quota;
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
AdminProviderQuotaRefreshRequest, OAUTH_ACCOUNT_BLOCK_PREFIX,
|
||||
};
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -19,27 +19,26 @@ use super::super::oauth::quota::kiro::refresh_kiro_provider_quota_locally;
|
||||
use super::super::oauth::quota::shared::normalize_string_id_list;
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.control_decision.as_ref() else {
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
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.method() != http::Method::POST
|
||||
|| !request_context
|
||||
.request_path
|
||||
.path()
|
||||
.starts_with("/api/admin/endpoints/providers/")
|
||||
|| !request_context.request_path.ends_with("/refresh-quota")
|
||||
|| !request_context.path().ends_with("/refresh-quota")
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let Some(provider_id) = admin_provider_id_for_refresh_quota(&request_context.request_path)
|
||||
else {
|
||||
let Some(provider_id) = admin_provider_id_for_refresh_quota(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
|
||||
@@ -1,10 +1,9 @@
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::{
|
||||
admin_export_key_id, admin_provider_id_for_keys, admin_reveal_key_id,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::{attach_admin_audit_response, query_param_value};
|
||||
use crate::handlers::public::build_admin_keys_grouped_by_format_payload;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -13,23 +12,20 @@ use axum::{
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::write::keys::build_admin_provider_keys_payload;
|
||||
use super::super::write::reveal::{build_admin_export_key_payload, build_admin_reveal_key_payload};
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
_request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.control_decision.as_ref() else {
|
||||
let Some(decision) = request_context.decision() 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"
|
||||
&& request_context.path() == "/api/admin/endpoints/keys/grouped-by-format"
|
||||
{
|
||||
let Some(payload) = build_admin_keys_grouped_by_format_payload(state).await else {
|
||||
let Some(payload) = state.build_admin_keys_grouped_by_format_payload().await else {
|
||||
return Ok(None);
|
||||
};
|
||||
return Ok(Some(Json(payload).into_response()));
|
||||
@@ -38,11 +34,11 @@ pub(super) async fn maybe_handle(
|
||||
if decision.route_family.as_deref() == Some("endpoints_manage")
|
||||
&& decision.route_kind.as_deref() == Some("reveal_key")
|
||||
&& request_context
|
||||
.request_path
|
||||
.path()
|
||||
.starts_with("/api/admin/endpoints/keys/")
|
||||
&& request_context.request_path.ends_with("/reveal")
|
||||
&& request_context.path().ends_with("/reveal")
|
||||
{
|
||||
let Some(key_id) = admin_reveal_key_id(&request_context.request_path) else {
|
||||
let Some(key_id) = admin_reveal_key_id(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
@@ -65,7 +61,7 @@ pub(super) async fn maybe_handle(
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
return Ok(Some(match build_admin_reveal_key_payload(state, &key) {
|
||||
return Ok(Some(match state.build_admin_reveal_key_payload(&key) {
|
||||
Ok(payload) => attach_admin_audit_response(
|
||||
Json(payload).into_response(),
|
||||
"admin_provider_key_revealed",
|
||||
@@ -84,11 +80,11 @@ pub(super) async fn maybe_handle(
|
||||
if decision.route_family.as_deref() == Some("endpoints_manage")
|
||||
&& decision.route_kind.as_deref() == Some("export_key")
|
||||
&& request_context
|
||||
.request_path
|
||||
.path()
|
||||
.starts_with("/api/admin/endpoints/keys/")
|
||||
&& request_context.request_path.ends_with("/export")
|
||||
&& request_context.path().ends_with("/export")
|
||||
{
|
||||
let Some(key_id) = admin_export_key_id(&request_context.request_path) else {
|
||||
let Some(key_id) = admin_export_key_id(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
@@ -112,7 +108,7 @@ pub(super) async fn maybe_handle(
|
||||
));
|
||||
};
|
||||
return Ok(Some(
|
||||
match build_admin_export_key_payload(state, &key).await {
|
||||
match state.build_admin_export_key_payload(&key).await {
|
||||
Ok(payload) => attach_admin_audit_response(
|
||||
Json(payload).into_response(),
|
||||
"admin_provider_key_exported",
|
||||
@@ -132,11 +128,11 @@ pub(super) async fn maybe_handle(
|
||||
if decision.route_family.as_deref() == Some("endpoints_manage")
|
||||
&& decision.route_kind.as_deref() == Some("list_provider_keys")
|
||||
&& request_context
|
||||
.request_path
|
||||
.path()
|
||||
.starts_with("/api/admin/endpoints/providers/")
|
||||
&& request_context.request_path.ends_with("/keys")
|
||||
&& request_context.path().ends_with("/keys")
|
||||
{
|
||||
let Some(provider_id) = admin_provider_id_for_keys(&request_context.request_path) else {
|
||||
let Some(provider_id) = admin_provider_id_for_keys(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
@@ -145,15 +141,18 @@ pub(super) async fn maybe_handle(
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let skip = query_param_value(request_context.request_query_string.as_deref(), "skip")
|
||||
let skip = query_param_value(request_context.query_string(), "skip")
|
||||
.and_then(|value| value.parse::<usize>().ok())
|
||||
.unwrap_or(0);
|
||||
let limit = query_param_value(request_context.request_query_string.as_deref(), "limit")
|
||||
let limit = query_param_value(request_context.query_string(), "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 {
|
||||
match state
|
||||
.build_admin_provider_keys_payload(&provider_id, skip, limit)
|
||||
.await
|
||||
{
|
||||
Some(payload) => Json(payload).into_response(),
|
||||
None => (
|
||||
http::StatusCode::NOT_FOUND,
|
||||
|
||||
@@ -1,355 +0,0 @@
|
||||
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)
|
||||
}
|
||||
@@ -1,9 +1,8 @@
|
||||
use super::builders::build_admin_create_provider_endpoint_record;
|
||||
use super::extractors::admin_provider_id_for_endpoints;
|
||||
use super::payloads::{build_admin_provider_endpoint_response, AdminProviderEndpointCreateRequest};
|
||||
use super::support::build_admin_endpoints_data_unavailable_response;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -14,21 +13,21 @@ use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.control_decision.as_ref() else {
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
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.method() != http::Method::POST
|
||||
|| !request_context
|
||||
.request_path
|
||||
.path()
|
||||
.starts_with("/api/admin/endpoints/providers/")
|
||||
|| !request_context.request_path.ends_with("/endpoints")
|
||||
|| !request_context.path().ends_with("/endpoints")
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
@@ -37,7 +36,7 @@ pub(super) async fn maybe_handle(
|
||||
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
|
||||
}
|
||||
|
||||
let Some(provider_id) = admin_provider_id_for_endpoints(&request_context.request_path) else {
|
||||
let Some(provider_id) = admin_provider_id_for_endpoints(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
@@ -81,7 +80,9 @@ pub(super) async fn maybe_handle(
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let record = match build_admin_create_provider_endpoint_record(state, &provider, payload).await
|
||||
let record = match state
|
||||
.build_admin_create_provider_endpoint_record(&provider, payload)
|
||||
.await
|
||||
{
|
||||
Ok(record) => record,
|
||||
Err(detail) => {
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
use super::extractors::admin_default_body_rules_api_format;
|
||||
use crate::api::ai::admin_default_body_rules_for_signature;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::query_param_value;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -12,26 +12,25 @@ use axum::{
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
_state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
_state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
_request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.control_decision.as_ref() else {
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
if decision.route_family.as_deref() != Some("endpoints_manage")
|
||||
|| decision.route_kind.as_deref() != Some("default_body_rules")
|
||||
|| !request_context
|
||||
.request_path
|
||||
.path()
|
||||
.starts_with("/api/admin/endpoints/defaults/")
|
||||
|| !request_context.request_path.ends_with("/body-rules")
|
||||
|| !request_context.path().ends_with("/body-rules")
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let Some(api_format) = admin_default_body_rules_api_format(&request_context.request_path)
|
||||
else {
|
||||
let Some(api_format) = admin_default_body_rules_api_format(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
@@ -40,10 +39,7 @@ pub(super) async fn maybe_handle(
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let provider_type = query_param_value(
|
||||
request_context.request_query_string.as_deref(),
|
||||
"provider_type",
|
||||
);
|
||||
let provider_type = query_param_value(request_context.query_string(), "provider_type");
|
||||
|
||||
Ok(Some(
|
||||
match admin_default_body_rules_for_signature(&api_format, provider_type.as_deref()) {
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
use super::extractors::admin_endpoint_id;
|
||||
use super::payloads::key_api_formats_without_entry;
|
||||
use super::support::build_admin_endpoints_data_unavailable_response;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -13,20 +13,18 @@ use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
_request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.control_decision.as_ref() else {
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
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/")
|
||||
|| request_context.method() != http::Method::DELETE
|
||||
|| !request_context.path().starts_with("/api/admin/endpoints/")
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
@@ -35,7 +33,7 @@ pub(super) async fn maybe_handle(
|
||||
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
|
||||
}
|
||||
|
||||
let Some(endpoint_id) = admin_endpoint_id(&request_context.request_path) else {
|
||||
let Some(endpoint_id) = admin_endpoint_id(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
use super::builders::build_admin_endpoint_payload;
|
||||
use super::extractors::admin_endpoint_id;
|
||||
use super::reads::build_admin_endpoint_payload;
|
||||
use super::support::build_admin_endpoints_data_unavailable_response;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -12,19 +12,17 @@ use axum::{
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
_request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.control_decision.as_ref() else {
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
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/")
|
||||
|| !request_context.path().starts_with("/api/admin/endpoints/")
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
@@ -33,7 +31,7 @@ pub(super) async fn maybe_handle(
|
||||
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
|
||||
}
|
||||
|
||||
let Some(endpoint_id) = admin_endpoint_id(&request_context.request_path) else {
|
||||
let Some(endpoint_id) = admin_endpoint_id(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
use super::builders::build_admin_provider_endpoints_payload;
|
||||
use super::extractors::admin_provider_id_for_endpoints;
|
||||
use super::reads::build_admin_provider_endpoints_payload;
|
||||
use super::support::build_admin_endpoints_data_unavailable_response;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::query_param_value;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -13,20 +13,20 @@ use axum::{
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
_request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.control_decision.as_ref() else {
|
||||
let Some(decision) = request_context.decision() 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
|
||||
.path()
|
||||
.starts_with("/api/admin/endpoints/providers/")
|
||||
|| !request_context.request_path.ends_with("/endpoints")
|
||||
|| !request_context.path().ends_with("/endpoints")
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
@@ -35,7 +35,7 @@ pub(super) async fn maybe_handle(
|
||||
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
|
||||
}
|
||||
|
||||
let Some(provider_id) = admin_provider_id_for_endpoints(&request_context.request_path) else {
|
||||
let Some(provider_id) = admin_provider_id_for_endpoints(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
@@ -44,10 +44,10 @@ pub(super) async fn maybe_handle(
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let skip = query_param_value(request_context.request_query_string.as_deref(), "skip")
|
||||
let skip = query_param_value(request_context.query_string(), "skip")
|
||||
.and_then(|value| value.parse::<usize>().ok())
|
||||
.unwrap_or(0);
|
||||
let limit = query_param_value(request_context.request_query_string.as_deref(), "limit")
|
||||
let limit = query_param_value(request_context.query_string(), "limit")
|
||||
.and_then(|value| value.parse::<usize>().ok())
|
||||
.filter(|value| *value > 0)
|
||||
.unwrap_or(100);
|
||||
|
||||
@@ -1,24 +1,24 @@
|
||||
mod builders;
|
||||
mod create;
|
||||
mod defaults;
|
||||
mod delete;
|
||||
mod detail;
|
||||
mod extractors;
|
||||
mod list;
|
||||
mod payloads;
|
||||
pub(crate) mod payloads;
|
||||
mod reads;
|
||||
mod support;
|
||||
mod update;
|
||||
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
response::Response,
|
||||
};
|
||||
|
||||
pub(crate) async fn maybe_build_local_admin_endpoints_routes_response(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
if let Some(response) = create::maybe_handle(state, request_context, request_body).await? {
|
||||
|
||||
@@ -1,52 +1,23 @@
|
||||
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
|
||||
use aether_admin::provider::endpoints as admin_provider_endpoints_pure;
|
||||
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(),
|
||||
)
|
||||
admin_provider_endpoints_pure::key_api_formats_without_entry(key, api_format)
|
||||
}
|
||||
|
||||
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)
|
||||
) -> (
|
||||
std::collections::BTreeMap<String, usize>,
|
||||
std::collections::BTreeMap<String, usize>,
|
||||
) {
|
||||
admin_provider_endpoints_pure::endpoint_key_counts_by_format(keys)
|
||||
}
|
||||
|
||||
pub(super) fn build_admin_provider_endpoint_response(
|
||||
@@ -56,46 +27,13 @@ pub(super) fn build_admin_provider_endpoint_response(
|
||||
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)
|
||||
admin_provider_endpoints_pure::build_admin_provider_endpoint_response(
|
||||
endpoint,
|
||||
provider_name,
|
||||
total_keys,
|
||||
active_keys,
|
||||
now_unix_secs,
|
||||
)
|
||||
}
|
||||
|
||||
fn default_admin_endpoint_max_retries() -> i32 {
|
||||
@@ -103,44 +41,44 @@ fn default_admin_endpoint_max_retries() -> i32 {
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub(super) struct AdminProviderEndpointCreateRequest {
|
||||
pub(super) provider_id: String,
|
||||
pub(super) api_format: String,
|
||||
pub(super) base_url: String,
|
||||
pub(crate) struct AdminProviderEndpointCreateRequest {
|
||||
pub(crate) provider_id: String,
|
||||
pub(crate) api_format: String,
|
||||
pub(crate) base_url: String,
|
||||
#[serde(default)]
|
||||
pub(super) custom_path: Option<String>,
|
||||
pub(crate) custom_path: Option<String>,
|
||||
#[serde(default)]
|
||||
pub(super) header_rules: Option<serde_json::Value>,
|
||||
pub(crate) header_rules: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub(super) body_rules: Option<serde_json::Value>,
|
||||
pub(crate) body_rules: Option<serde_json::Value>,
|
||||
#[serde(default = "default_admin_endpoint_max_retries")]
|
||||
pub(super) max_retries: i32,
|
||||
pub(crate) max_retries: i32,
|
||||
#[serde(default)]
|
||||
pub(super) config: Option<serde_json::Value>,
|
||||
pub(crate) config: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub(super) proxy: Option<serde_json::Value>,
|
||||
pub(crate) proxy: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub(super) format_acceptance_config: Option<serde_json::Value>,
|
||||
pub(crate) format_acceptance_config: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub(super) struct AdminProviderEndpointUpdateRequest {
|
||||
pub(crate) struct AdminProviderEndpointUpdateRequest {
|
||||
#[serde(default)]
|
||||
pub(super) base_url: Option<String>,
|
||||
pub(crate) base_url: Option<String>,
|
||||
#[serde(default)]
|
||||
pub(super) custom_path: Option<String>,
|
||||
pub(crate) custom_path: Option<String>,
|
||||
#[serde(default)]
|
||||
pub(super) header_rules: Option<serde_json::Value>,
|
||||
pub(crate) header_rules: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub(super) body_rules: Option<serde_json::Value>,
|
||||
pub(crate) body_rules: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub(super) max_retries: Option<i32>,
|
||||
pub(crate) max_retries: Option<i32>,
|
||||
#[serde(default)]
|
||||
pub(super) is_active: Option<bool>,
|
||||
pub(crate) is_active: Option<bool>,
|
||||
#[serde(default)]
|
||||
pub(super) config: Option<serde_json::Value>,
|
||||
pub(crate) config: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub(super) proxy: Option<serde_json::Value>,
|
||||
pub(crate) proxy: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
pub(super) format_acceptance_config: Option<serde_json::Value>,
|
||||
pub(crate) format_acceptance_config: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
|
||||
};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use super::payloads::{build_admin_provider_endpoint_response, endpoint_key_counts_by_format};
|
||||
|
||||
pub(crate) async fn build_admin_provider_endpoints_payload(
|
||||
state: &AdminAppState<'_>,
|
||||
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: &AdminAppState<'_>,
|
||||
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,
|
||||
))
|
||||
}
|
||||
@@ -1,12 +1,11 @@
|
||||
use super::builders::build_admin_update_provider_endpoint_record;
|
||||
use super::extractors::admin_endpoint_id;
|
||||
use super::payloads::{
|
||||
build_admin_provider_endpoint_response, endpoint_key_counts_by_format,
|
||||
AdminProviderEndpointUpdateRequest,
|
||||
};
|
||||
use super::support::build_admin_endpoints_data_unavailable_response;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -17,20 +16,18 @@ use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.control_decision.as_ref() else {
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
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/")
|
||||
|| request_context.method() != http::Method::PUT
|
||||
|| !request_context.path().starts_with("/api/admin/endpoints/")
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
@@ -39,7 +36,7 @@ pub(super) async fn maybe_handle(
|
||||
return Ok(Some(build_admin_endpoints_data_unavailable_response()));
|
||||
}
|
||||
|
||||
let Some(endpoint_id) = admin_endpoint_id(&request_context.request_path) else {
|
||||
let Some(endpoint_id) = admin_endpoint_id(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
@@ -120,14 +117,14 @@ pub(super) async fn maybe_handle(
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let updated_record = match build_admin_update_provider_endpoint_record(
|
||||
state,
|
||||
&provider,
|
||||
&existing_endpoint,
|
||||
&raw_payload,
|
||||
payload,
|
||||
)
|
||||
.await
|
||||
let updated_record = match state
|
||||
.build_admin_update_provider_endpoint_record(
|
||||
&provider,
|
||||
&existing_endpoint,
|
||||
&raw_payload,
|
||||
payload,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(record) => record,
|
||||
Err(detail) => {
|
||||
|
||||
@@ -5,18 +5,20 @@ pub(crate) mod ops;
|
||||
pub(crate) mod pool;
|
||||
pub(crate) mod pool_admin;
|
||||
pub(crate) mod shared;
|
||||
pub(crate) mod summary;
|
||||
pub(crate) mod write;
|
||||
|
||||
mod crud;
|
||||
mod delete_task;
|
||||
pub(crate) mod crud;
|
||||
pub(crate) mod delete_task;
|
||||
mod models;
|
||||
mod query;
|
||||
mod strategy;
|
||||
mod summary;
|
||||
pub(crate) mod query;
|
||||
mod routes;
|
||||
pub(crate) mod strategy;
|
||||
|
||||
pub(crate) use self::crud::maybe_build_local_admin_providers_response;
|
||||
pub(crate) use self::models::maybe_build_local_admin_provider_models_response;
|
||||
pub(super) 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;
|
||||
pub(super) use self::ops::maybe_build_local_admin_provider_ops_response;
|
||||
pub(super) use self::query::maybe_build_local_admin_provider_query_response;
|
||||
pub(super) use self::routes::maybe_build_local_admin_provider_response;
|
||||
pub(super) use self::strategy::maybe_build_local_admin_provider_strategy_response;
|
||||
|
||||
@@ -1,9 +1,7 @@
|
||||
use super::write::build_admin_batch_assign_global_models_payload;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_assign_global_models_path;
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminBatchAssignGlobalModelsRequest;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -13,17 +11,15 @@ use axum::{
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
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
|
||||
if request_context.route_family() == Some("provider_models_manage")
|
||||
&& request_context.route_kind() == Some("assign_global_models")
|
||||
&& request_context.method() == http::Method::POST
|
||||
{
|
||||
let Some(provider_id) =
|
||||
admin_provider_assign_global_models_path(&request_context.request_path)
|
||||
let Some(provider_id) = admin_provider_assign_global_models_path(request_context.path())
|
||||
else {
|
||||
return Ok(Some(
|
||||
(
|
||||
@@ -69,12 +65,9 @@ pub(super) async fn maybe_handle(
|
||||
));
|
||||
}
|
||||
};
|
||||
let payload = match build_admin_batch_assign_global_models_payload(
|
||||
state,
|
||||
&provider_id,
|
||||
payload.global_model_ids,
|
||||
)
|
||||
.await
|
||||
let payload = match state
|
||||
.build_admin_batch_assign_global_models_payload(&provider_id, payload.global_model_ids)
|
||||
.await
|
||||
{
|
||||
Ok(payload) => payload,
|
||||
Err(detail) => {
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
use super::write::build_admin_provider_available_source_models_payload;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_available_source_models_path;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -12,17 +10,15 @@ use axum::{
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
_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
|
||||
if request_context.route_family() == Some("provider_models_manage")
|
||||
&& request_context.route_kind() == Some("available_source_models")
|
||||
&& request_context.method() == http::Method::GET
|
||||
{
|
||||
let Some(provider_id) =
|
||||
admin_provider_available_source_models_path(&request_context.request_path)
|
||||
let Some(provider_id) = admin_provider_available_source_models_path(request_context.path())
|
||||
else {
|
||||
return Ok(Some(
|
||||
(
|
||||
@@ -33,7 +29,10 @@ pub(super) async fn maybe_handle(
|
||||
));
|
||||
};
|
||||
return Ok(Some(
|
||||
match build_admin_provider_available_source_models_payload(state, &provider_id).await {
|
||||
match state
|
||||
.build_admin_provider_available_source_models_payload(&provider_id)
|
||||
.await
|
||||
{
|
||||
Some(payload) => Json(payload).into_response(),
|
||||
None => (
|
||||
http::StatusCode::NOT_FOUND,
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
use super::payloads::{admin_provider_model_name_exists, build_admin_provider_model_response};
|
||||
use super::write::build_admin_provider_model_create_record;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_models_batch_path;
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminProviderModelCreateRequest;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -16,18 +14,16 @@ use std::collections::BTreeSet;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
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")
|
||||
if request_context.route_family() == Some("provider_models_manage")
|
||||
&& request_context.route_kind() == Some("batch_create_provider_models")
|
||||
&& request_context.method() == http::Method::POST
|
||||
&& request_context.path().ends_with("/models/batch")
|
||||
{
|
||||
let Some(provider_id) = admin_provider_models_batch_path(&request_context.request_path)
|
||||
else {
|
||||
let Some(provider_id) = admin_provider_models_batch_path(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
@@ -98,12 +94,9 @@ pub(super) async fn maybe_handle(
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let record = match build_admin_provider_model_create_record(
|
||||
state,
|
||||
&provider_id,
|
||||
payload,
|
||||
)
|
||||
.await
|
||||
let record = match state
|
||||
.build_admin_provider_model_create_record(&provider_id, payload)
|
||||
.await
|
||||
{
|
||||
Ok(record) => record,
|
||||
Err(detail) => {
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
use super::payloads::build_admin_provider_model_response;
|
||||
use super::write::build_admin_provider_model_create_record;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_id_for_models_list;
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminProviderModelCreateRequest;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -15,18 +13,16 @@ use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
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")
|
||||
if request_context.route_family() == Some("provider_models_manage")
|
||||
&& request_context.route_kind() == Some("create_provider_model")
|
||||
&& request_context.method() == http::Method::POST
|
||||
&& request_context.path().ends_with("/models")
|
||||
{
|
||||
let Some(provider_id) = admin_provider_id_for_models_list(&request_context.request_path)
|
||||
else {
|
||||
let Some(provider_id) = admin_provider_id_for_models_list(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
@@ -71,19 +67,21 @@ pub(super) async fn maybe_handle(
|
||||
));
|
||||
}
|
||||
};
|
||||
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 record = match state
|
||||
.build_admin_provider_model_create_record(&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) => {
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_model_route_parts;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -11,18 +10,17 @@ use axum::{
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
_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/")
|
||||
if request_context.route_family() == Some("provider_models_manage")
|
||||
&& request_context.route_kind() == Some("delete_provider_model")
|
||||
&& request_context.method() == http::Method::DELETE
|
||||
&& request_context.path().contains("/models/")
|
||||
{
|
||||
let Some((provider_id, model_id)) =
|
||||
admin_provider_model_route_parts(&request_context.request_path)
|
||||
admin_provider_model_route_parts(request_context.path())
|
||||
else {
|
||||
return Ok(Some(
|
||||
(
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
use super::payloads::build_admin_provider_model_payload;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_model_route_parts;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -12,21 +11,18 @@ use axum::{
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
_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/")
|
||||
if request_context.route_family() == Some("provider_models_manage")
|
||||
&& request_context.route_kind() == Some("get_provider_model")
|
||||
&& request_context.method() == http::Method::GET
|
||||
&& request_context.path().starts_with("/api/admin/providers/")
|
||||
&& request_context.path().contains("/models/")
|
||||
{
|
||||
let Some((provider_id, model_id)) =
|
||||
admin_provider_model_route_parts(&request_context.request_path)
|
||||
admin_provider_model_route_parts(request_context.path())
|
||||
else {
|
||||
return Ok(Some(
|
||||
(
|
||||
|
||||
@@ -1,9 +1,7 @@
|
||||
use super::write::build_admin_import_provider_models_payload;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_import_models_path;
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminImportProviderModelsRequest;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -13,17 +11,15 @@ use axum::{
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
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
|
||||
if request_context.route_family() == Some("provider_models_manage")
|
||||
&& request_context.route_kind() == Some("import_from_upstream")
|
||||
&& request_context.method() == http::Method::POST
|
||||
{
|
||||
let Some(provider_id) = admin_provider_import_models_path(&request_context.request_path)
|
||||
else {
|
||||
let Some(provider_id) = admin_provider_import_models_path(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
@@ -68,19 +64,21 @@ pub(super) async fn maybe_handle(
|
||||
));
|
||||
}
|
||||
};
|
||||
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(),
|
||||
));
|
||||
}
|
||||
};
|
||||
let payload = match state
|
||||
.build_admin_import_provider_models_payload(&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()));
|
||||
}
|
||||
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
use super::payloads::build_admin_provider_models_payload;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_id_for_models_list;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::{query_param_optional_bool, query_param_value};
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -13,21 +12,17 @@ use axum::{
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
_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")
|
||||
if request_context.route_family() == Some("provider_models_manage")
|
||||
&& request_context.route_kind() == Some("list_provider_models")
|
||||
&& request_context.method() == http::Method::GET
|
||||
&& request_context.path().starts_with("/api/admin/providers/")
|
||||
&& request_context.path().ends_with("/models")
|
||||
{
|
||||
let Some(provider_id) = admin_provider_id_for_models_list(&request_context.request_path)
|
||||
else {
|
||||
let Some(provider_id) = admin_provider_id_for_models_list(request_context.path()) else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
@@ -36,15 +31,14 @@ pub(super) async fn maybe_handle(
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let skip = query_param_value(request_context.request_query_string.as_deref(), "skip")
|
||||
let skip = query_param_value(request_context.query_string(), "skip")
|
||||
.and_then(|value| value.parse::<usize>().ok())
|
||||
.unwrap_or(0);
|
||||
let limit = query_param_value(request_context.request_query_string.as_deref(), "limit")
|
||||
let limit = query_param_value(request_context.query_string(), "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 is_active = query_param_optional_bool(request_context.query_string(), "is_active");
|
||||
return Ok(Some(
|
||||
match build_admin_provider_models_payload(state, &provider_id, skip, limit, is_active)
|
||||
.await
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::body::{Body, Bytes};
|
||||
use axum::http::Response;
|
||||
|
||||
@@ -14,68 +13,53 @@ mod import;
|
||||
mod list;
|
||||
mod payloads;
|
||||
mod update;
|
||||
mod write;
|
||||
|
||||
pub(crate) async fn maybe_build_local_admin_provider_models_response(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.control_decision.as_ref() else {
|
||||
if request_context.route_family() != Some("provider_models_manage") {
|
||||
return Ok(None);
|
||||
};
|
||||
}
|
||||
|
||||
if let Some(response) = list::maybe_handle(state, request_context, request_body).await? {
|
||||
return Ok(Some(response));
|
||||
}
|
||||
|
||||
if let Some(response) = detail::maybe_handle(state, request_context, request_body).await? {
|
||||
return Ok(Some(response));
|
||||
}
|
||||
|
||||
if let Some(response) = create::maybe_handle(state, request_context, request_body).await? {
|
||||
return Ok(Some(response));
|
||||
}
|
||||
|
||||
if let Some(response) = update::maybe_handle(state, request_context, request_body).await? {
|
||||
return Ok(Some(response));
|
||||
}
|
||||
|
||||
if let Some(response) = delete::maybe_handle(state, request_context, request_body).await? {
|
||||
return Ok(Some(response));
|
||||
}
|
||||
|
||||
if let Some(response) = batch::maybe_handle(state, request_context, request_body).await? {
|
||||
return Ok(Some(response));
|
||||
}
|
||||
|
||||
if let Some(response) =
|
||||
list::maybe_handle(state, request_context, request_body, decision).await?
|
||||
available_source::maybe_handle(state, request_context, request_body).await?
|
||||
{
|
||||
return Ok(Some(response));
|
||||
}
|
||||
|
||||
if let Some(response) =
|
||||
detail::maybe_handle(state, request_context, request_body, decision).await?
|
||||
assign_global::maybe_handle(state, request_context, request_body).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?
|
||||
{
|
||||
if let Some(response) = import::maybe_handle(state, request_context, request_body).await? {
|
||||
return Ok(Some(response));
|
||||
}
|
||||
|
||||
|
||||
@@ -1,192 +1,39 @@
|
||||
use crate::handlers::admin::shared::unix_secs_to_rfc3339;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::models as admin_provider_models_pure;
|
||||
use aether_data_contracts::repository::global_models::{
|
||||
AdminProviderModelListQuery, StoredAdminProviderModel,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
fn model_tiered_pricing_first_tier_value(
|
||||
tiered_pricing: Option<&serde_json::Value>,
|
||||
field_name: &str,
|
||||
) -> Option<f64> {
|
||||
tiered_pricing
|
||||
.and_then(|value| value.get("tiers"))
|
||||
.and_then(serde_json::Value::as_array)
|
||||
.and_then(|tiers| tiers.first())
|
||||
.and_then(|tier| tier.get(field_name))
|
||||
.and_then(serde_json::Value::as_f64)
|
||||
}
|
||||
|
||||
fn model_effective_capability(
|
||||
explicit: Option<bool>,
|
||||
global_model_config: Option<&serde_json::Value>,
|
||||
config_key: &str,
|
||||
) -> bool {
|
||||
explicit.unwrap_or_else(|| {
|
||||
global_model_config
|
||||
.and_then(|value| value.get(config_key))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
})
|
||||
}
|
||||
|
||||
fn merge_json_values(base: &mut serde_json::Value, overlay: serde_json::Value) {
|
||||
match (base, overlay) {
|
||||
(serde_json::Value::Object(base_map), serde_json::Value::Object(overlay_map)) => {
|
||||
for (key, value) in overlay_map {
|
||||
match base_map.get_mut(&key) {
|
||||
Some(existing) => merge_json_values(existing, value),
|
||||
None => {
|
||||
base_map.insert(key, value);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
(base, overlay) => *base = overlay,
|
||||
}
|
||||
}
|
||||
|
||||
fn merge_admin_provider_model_effective_config(
|
||||
model: &StoredAdminProviderModel,
|
||||
) -> Option<serde_json::Value> {
|
||||
let mut merged = match model.global_model_config.clone() {
|
||||
Some(serde_json::Value::Object(map)) => serde_json::Value::Object(map),
|
||||
Some(other) => other,
|
||||
None => serde_json::Value::Object(serde_json::Map::new()),
|
||||
};
|
||||
|
||||
if let Some(config) = model.config.clone() {
|
||||
merge_json_values(&mut merged, config);
|
||||
}
|
||||
|
||||
match merged {
|
||||
serde_json::Value::Null => None,
|
||||
serde_json::Value::Object(ref map) if map.is_empty() => None,
|
||||
value => Some(value),
|
||||
}
|
||||
}
|
||||
|
||||
fn 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(super) fn admin_provider_model_effective_input_price(
|
||||
model: &StoredAdminProviderModel,
|
||||
) -> Option<f64> {
|
||||
model_tiered_pricing_first_tier_value(model.tiered_pricing.as_ref(), "input_price_per_1m")
|
||||
.or_else(|| {
|
||||
model_tiered_pricing_first_tier_value(
|
||||
model.global_model_default_tiered_pricing.as_ref(),
|
||||
"input_price_per_1m",
|
||||
)
|
||||
})
|
||||
admin_provider_models_pure::admin_provider_model_effective_input_price(model)
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_model_effective_output_price(
|
||||
model: &StoredAdminProviderModel,
|
||||
) -> Option<f64> {
|
||||
model_tiered_pricing_first_tier_value(model.tiered_pricing.as_ref(), "output_price_per_1m")
|
||||
.or_else(|| {
|
||||
model_tiered_pricing_first_tier_value(
|
||||
model.global_model_default_tiered_pricing.as_ref(),
|
||||
"output_price_per_1m",
|
||||
)
|
||||
})
|
||||
admin_provider_models_pure::admin_provider_model_effective_output_price(model)
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_model_effective_capability(
|
||||
model: &StoredAdminProviderModel,
|
||||
capability: &str,
|
||||
) -> bool {
|
||||
match capability {
|
||||
"vision" => model_effective_capability(
|
||||
model.supports_vision,
|
||||
model.global_model_config.as_ref(),
|
||||
"vision",
|
||||
),
|
||||
"function_calling" => model_effective_capability(
|
||||
model.supports_function_calling,
|
||||
model.global_model_config.as_ref(),
|
||||
"function_calling",
|
||||
),
|
||||
"streaming" => model_effective_capability(
|
||||
model.supports_streaming,
|
||||
model.global_model_config.as_ref(),
|
||||
"streaming",
|
||||
),
|
||||
"extended_thinking" => model_effective_capability(
|
||||
model.supports_extended_thinking,
|
||||
model.global_model_config.as_ref(),
|
||||
"extended_thinking",
|
||||
),
|
||||
"image_generation" => model_effective_capability(
|
||||
model.supports_image_generation,
|
||||
model.global_model_config.as_ref(),
|
||||
"image_generation",
|
||||
),
|
||||
_ => false,
|
||||
}
|
||||
admin_provider_models_pure::admin_provider_model_effective_capability(model, capability)
|
||||
}
|
||||
|
||||
pub(super) fn build_admin_provider_model_response(
|
||||
model: &StoredAdminProviderModel,
|
||||
now_unix_secs: u64,
|
||||
) -> serde_json::Value {
|
||||
let effective_tiered_pricing = model
|
||||
.tiered_pricing
|
||||
.clone()
|
||||
.or_else(|| model.global_model_default_tiered_pricing.clone());
|
||||
let effective_config = merge_admin_provider_model_effective_config(model);
|
||||
|
||||
json!({
|
||||
"id": &model.id,
|
||||
"provider_id": &model.provider_id,
|
||||
"global_model_id": &model.global_model_id,
|
||||
"provider_model_name": &model.provider_model_name,
|
||||
"provider_model_mappings": model.provider_model_mappings.clone(),
|
||||
"price_per_request": model.price_per_request,
|
||||
"tiered_pricing": model.tiered_pricing.clone(),
|
||||
"effective_tiered_pricing": effective_tiered_pricing,
|
||||
"effective_input_price": admin_provider_model_effective_input_price(model),
|
||||
"effective_output_price": admin_provider_model_effective_output_price(model),
|
||||
"effective_price_per_request": model
|
||||
.price_per_request
|
||||
.or(model.global_model_default_price_per_request),
|
||||
"supports_vision": model.supports_vision,
|
||||
"supports_function_calling": model.supports_function_calling,
|
||||
"supports_streaming": model.supports_streaming,
|
||||
"supports_extended_thinking": model.supports_extended_thinking,
|
||||
"supports_image_generation": model.supports_image_generation,
|
||||
"effective_supports_vision": admin_provider_model_effective_capability(model, "vision"),
|
||||
"effective_supports_function_calling": admin_provider_model_effective_capability(
|
||||
model,
|
||||
"function_calling",
|
||||
),
|
||||
"effective_supports_streaming": admin_provider_model_effective_capability(model, "streaming"),
|
||||
"effective_supports_extended_thinking": admin_provider_model_effective_capability(
|
||||
model,
|
||||
"extended_thinking",
|
||||
),
|
||||
"effective_supports_image_generation": admin_provider_model_effective_capability(
|
||||
model,
|
||||
"image_generation",
|
||||
),
|
||||
"is_active": model.is_active,
|
||||
"is_available": model.is_available,
|
||||
"config": model.config.clone(),
|
||||
"effective_config": effective_config,
|
||||
"global_model_name": model.global_model_name.clone(),
|
||||
"global_model_display_name": model.global_model_display_name.clone(),
|
||||
"created_at": timestamp_or_now(model.created_at_unix_secs, now_unix_secs),
|
||||
"updated_at": timestamp_or_now(model.updated_at_unix_secs, now_unix_secs),
|
||||
})
|
||||
admin_provider_models_pure::build_admin_provider_model_response(model, now_unix_secs)
|
||||
}
|
||||
|
||||
pub(super) async fn build_admin_provider_models_payload(
|
||||
state: &AppState,
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
skip: usize,
|
||||
limit: usize,
|
||||
@@ -229,7 +76,7 @@ pub(super) async fn build_admin_provider_models_payload(
|
||||
}
|
||||
|
||||
pub(super) async fn build_admin_provider_model_payload(
|
||||
state: &AppState,
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
model_id: &str,
|
||||
) -> Option<serde_json::Value> {
|
||||
@@ -249,25 +96,12 @@ pub(super) async fn build_admin_provider_model_payload(
|
||||
}
|
||||
|
||||
pub(super) async fn admin_provider_model_name_exists(
|
||||
state: &AppState,
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
provider_model_name: &str,
|
||||
exclude_model_id: Option<&str>,
|
||||
) -> Result<bool, GatewayError> {
|
||||
let target = provider_model_name.trim();
|
||||
if target.is_empty() {
|
||||
return Ok(false);
|
||||
}
|
||||
let models = state
|
||||
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||
provider_id: provider_id.to_string(),
|
||||
is_active: None,
|
||||
offset: 0,
|
||||
limit: 10_000,
|
||||
})
|
||||
.await?;
|
||||
Ok(models.into_iter().any(|model| {
|
||||
model.provider_model_name == target
|
||||
&& exclude_model_id.is_none_or(|exclude| model.id != exclude)
|
||||
}))
|
||||
state
|
||||
.admin_provider_model_name_exists(provider_id, provider_model_name, exclude_model_id)
|
||||
.await
|
||||
}
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
use super::payloads::build_admin_provider_model_response;
|
||||
use super::write::build_admin_provider_model_update_record;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_model_route_parts;
|
||||
use crate::handlers::admin::provider::shared::payloads::AdminProviderModelUpdateRequest;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -15,18 +13,17 @@ use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(super) async fn maybe_handle(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
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/")
|
||||
if request_context.route_family() == Some("provider_models_manage")
|
||||
&& request_context.route_kind() == Some("update_provider_model")
|
||||
&& request_context.method() == http::Method::PATCH
|
||||
&& request_context.path().contains("/models/")
|
||||
{
|
||||
let Some((provider_id, model_id)) =
|
||||
admin_provider_model_route_parts(&request_context.request_path)
|
||||
admin_provider_model_route_parts(request_context.path())
|
||||
else {
|
||||
return Ok(Some(
|
||||
(
|
||||
@@ -90,21 +87,21 @@ pub(super) async fn maybe_handle(
|
||||
));
|
||||
}
|
||||
};
|
||||
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(),
|
||||
));
|
||||
}
|
||||
};
|
||||
let record = match state
|
||||
.build_admin_provider_model_update_record(&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) => {
|
||||
|
||||
@@ -1,472 +0,0 @@
|
||||
use super::payloads::{
|
||||
admin_provider_model_effective_capability, admin_provider_model_effective_input_price,
|
||||
admin_provider_model_effective_output_price, admin_provider_model_name_exists,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
AdminImportProviderModelsRequest, AdminProviderModelCreateRequest,
|
||||
AdminProviderModelUpdateRequest,
|
||||
};
|
||||
use crate::handlers::admin::shared::{
|
||||
normalize_json_array, normalize_json_object, normalize_string_list,
|
||||
};
|
||||
use crate::AppState;
|
||||
use aether_data_contracts::repository::global_models::{
|
||||
AdminProviderModelListQuery, CreateAdminGlobalModelRecord, UpdateAdminGlobalModelRecord,
|
||||
UpsertAdminProviderModelRecord,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use uuid::Uuid;
|
||||
|
||||
fn normalize_required_trimmed_string(value: &str, field_name: &str) -> Result<String, String> {
|
||||
let trimmed = value.trim();
|
||||
if trimmed.is_empty() {
|
||||
return Err(format!("{field_name} 不能为空"));
|
||||
}
|
||||
Ok(trimmed.to_string())
|
||||
}
|
||||
|
||||
fn normalize_optional_price(value: Option<f64>, field_name: &str) -> Result<Option<f64>, String> {
|
||||
let Some(value) = value else {
|
||||
return Ok(None);
|
||||
};
|
||||
if !value.is_finite() || value < 0.0 {
|
||||
return Err(format!("{field_name} 必须是非负数"));
|
||||
}
|
||||
Ok(Some(value))
|
||||
}
|
||||
|
||||
async fn resolve_admin_global_model_by_id_or_err(
|
||||
state: &AppState,
|
||||
global_model_id: &str,
|
||||
) -> Result<aether_data_contracts::repository::global_models::StoredAdminGlobalModel, String> {
|
||||
state
|
||||
.get_admin_global_model_by_id(global_model_id)
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?
|
||||
.ok_or_else(|| format!("GlobalModel {global_model_id} 不存在"))
|
||||
}
|
||||
|
||||
pub(super) async fn build_admin_provider_model_create_record(
|
||||
state: &AppState,
|
||||
provider_id: &str,
|
||||
payload: AdminProviderModelCreateRequest,
|
||||
) -> Result<UpsertAdminProviderModelRecord, String> {
|
||||
let provider_model_name =
|
||||
normalize_required_trimmed_string(&payload.provider_model_name, "provider_model_name")?;
|
||||
if admin_provider_model_name_exists(state, provider_id, &provider_model_name, None)
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?
|
||||
{
|
||||
return Err(format!("模型 '{provider_model_name}' 已存在"));
|
||||
}
|
||||
let global_model_id =
|
||||
normalize_required_trimmed_string(&payload.global_model_id, "global_model_id")?;
|
||||
resolve_admin_global_model_by_id_or_err(state, &global_model_id).await?;
|
||||
let price_per_request =
|
||||
normalize_optional_price(payload.price_per_request, "price_per_request")?;
|
||||
let tiered_pricing = normalize_json_object(payload.tiered_pricing, "tiered_pricing")?;
|
||||
let provider_model_mappings =
|
||||
normalize_json_array(payload.provider_model_mappings, "provider_model_mappings")?;
|
||||
let config = normalize_json_object(payload.config, "config")?;
|
||||
UpsertAdminProviderModelRecord::new(
|
||||
Uuid::new_v4().to_string(),
|
||||
provider_id.to_string(),
|
||||
global_model_id,
|
||||
provider_model_name,
|
||||
provider_model_mappings,
|
||||
price_per_request,
|
||||
tiered_pricing,
|
||||
payload.supports_vision,
|
||||
payload.supports_function_calling,
|
||||
payload.supports_streaming,
|
||||
payload.supports_extended_thinking,
|
||||
None,
|
||||
payload.is_active.unwrap_or(true),
|
||||
true,
|
||||
config,
|
||||
)
|
||||
.map_err(|err| err.to_string())
|
||||
}
|
||||
|
||||
pub(super) async fn build_admin_provider_model_update_record(
|
||||
state: &AppState,
|
||||
existing: &aether_data_contracts::repository::global_models::StoredAdminProviderModel,
|
||||
raw_payload: &serde_json::Map<String, serde_json::Value>,
|
||||
payload: AdminProviderModelUpdateRequest,
|
||||
) -> Result<UpsertAdminProviderModelRecord, String> {
|
||||
let provider_model_name = if let Some(value) = raw_payload.get("provider_model_name") {
|
||||
let Some(name) = payload.provider_model_name.as_deref() else {
|
||||
return Err(if value.is_null() {
|
||||
"provider_model_name 不能为空".to_string()
|
||||
} else {
|
||||
"provider_model_name 必须是字符串".to_string()
|
||||
});
|
||||
};
|
||||
let name = normalize_required_trimmed_string(name, "provider_model_name")?;
|
||||
if admin_provider_model_name_exists(state, &existing.provider_id, &name, Some(&existing.id))
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?
|
||||
{
|
||||
return Err(format!("模型 '{name}' 已存在"));
|
||||
}
|
||||
name
|
||||
} else {
|
||||
existing.provider_model_name.clone()
|
||||
};
|
||||
|
||||
let global_model_id = if let Some(value) = raw_payload.get("global_model_id") {
|
||||
let Some(global_model_id) = payload.global_model_id.as_deref() else {
|
||||
return Err(if value.is_null() {
|
||||
"global_model_id 不能为空".to_string()
|
||||
} else {
|
||||
"global_model_id 必须是字符串".to_string()
|
||||
});
|
||||
};
|
||||
let global_model_id =
|
||||
normalize_required_trimmed_string(global_model_id, "global_model_id")?;
|
||||
resolve_admin_global_model_by_id_or_err(state, &global_model_id).await?;
|
||||
global_model_id
|
||||
} else {
|
||||
existing.global_model_id.clone()
|
||||
};
|
||||
|
||||
let price_per_request = if raw_payload.contains_key("price_per_request") {
|
||||
normalize_optional_price(payload.price_per_request, "price_per_request")?
|
||||
} else {
|
||||
existing.price_per_request
|
||||
};
|
||||
let tiered_pricing = if raw_payload.contains_key("tiered_pricing") {
|
||||
normalize_json_object(payload.tiered_pricing, "tiered_pricing")?
|
||||
} else {
|
||||
existing.tiered_pricing.clone()
|
||||
};
|
||||
let provider_model_mappings = if raw_payload.contains_key("provider_model_mappings") {
|
||||
normalize_json_array(payload.provider_model_mappings, "provider_model_mappings")?
|
||||
} else {
|
||||
existing.provider_model_mappings.clone()
|
||||
};
|
||||
let config = if raw_payload.contains_key("config") {
|
||||
normalize_json_object(payload.config, "config")?
|
||||
} else {
|
||||
existing.config.clone()
|
||||
};
|
||||
|
||||
UpsertAdminProviderModelRecord::new(
|
||||
existing.id.clone(),
|
||||
existing.provider_id.clone(),
|
||||
global_model_id,
|
||||
provider_model_name,
|
||||
provider_model_mappings,
|
||||
price_per_request,
|
||||
tiered_pricing,
|
||||
if raw_payload.contains_key("supports_vision") {
|
||||
payload.supports_vision
|
||||
} else {
|
||||
existing.supports_vision
|
||||
},
|
||||
if raw_payload.contains_key("supports_function_calling") {
|
||||
payload.supports_function_calling
|
||||
} else {
|
||||
existing.supports_function_calling
|
||||
},
|
||||
if raw_payload.contains_key("supports_streaming") {
|
||||
payload.supports_streaming
|
||||
} else {
|
||||
existing.supports_streaming
|
||||
},
|
||||
if raw_payload.contains_key("supports_extended_thinking") {
|
||||
payload.supports_extended_thinking
|
||||
} else {
|
||||
existing.supports_extended_thinking
|
||||
},
|
||||
existing.supports_image_generation,
|
||||
payload.is_active.unwrap_or(existing.is_active),
|
||||
payload.is_available.unwrap_or(existing.is_available),
|
||||
config,
|
||||
)
|
||||
.map_err(|err| err.to_string())
|
||||
}
|
||||
|
||||
pub(super) async fn build_admin_provider_available_source_models_payload(
|
||||
state: &AppState,
|
||||
provider_id: &str,
|
||||
) -> Option<serde_json::Value> {
|
||||
if !state.has_global_model_data_reader() || !state.has_provider_catalog_data_reader() {
|
||||
return None;
|
||||
}
|
||||
let provider = state
|
||||
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
|
||||
.await
|
||||
.ok()?
|
||||
.into_iter()
|
||||
.next()?;
|
||||
let models = state
|
||||
.list_admin_provider_available_source_models(&provider.id)
|
||||
.await
|
||||
.ok()?;
|
||||
let mut by_global_model = BTreeMap::<
|
||||
String,
|
||||
aether_data_contracts::repository::global_models::StoredAdminProviderModel,
|
||||
>::new();
|
||||
for model in models {
|
||||
by_global_model
|
||||
.entry(model.global_model_id.clone())
|
||||
.or_insert(model);
|
||||
}
|
||||
let mut payload_models = by_global_model
|
||||
.into_values()
|
||||
.map(|model| {
|
||||
json!({
|
||||
"global_model_name": model.global_model_name,
|
||||
"display_name": model.global_model_display_name,
|
||||
"provider_model_name": model.provider_model_name,
|
||||
"model_id": model.id,
|
||||
"price": {
|
||||
"input_price_per_1m": admin_provider_model_effective_input_price(&model),
|
||||
"output_price_per_1m": admin_provider_model_effective_output_price(&model),
|
||||
"cache_creation_price_per_1m": serde_json::Value::Null,
|
||||
"cache_read_price_per_1m": serde_json::Value::Null,
|
||||
"price_per_request": model.price_per_request.or(model.global_model_default_price_per_request),
|
||||
},
|
||||
"capabilities": json!({
|
||||
"supports_vision": admin_provider_model_effective_capability(&model, "vision"),
|
||||
"supports_function_calling": admin_provider_model_effective_capability(&model, "function_calling"),
|
||||
"supports_streaming": admin_provider_model_effective_capability(&model, "streaming"),
|
||||
}),
|
||||
"is_active": model.is_active,
|
||||
})
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let total = payload_models.len();
|
||||
payload_models.sort_by(|left, right| {
|
||||
left.get("global_model_name")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.cmp(
|
||||
&right
|
||||
.get("global_model_name")
|
||||
.and_then(serde_json::Value::as_str),
|
||||
)
|
||||
});
|
||||
Some(json!({
|
||||
"models": payload_models,
|
||||
"total": total,
|
||||
}))
|
||||
}
|
||||
|
||||
pub(super) async fn build_admin_batch_assign_global_models_payload(
|
||||
state: &AppState,
|
||||
provider_id: &str,
|
||||
global_model_ids: Vec<String>,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let existing_models = state
|
||||
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||
provider_id: provider_id.to_string(),
|
||||
is_active: None,
|
||||
offset: 0,
|
||||
limit: 10_000,
|
||||
})
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?;
|
||||
let existing_global_model_ids = existing_models
|
||||
.into_iter()
|
||||
.map(|model| model.global_model_id)
|
||||
.collect::<BTreeSet<_>>();
|
||||
|
||||
let mut success = Vec::new();
|
||||
let mut errors = Vec::new();
|
||||
for global_model_id in global_model_ids {
|
||||
let global_model_id = global_model_id.trim().to_string();
|
||||
if global_model_id.is_empty() {
|
||||
continue;
|
||||
}
|
||||
let global_model =
|
||||
match resolve_admin_global_model_by_id_or_err(state, &global_model_id).await {
|
||||
Ok(model) => model,
|
||||
Err(detail) => {
|
||||
errors.push(json!({
|
||||
"global_model_id": global_model_id,
|
||||
"error": detail,
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
};
|
||||
if existing_global_model_ids.contains(&global_model.id) {
|
||||
errors.push(json!({
|
||||
"global_model_id": global_model.id,
|
||||
"error": "Model already exists",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
let record = UpsertAdminProviderModelRecord::new(
|
||||
Uuid::new_v4().to_string(),
|
||||
provider_id.to_string(),
|
||||
global_model.id.clone(),
|
||||
global_model.name.clone(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
true,
|
||||
None,
|
||||
)
|
||||
.map_err(|err| err.to_string())?;
|
||||
match state.create_admin_provider_model(&record).await {
|
||||
Ok(Some(created)) => success.push(json!({
|
||||
"global_model_id": global_model.id,
|
||||
"global_model_name": global_model.name,
|
||||
"provider_model_id": created.id,
|
||||
})),
|
||||
Ok(None) => errors.push(json!({
|
||||
"global_model_id": global_model.id,
|
||||
"error": "Create provider model failed",
|
||||
})),
|
||||
Err(err) => errors.push(json!({
|
||||
"global_model_id": global_model.id,
|
||||
"error": format!("{err:?}"),
|
||||
})),
|
||||
}
|
||||
}
|
||||
Ok(json!({
|
||||
"success": success,
|
||||
"errors": errors,
|
||||
}))
|
||||
}
|
||||
|
||||
pub(super) async fn build_admin_import_provider_models_payload(
|
||||
state: &AppState,
|
||||
provider_id: &str,
|
||||
payload: AdminImportProviderModelsRequest,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let default_pricing = json!({
|
||||
"tiers": [{
|
||||
"up_to": null,
|
||||
"input_price_per_1m": 0.0,
|
||||
"output_price_per_1m": 0.0,
|
||||
}]
|
||||
});
|
||||
let tiered_pricing = normalize_json_object(payload.tiered_pricing, "tiered_pricing")?;
|
||||
|
||||
let existing_models = state
|
||||
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||
provider_id: provider_id.to_string(),
|
||||
is_active: None,
|
||||
offset: 0,
|
||||
limit: 10_000,
|
||||
})
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?;
|
||||
let mut existing_by_name = existing_models
|
||||
.iter()
|
||||
.map(|model| (model.provider_model_name.clone(), model.clone()))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
|
||||
let mut success = Vec::new();
|
||||
let mut errors = Vec::new();
|
||||
|
||||
for model_id in payload.model_ids {
|
||||
let trimmed = model_id.trim();
|
||||
if trimmed.is_empty() || trimmed.len() > 100 {
|
||||
errors.push(json!({
|
||||
"model_id": if trimmed.is_empty() { "<empty>" } else { trimmed },
|
||||
"error": "Invalid model_id: must be 1-100 characters",
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
|
||||
if let Some(existing) = existing_by_name.get(trimmed) {
|
||||
success.push(json!({
|
||||
"model_id": trimmed,
|
||||
"global_model_id": existing.global_model_id,
|
||||
"global_model_name": existing.global_model_name,
|
||||
"provider_model_id": existing.id,
|
||||
"created_global_model": false,
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
|
||||
let mut created_global_model = false;
|
||||
let global_model = if let Some(existing) = state
|
||||
.get_admin_global_model_by_name(trimmed)
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?
|
||||
{
|
||||
existing
|
||||
} else {
|
||||
let created = state
|
||||
.create_admin_global_model(
|
||||
&CreateAdminGlobalModelRecord::new(
|
||||
Uuid::new_v4().to_string(),
|
||||
trimmed.to_string(),
|
||||
trimmed.to_string(),
|
||||
true,
|
||||
payload.price_per_request,
|
||||
tiered_pricing
|
||||
.clone()
|
||||
.or_else(|| Some(default_pricing.clone())),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.map_err(|err| err.to_string())?,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?;
|
||||
let Some(created) = created else {
|
||||
errors.push(json!({"model_id": trimmed, "error": "Create GlobalModel failed"}));
|
||||
continue;
|
||||
};
|
||||
created_global_model = true;
|
||||
created
|
||||
};
|
||||
|
||||
let record = UpsertAdminProviderModelRecord::new(
|
||||
Uuid::new_v4().to_string(),
|
||||
provider_id.to_string(),
|
||||
global_model.id.clone(),
|
||||
trimmed.to_string(),
|
||||
None,
|
||||
payload.price_per_request,
|
||||
tiered_pricing.clone(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
true,
|
||||
None,
|
||||
)
|
||||
.map_err(|err| err.to_string())?;
|
||||
|
||||
match state.create_admin_provider_model(&record).await {
|
||||
Ok(Some(created)) => {
|
||||
existing_by_name.insert(trimmed.to_string(), created.clone());
|
||||
success.push(json!({
|
||||
"model_id": trimmed,
|
||||
"global_model_id": global_model.id,
|
||||
"global_model_name": global_model.name,
|
||||
"provider_model_id": created.id,
|
||||
"created_global_model": created_global_model,
|
||||
}));
|
||||
}
|
||||
Ok(None) => errors.push(json!({
|
||||
"model_id": trimmed,
|
||||
"error": "Create provider model failed",
|
||||
})),
|
||||
Err(err) => errors.push(json!({
|
||||
"model_id": trimmed,
|
||||
"error": format!("{err:?}"),
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
Ok(json!({
|
||||
"success": success,
|
||||
"errors": errors,
|
||||
}))
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,274 @@
|
||||
use super::kiro_import::execute_admin_provider_oauth_kiro_batch_import;
|
||||
use super::parse::{
|
||||
apply_admin_provider_oauth_batch_import_hints, extract_admin_provider_oauth_batch_error_detail,
|
||||
parse_admin_provider_oauth_batch_import_entries, AdminProviderOAuthBatchImportEntry,
|
||||
AdminProviderOAuthBatchImportOutcome,
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::duplicates::find_duplicate_provider_oauth_key;
|
||||
use crate::handlers::admin::provider::oauth::provisioning::build_provider_oauth_auth_config_from_token_payload;
|
||||
use crate::handlers::admin::provider::oauth::provisioning::{
|
||||
create_provider_oauth_catalog_key, provider_oauth_active_api_formats,
|
||||
provider_oauth_key_proxy_value, update_existing_provider_oauth_catalog_key,
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::runtime::refresh_provider_oauth_account_state_after_update;
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
admin_provider_oauth_template, exchange_admin_provider_oauth_refresh_token,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::support::ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::oauth::parse_admin_provider_oauth_kiro_batch_import_entries;
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(super) fn estimate_admin_provider_oauth_batch_import_total(
|
||||
provider_type: &str,
|
||||
raw_credentials: &str,
|
||||
) -> usize {
|
||||
if provider_type.eq_ignore_ascii_case("kiro") {
|
||||
parse_admin_provider_oauth_kiro_batch_import_entries(raw_credentials).len()
|
||||
} else {
|
||||
parse_admin_provider_oauth_batch_import_entries(raw_credentials).len()
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn execute_admin_provider_oauth_batch_import_for_provider_type(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
provider_type: &str,
|
||||
raw_credentials: &str,
|
||||
proxy_node_id: Option<&str>,
|
||||
) -> Result<AdminProviderOAuthBatchImportOutcome, GatewayError> {
|
||||
if provider_type.eq_ignore_ascii_case("kiro") {
|
||||
execute_admin_provider_oauth_kiro_batch_import(
|
||||
state,
|
||||
provider_id,
|
||||
raw_credentials,
|
||||
proxy_node_id,
|
||||
)
|
||||
.await
|
||||
} else {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(raw_credentials);
|
||||
execute_admin_provider_oauth_batch_import(
|
||||
state,
|
||||
provider_id,
|
||||
provider_type,
|
||||
&entries,
|
||||
proxy_node_id,
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
provider_type: &str,
|
||||
entries: &[AdminProviderOAuthBatchImportEntry],
|
||||
proxy_node_id: Option<&str>,
|
||||
) -> Result<AdminProviderOAuthBatchImportOutcome, GatewayError> {
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(AdminProviderOAuthBatchImportOutcome {
|
||||
total: entries.len(),
|
||||
success: 0,
|
||||
failed: entries.len(),
|
||||
results: entries
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, _)| {
|
||||
json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": "Provider 不存在",
|
||||
"replaced": false,
|
||||
})
|
||||
})
|
||||
.collect(),
|
||||
});
|
||||
};
|
||||
|
||||
let Some(template) = admin_provider_oauth_template(provider_type) else {
|
||||
return Ok(AdminProviderOAuthBatchImportOutcome {
|
||||
total: entries.len(),
|
||||
success: 0,
|
||||
failed: entries.len(),
|
||||
results: entries
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, _)| {
|
||||
json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL,
|
||||
"replaced": false,
|
||||
})
|
||||
})
|
||||
.collect(),
|
||||
});
|
||||
};
|
||||
|
||||
let endpoints = state
|
||||
.list_provider_catalog_endpoints_by_provider_ids(&[provider_id.to_string()])
|
||||
.await?;
|
||||
let api_formats = provider_oauth_active_api_formats(&endpoints);
|
||||
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id);
|
||||
let mut results = Vec::with_capacity(entries.len());
|
||||
let mut success = 0usize;
|
||||
let mut failed = 0usize;
|
||||
|
||||
for (index, entry) in entries.iter().enumerate() {
|
||||
let token_payload = match exchange_admin_provider_oauth_refresh_token(
|
||||
state,
|
||||
template,
|
||||
entry.refresh_token.as_str(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(payload) => payload,
|
||||
Err(response) => {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": format!(
|
||||
"Token 验证失败: {}",
|
||||
extract_admin_provider_oauth_batch_error_detail(response).await
|
||||
),
|
||||
"replaced": false,
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
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 {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": "Token 刷新返回缺少 access_token",
|
||||
"replaced": false,
|
||||
}));
|
||||
continue;
|
||||
};
|
||||
|
||||
let refresh_token = returned_refresh_token
|
||||
.or_else(|| Some(entry.refresh_token.clone()))
|
||||
.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));
|
||||
}
|
||||
apply_admin_provider_oauth_batch_import_hints(provider_type, entry, &mut auth_config);
|
||||
|
||||
let duplicate =
|
||||
match find_duplicate_provider_oauth_key(state, provider_id, &auth_config, None).await {
|
||||
Ok(value) => value,
|
||||
Err(detail) => {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": detail,
|
||||
"replaced": false,
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let replaced = duplicate.is_some();
|
||||
let (persisted_key, key_name) = 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, existing_key.name.clone()),
|
||||
None => {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": "provider oauth write unavailable",
|
||||
"replaced": true,
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
let key_name = auth_config
|
||||
.get("email")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|email| format!("{provider_type}_{email}"))
|
||||
.unwrap_or_else(|| {
|
||||
format!(
|
||||
"{}_{}_{}",
|
||||
provider_type,
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0),
|
||||
index
|
||||
)
|
||||
});
|
||||
match create_provider_oauth_catalog_key(
|
||||
state,
|
||||
provider_id,
|
||||
key_name.as_str(),
|
||||
&access_token,
|
||||
&auth_config,
|
||||
&api_formats,
|
||||
key_proxy.clone(),
|
||||
expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => (key, key_name),
|
||||
None => {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": "provider oauth write unavailable",
|
||||
"replaced": false,
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
let _ =
|
||||
refresh_provider_oauth_account_state_after_update(state, &provider, &persisted_key.id)
|
||||
.await;
|
||||
|
||||
success += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "success",
|
||||
"key_id": persisted_key.id,
|
||||
"key_name": key_name,
|
||||
"error": serde_json::Value::Null,
|
||||
"replaced": replaced,
|
||||
}));
|
||||
}
|
||||
|
||||
Ok(AdminProviderOAuthBatchImportOutcome {
|
||||
total: entries.len(),
|
||||
success,
|
||||
failed,
|
||||
results,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,266 @@
|
||||
use super::parse::{AdminProviderOAuthBatchImportEntry, AdminProviderOAuthBatchImportOutcome};
|
||||
use crate::handlers::admin::provider::oauth::duplicates::find_duplicate_provider_oauth_key;
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::provisioning::{
|
||||
create_provider_oauth_catalog_key, provider_oauth_active_api_formats,
|
||||
provider_oauth_key_proxy_value, update_existing_provider_oauth_catalog_key,
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::runtime::refresh_provider_oauth_account_state_after_update;
|
||||
use crate::handlers::admin::provider::oauth::state::decode_jwt_claims;
|
||||
use crate::handlers::admin::provider::shared::support::ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL;
|
||||
use crate::handlers::admin::request::{
|
||||
AdminAppState, AdminKiroAuthConfig, AdminKiroOAuthRefreshAdapter,
|
||||
};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::oauth::{
|
||||
build_kiro_batch_import_key_name, coerce_admin_provider_oauth_import_str,
|
||||
parse_admin_provider_oauth_kiro_batch_import_entries,
|
||||
};
|
||||
use serde_json::{json, Map, Value};
|
||||
use std::collections::BTreeSet;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
fn admin_provider_oauth_kiro_refresh_base_url_override(
|
||||
state: &AdminAppState<'_>,
|
||||
override_key: &str,
|
||||
) -> Option<String> {
|
||||
let override_url = state.provider_oauth_token_url(override_key, "");
|
||||
let normalized = override_url.trim();
|
||||
(!normalized.is_empty()).then(|| normalized.to_string())
|
||||
}
|
||||
|
||||
pub(super) async fn execute_admin_provider_oauth_kiro_batch_import(
|
||||
state: &AdminAppState<'_>,
|
||||
provider_id: &str,
|
||||
raw_credentials: &str,
|
||||
proxy_node_id: Option<&str>,
|
||||
) -> Result<AdminProviderOAuthBatchImportOutcome, GatewayError> {
|
||||
let entries = parse_admin_provider_oauth_kiro_batch_import_entries(raw_credentials);
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(AdminProviderOAuthBatchImportOutcome {
|
||||
total: entries.len(),
|
||||
success: 0,
|
||||
failed: entries.len(),
|
||||
results: entries
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, _)| {
|
||||
json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": "Provider 不存在",
|
||||
"replaced": false,
|
||||
})
|
||||
})
|
||||
.collect(),
|
||||
});
|
||||
};
|
||||
|
||||
let endpoints = state
|
||||
.list_provider_catalog_endpoints_by_provider_ids(&[provider_id.to_string()])
|
||||
.await?;
|
||||
let key_proxy = provider_oauth_key_proxy_value(proxy_node_id);
|
||||
let adapter = AdminKiroOAuthRefreshAdapter::default().with_refresh_base_urls(
|
||||
admin_provider_oauth_kiro_refresh_base_url_override(state, "kiro_social_refresh"),
|
||||
admin_provider_oauth_kiro_refresh_base_url_override(state, "kiro_idc_refresh"),
|
||||
);
|
||||
let mut results = Vec::with_capacity(entries.len());
|
||||
let mut success = 0usize;
|
||||
let mut failed = 0usize;
|
||||
|
||||
for (index, entry) in entries.iter().enumerate() {
|
||||
let Some(mut refreshed_auth_config) = AdminKiroAuthConfig::from_json_value(entry) else {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": "未找到有效的凭据数据",
|
||||
"replaced": false,
|
||||
}));
|
||||
continue;
|
||||
};
|
||||
|
||||
let has_refresh_token = refreshed_auth_config
|
||||
.refresh_token
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty());
|
||||
if !has_refresh_token {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": "缺少可用的 Kiro refresh 凭据",
|
||||
"replaced": false,
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
|
||||
refreshed_auth_config = match adapter
|
||||
.refresh_auth_config(state.http_client(), &refreshed_auth_config)
|
||||
.await
|
||||
{
|
||||
Ok(config) => config,
|
||||
Err(err) => {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": format!("Token 验证失败: {err:?}"),
|
||||
"replaced": false,
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
if refreshed_auth_config.auth_method.is_none() {
|
||||
refreshed_auth_config.auth_method = Some(if refreshed_auth_config.is_idc_auth() {
|
||||
"idc".to_string()
|
||||
} else {
|
||||
"social".to_string()
|
||||
});
|
||||
}
|
||||
|
||||
let mut auth_config = refreshed_auth_config
|
||||
.to_json_value()
|
||||
.as_object()
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
auth_config.insert("provider_type".to_string(), json!("kiro"));
|
||||
let email = decode_jwt_claims(
|
||||
refreshed_auth_config
|
||||
.access_token
|
||||
.as_deref()
|
||||
.unwrap_or_default(),
|
||||
)
|
||||
.and_then(|claims: Map<String, Value>| claims.get("email").cloned())
|
||||
.and_then(|value: Value| value.as_str().map(ToOwned::to_owned))
|
||||
.or_else(|| coerce_admin_provider_oauth_import_str(entry.get("email")));
|
||||
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(value) => value,
|
||||
Err(detail) => {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": detail,
|
||||
"replaced": false,
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let access_token = refreshed_auth_config
|
||||
.access_token
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
let Some(access_token) = access_token else {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": "Token 验证失败: accessToken 为空",
|
||||
"replaced": false,
|
||||
}));
|
||||
continue;
|
||||
};
|
||||
|
||||
let replaced = duplicate.is_some();
|
||||
let (persisted_key, key_name) = if let Some(existing_key) = duplicate {
|
||||
match update_existing_provider_oauth_catalog_key(
|
||||
state,
|
||||
&existing_key,
|
||||
&access_token,
|
||||
&auth_config,
|
||||
key_proxy.clone(),
|
||||
refreshed_auth_config.expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => (key, existing_key.name.clone()),
|
||||
None => {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": "provider oauth write unavailable",
|
||||
"replaced": true,
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
let key_name = build_kiro_batch_import_key_name(
|
||||
auth_config.get("email").and_then(serde_json::Value::as_str),
|
||||
auth_config
|
||||
.get("auth_method")
|
||||
.and_then(serde_json::Value::as_str),
|
||||
auth_config
|
||||
.get("refresh_token")
|
||||
.and_then(serde_json::Value::as_str),
|
||||
);
|
||||
match create_provider_oauth_catalog_key(
|
||||
state,
|
||||
provider_id,
|
||||
&key_name,
|
||||
&access_token,
|
||||
&auth_config,
|
||||
&provider_oauth_active_api_formats(&endpoints),
|
||||
key_proxy.clone(),
|
||||
refreshed_auth_config.expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => (key, key_name),
|
||||
None => {
|
||||
failed += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": "provider oauth write unavailable",
|
||||
"replaced": false,
|
||||
}));
|
||||
continue;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
let auth_method = auth_config
|
||||
.get("auth_method")
|
||||
.cloned()
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
let _ =
|
||||
refresh_provider_oauth_account_state_after_update(state, &provider, &persisted_key.id)
|
||||
.await;
|
||||
|
||||
success += 1;
|
||||
results.push(json!({
|
||||
"index": index,
|
||||
"status": "success",
|
||||
"key_id": persisted_key.id,
|
||||
"key_name": key_name,
|
||||
"auth_method": auth_method,
|
||||
"error": serde_json::Value::Null,
|
||||
"replaced": replaced,
|
||||
}));
|
||||
}
|
||||
|
||||
Ok(AdminProviderOAuthBatchImportOutcome {
|
||||
total: entries.len(),
|
||||
success,
|
||||
failed,
|
||||
results,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
mod execution;
|
||||
mod kiro_import;
|
||||
mod orchestration;
|
||||
mod parse;
|
||||
mod task;
|
||||
|
||||
pub(super) use orchestration::handle_admin_provider_oauth_batch_import;
|
||||
pub(super) use task::handle_admin_provider_oauth_start_batch_import_task;
|
||||
@@ -0,0 +1,87 @@
|
||||
use super::execution::{
|
||||
estimate_admin_provider_oauth_batch_import_total,
|
||||
execute_admin_provider_oauth_batch_import_for_provider_type,
|
||||
};
|
||||
use super::parse::{
|
||||
build_admin_provider_oauth_batch_import_response,
|
||||
parse_admin_provider_oauth_batch_import_request, AdminProviderOAuthBatchImportRequest,
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
|
||||
is_fixed_provider_type_for_provider_oauth,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_provider_id;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::Bytes,
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
|
||||
pub(in super::super) async fn handle_admin_provider_oauth_batch_import(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Response, GatewayError> {
|
||||
let raw_state = state.cloned_app();
|
||||
if !state.has_provider_catalog_data_reader() {
|
||||
return Ok(build_admin_provider_oauth_backend_unavailable_response());
|
||||
}
|
||||
let Some(provider_id) = admin_provider_oauth_batch_import_provider_id(request_context.path())
|
||||
else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Provider 不存在",
|
||||
));
|
||||
};
|
||||
let payload = match parse_admin_provider_oauth_batch_import_request(request_body) {
|
||||
Ok(payload) => payload,
|
||||
Err(response) => return Ok(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 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" && admin_provider_oauth_template(&provider_type).is_none() {
|
||||
return Ok(build_admin_provider_oauth_backend_unavailable_response());
|
||||
}
|
||||
|
||||
let total = estimate_admin_provider_oauth_batch_import_total(
|
||||
&provider_type,
|
||||
payload.credentials.as_str(),
|
||||
);
|
||||
if total == 0 {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"未找到有效的 Token 数据",
|
||||
));
|
||||
}
|
||||
|
||||
let outcome = execute_admin_provider_oauth_batch_import_for_provider_type(
|
||||
&AdminAppState::new(&raw_state),
|
||||
&provider_id,
|
||||
&provider_type,
|
||||
payload.credentials.as_str(),
|
||||
payload.proxy_node_id.as_deref(),
|
||||
)
|
||||
.await?;
|
||||
Ok(build_admin_provider_oauth_batch_import_response(&outcome).into_response())
|
||||
}
|
||||
@@ -0,0 +1,289 @@
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::state::current_unix_secs;
|
||||
use axum::{
|
||||
body::{to_bytes, Body, Bytes},
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde::Deserialize;
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
#[derive(Debug, Clone, Deserialize)]
|
||||
pub(super) struct AdminProviderOAuthBatchImportRequest {
|
||||
pub credentials: String,
|
||||
pub proxy_node_id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(super) struct AdminProviderOAuthBatchImportEntry {
|
||||
pub refresh_token: String,
|
||||
pub account_id: Option<String>,
|
||||
pub account_user_id: Option<String>,
|
||||
pub plan_type: Option<String>,
|
||||
pub user_id: Option<String>,
|
||||
pub email: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(super) struct AdminProviderOAuthBatchImportOutcome {
|
||||
pub total: usize,
|
||||
pub success: usize,
|
||||
pub failed: usize,
|
||||
pub results: Vec<serde_json::Value>,
|
||||
}
|
||||
|
||||
pub(super) fn parse_admin_provider_oauth_batch_import_request(
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<AdminProviderOAuthBatchImportRequest, Response<Body>> {
|
||||
let Some(request_body) = request_body else {
|
||||
return Err(
|
||||
crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"请求体必须是合法的 JSON 对象",
|
||||
),
|
||||
);
|
||||
};
|
||||
match serde_json::from_slice::<AdminProviderOAuthBatchImportRequest>(request_body) {
|
||||
Ok(payload) if !payload.credentials.trim().is_empty() => Ok(payload),
|
||||
_ => Err(
|
||||
crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"请求体必须是合法的 JSON 对象",
|
||||
),
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
fn coerce_admin_provider_oauth_import_str(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 extract_admin_provider_oauth_batch_import_entry(
|
||||
item: &serde_json::Value,
|
||||
) -> Option<AdminProviderOAuthBatchImportEntry> {
|
||||
match item {
|
||||
serde_json::Value::String(value) => {
|
||||
let refresh_token = value.trim();
|
||||
if refresh_token.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(AdminProviderOAuthBatchImportEntry {
|
||||
refresh_token: refresh_token.to_string(),
|
||||
account_id: None,
|
||||
account_user_id: None,
|
||||
plan_type: None,
|
||||
user_id: None,
|
||||
email: None,
|
||||
})
|
||||
}
|
||||
}
|
||||
serde_json::Value::Object(object) => {
|
||||
let refresh_token = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("refresh_token")
|
||||
.or_else(|| object.get("refreshToken")),
|
||||
)?;
|
||||
let account_id = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("account_id")
|
||||
.or_else(|| object.get("accountId"))
|
||||
.or_else(|| object.get("chatgpt_account_id"))
|
||||
.or_else(|| object.get("chatgptAccountId")),
|
||||
);
|
||||
let account_user_id = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("account_user_id")
|
||||
.or_else(|| object.get("accountUserId"))
|
||||
.or_else(|| object.get("chatgpt_account_user_id"))
|
||||
.or_else(|| object.get("chatgptAccountUserId")),
|
||||
);
|
||||
let plan_type = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("plan_type")
|
||||
.or_else(|| object.get("planType"))
|
||||
.or_else(|| object.get("chatgpt_plan_type"))
|
||||
.or_else(|| object.get("chatgptPlanType")),
|
||||
)
|
||||
.map(|value| value.to_ascii_lowercase());
|
||||
let user_id = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("user_id")
|
||||
.or_else(|| object.get("userId"))
|
||||
.or_else(|| object.get("chatgpt_user_id"))
|
||||
.or_else(|| object.get("chatgptUserId")),
|
||||
);
|
||||
let email = coerce_admin_provider_oauth_import_str(object.get("email"));
|
||||
Some(AdminProviderOAuthBatchImportEntry {
|
||||
refresh_token,
|
||||
account_id,
|
||||
account_user_id,
|
||||
plan_type,
|
||||
user_id,
|
||||
email,
|
||||
})
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn parse_admin_provider_oauth_batch_import_entries(
|
||||
raw_credentials: &str,
|
||||
) -> Vec<AdminProviderOAuthBatchImportEntry> {
|
||||
let raw = raw_credentials.trim();
|
||||
if raw.is_empty() {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
if raw.starts_with('[') {
|
||||
if let Ok(serde_json::Value::Array(items)) = serde_json::from_str::<serde_json::Value>(raw)
|
||||
{
|
||||
return items
|
||||
.iter()
|
||||
.filter_map(extract_admin_provider_oauth_batch_import_entry)
|
||||
.collect();
|
||||
}
|
||||
}
|
||||
|
||||
if raw.starts_with('{') {
|
||||
if let Ok(value @ serde_json::Value::Object(_)) =
|
||||
serde_json::from_str::<serde_json::Value>(raw)
|
||||
{
|
||||
return extract_admin_provider_oauth_batch_import_entry(&value)
|
||||
.into_iter()
|
||||
.collect();
|
||||
}
|
||||
}
|
||||
|
||||
raw.lines()
|
||||
.map(str::trim)
|
||||
.filter(|line| !line.is_empty() && !line.starts_with('#'))
|
||||
.map(|refresh_token| AdminProviderOAuthBatchImportEntry {
|
||||
refresh_token: refresh_token.to_string(),
|
||||
account_id: None,
|
||||
account_user_id: None,
|
||||
plan_type: None,
|
||||
user_id: None,
|
||||
email: None,
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(super) fn apply_admin_provider_oauth_batch_import_hints(
|
||||
provider_type: &str,
|
||||
entry: &AdminProviderOAuthBatchImportEntry,
|
||||
auth_config: &mut serde_json::Map<String, serde_json::Value>,
|
||||
) {
|
||||
if !provider_type.eq_ignore_ascii_case("codex") {
|
||||
return;
|
||||
}
|
||||
if let Some(account_id) = entry.account_id.as_ref() {
|
||||
auth_config
|
||||
.entry("account_id".to_string())
|
||||
.or_insert_with(|| json!(account_id));
|
||||
}
|
||||
if let Some(account_user_id) = entry.account_user_id.as_ref() {
|
||||
auth_config
|
||||
.entry("account_user_id".to_string())
|
||||
.or_insert_with(|| json!(account_user_id));
|
||||
}
|
||||
if let Some(plan_type) = entry.plan_type.as_ref() {
|
||||
auth_config
|
||||
.entry("plan_type".to_string())
|
||||
.or_insert_with(|| json!(plan_type));
|
||||
}
|
||||
if let Some(user_id) = entry.user_id.as_ref() {
|
||||
auth_config
|
||||
.entry("user_id".to_string())
|
||||
.or_insert_with(|| json!(user_id));
|
||||
}
|
||||
if let Some(email) = entry.email.as_ref() {
|
||||
auth_config
|
||||
.entry("email".to_string())
|
||||
.or_insert_with(|| json!(email));
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn extract_admin_provider_oauth_batch_error_detail(
|
||||
response: Response<Body>,
|
||||
) -> String {
|
||||
let status = response.status();
|
||||
let raw_body = to_bytes(response.into_body(), usize::MAX).await.ok();
|
||||
if let Some(raw_body) = raw_body {
|
||||
if let Ok(value) = serde_json::from_slice::<serde_json::Value>(&raw_body) {
|
||||
if let Some(detail) = value.get("detail").and_then(serde_json::Value::as_str) {
|
||||
let normalized = detail.trim();
|
||||
if !normalized.is_empty() {
|
||||
return normalized.to_string();
|
||||
}
|
||||
}
|
||||
}
|
||||
let normalized = String::from_utf8_lossy(&raw_body).trim().to_string();
|
||||
if !normalized.is_empty() {
|
||||
return normalized;
|
||||
}
|
||||
}
|
||||
format!("HTTP {}", status.as_u16())
|
||||
}
|
||||
|
||||
pub(super) fn build_admin_provider_oauth_batch_import_response(
|
||||
outcome: &AdminProviderOAuthBatchImportOutcome,
|
||||
) -> Json<serde_json::Value> {
|
||||
Json(json!({
|
||||
"total": outcome.total,
|
||||
"success": outcome.success,
|
||||
"failed": outcome.failed,
|
||||
"results": outcome.results,
|
||||
}))
|
||||
}
|
||||
|
||||
pub(super) fn build_admin_provider_oauth_batch_task_state(
|
||||
task_id: &str,
|
||||
provider_id: &str,
|
||||
provider_type: &str,
|
||||
status: &str,
|
||||
total: usize,
|
||||
processed: usize,
|
||||
success: usize,
|
||||
failed: usize,
|
||||
message: Option<&str>,
|
||||
error: Option<&str>,
|
||||
error_samples: Vec<serde_json::Value>,
|
||||
created_at: u64,
|
||||
started_at: Option<u64>,
|
||||
finished_at: Option<u64>,
|
||||
) -> serde_json::Value {
|
||||
let updated_at = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(created_at);
|
||||
let progress_percent = if total == 0 {
|
||||
0
|
||||
} else {
|
||||
((processed * 100) / total).min(100) as u64
|
||||
};
|
||||
json!({
|
||||
"task_id": task_id,
|
||||
"provider_id": provider_id,
|
||||
"provider_type": provider_type,
|
||||
"status": status,
|
||||
"total": total,
|
||||
"processed": processed,
|
||||
"success": success,
|
||||
"failed": failed,
|
||||
"progress_percent": progress_percent,
|
||||
"message": message,
|
||||
"error": error,
|
||||
"error_samples": error_samples,
|
||||
"created_at": created_at,
|
||||
"started_at": started_at,
|
||||
"finished_at": finished_at,
|
||||
"updated_at": updated_at,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,249 @@
|
||||
use super::execution::{
|
||||
estimate_admin_provider_oauth_batch_import_total,
|
||||
execute_admin_provider_oauth_batch_import_for_provider_type,
|
||||
};
|
||||
use super::parse::{
|
||||
build_admin_provider_oauth_batch_task_state, parse_admin_provider_oauth_batch_import_request,
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
|
||||
is_fixed_provider_type_for_provider_oauth,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_task_provider_id;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::Bytes,
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use tokio::task;
|
||||
use uuid::Uuid;
|
||||
|
||||
const PROVIDER_OAUTH_BATCH_TASK_MAX_ERROR_SAMPLES: usize = 20;
|
||||
|
||||
pub(in super::super) async fn handle_admin_provider_oauth_start_batch_import_task(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Response, GatewayError> {
|
||||
if !state.has_provider_catalog_data_reader() {
|
||||
return Ok(build_admin_provider_oauth_backend_unavailable_response());
|
||||
}
|
||||
let Some(provider_id) =
|
||||
admin_provider_oauth_batch_import_task_provider_id(request_context.path())
|
||||
else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Provider 不存在",
|
||||
));
|
||||
};
|
||||
let payload = match parse_admin_provider_oauth_batch_import_request(request_body) {
|
||||
Ok(payload) => payload,
|
||||
Err(response) => return Ok(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 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" && admin_provider_oauth_template(&provider_type).is_none() {
|
||||
return Ok(build_admin_provider_oauth_backend_unavailable_response());
|
||||
}
|
||||
|
||||
let total = estimate_admin_provider_oauth_batch_import_total(
|
||||
&provider_type,
|
||||
payload.credentials.as_str(),
|
||||
);
|
||||
if total == 0 {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"未找到有效的 Token 数据",
|
||||
));
|
||||
}
|
||||
|
||||
let task_id = Uuid::new_v4().to_string();
|
||||
let created_at = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
let submitted_state = build_admin_provider_oauth_batch_task_state(
|
||||
&task_id,
|
||||
&provider_id,
|
||||
&provider_type,
|
||||
"submitted",
|
||||
total,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
Some("任务已提交,等待执行"),
|
||||
None,
|
||||
Vec::new(),
|
||||
created_at,
|
||||
None,
|
||||
None,
|
||||
);
|
||||
if state
|
||||
.save_provider_oauth_batch_task_payload(&task_id, &submitted_state)
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"provider oauth batch task redis unavailable",
|
||||
));
|
||||
}
|
||||
|
||||
let task_state = state.cloned_app();
|
||||
let task_id_for_worker = task_id.clone();
|
||||
let provider_id_for_worker = provider_id.clone();
|
||||
let provider_type_for_worker = provider_type.clone();
|
||||
let proxy_node_id = payload.proxy_node_id.clone();
|
||||
let raw_credentials = payload.credentials.clone();
|
||||
task::spawn(async move {
|
||||
let started_at = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(created_at);
|
||||
let processing_state = build_admin_provider_oauth_batch_task_state(
|
||||
&task_id_for_worker,
|
||||
&provider_id_for_worker,
|
||||
&provider_type_for_worker,
|
||||
"processing",
|
||||
total,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
Some("任务开始执行"),
|
||||
None,
|
||||
Vec::new(),
|
||||
created_at,
|
||||
Some(started_at),
|
||||
None,
|
||||
);
|
||||
let _ = AdminAppState::new(&task_state)
|
||||
.save_provider_oauth_batch_task_payload(&task_id_for_worker, &processing_state)
|
||||
.await;
|
||||
|
||||
match execute_admin_provider_oauth_batch_import_for_provider_type(
|
||||
&AdminAppState::new(&task_state),
|
||||
&provider_id_for_worker,
|
||||
&provider_type_for_worker,
|
||||
raw_credentials.as_str(),
|
||||
proxy_node_id.as_deref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(outcome) => {
|
||||
let finished_at = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(started_at);
|
||||
let error_samples = outcome
|
||||
.results
|
||||
.iter()
|
||||
.filter(|item| {
|
||||
item.get("status").and_then(serde_json::Value::as_str) == Some("error")
|
||||
})
|
||||
.take(PROVIDER_OAUTH_BATCH_TASK_MAX_ERROR_SAMPLES)
|
||||
.cloned()
|
||||
.collect::<Vec<_>>();
|
||||
let message = format!(
|
||||
"导入完成:成功 {},失败 {}",
|
||||
outcome.success, outcome.failed
|
||||
);
|
||||
let completed_state = build_admin_provider_oauth_batch_task_state(
|
||||
&task_id_for_worker,
|
||||
&provider_id_for_worker,
|
||||
&provider_type_for_worker,
|
||||
"completed",
|
||||
outcome.total,
|
||||
outcome.total,
|
||||
outcome.success,
|
||||
outcome.failed,
|
||||
Some(message.as_str()),
|
||||
None,
|
||||
error_samples,
|
||||
created_at,
|
||||
Some(started_at),
|
||||
Some(finished_at),
|
||||
);
|
||||
let _ = AdminAppState::new(&task_state)
|
||||
.save_provider_oauth_batch_task_payload(&task_id_for_worker, &completed_state)
|
||||
.await;
|
||||
}
|
||||
Err(err) => {
|
||||
let finished_at = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(started_at);
|
||||
let error_message = format!("{err:?}");
|
||||
let failed_state = build_admin_provider_oauth_batch_task_state(
|
||||
&task_id_for_worker,
|
||||
&provider_id_for_worker,
|
||||
&provider_type_for_worker,
|
||||
"failed",
|
||||
total,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
Some("导入任务执行失败"),
|
||||
Some(error_message.as_str()),
|
||||
Vec::new(),
|
||||
created_at,
|
||||
Some(started_at),
|
||||
Some(finished_at),
|
||||
);
|
||||
let _ = AdminAppState::new(&task_state)
|
||||
.save_provider_oauth_batch_task_payload(&task_id_for_worker, &failed_state)
|
||||
.await;
|
||||
tracing::warn!(
|
||||
task_id = %task_id_for_worker,
|
||||
provider_id = %provider_id_for_worker,
|
||||
error = %error_message,
|
||||
"provider oauth batch import task failed"
|
||||
);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
let submitted_response = build_admin_provider_oauth_batch_task_state(
|
||||
&task_id,
|
||||
&provider_id,
|
||||
&provider_type,
|
||||
"submitted",
|
||||
total,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
Some("任务已提交,等待执行"),
|
||||
None,
|
||||
Vec::new(),
|
||||
created_at,
|
||||
None,
|
||||
None,
|
||||
);
|
||||
Ok(Json(submitted_response).into_response())
|
||||
}
|
||||
@@ -1,562 +0,0 @@
|
||||
use super::super::quota::codex::refresh_codex_provider_quota_locally;
|
||||
use super::super::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::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::paths::{
|
||||
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,257 @@
|
||||
use super::super::super::errors::build_internal_control_error_response;
|
||||
use super::super::super::quota::codex::refresh_codex_provider_quota_locally;
|
||||
use super::super::super::state::{
|
||||
admin_provider_oauth_template, enrich_admin_provider_oauth_auth_config,
|
||||
is_fixed_provider_type_for_provider_oauth, json_non_empty_string, json_u64_value,
|
||||
};
|
||||
use super::shared::{
|
||||
parse_admin_provider_oauth_complete_callback, parse_admin_provider_oauth_complete_request_body,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_complete_key_id;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::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: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let Some(key_id) = admin_provider_oauth_complete_key_id(request_context.path()) else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Key 不存在",
|
||||
));
|
||||
};
|
||||
let payload = match parse_admin_provider_oauth_complete_request_body(request_body) {
|
||||
Ok(payload) => payload,
|
||||
Err(response) => return Ok(response),
|
||||
};
|
||||
let callback = match parse_admin_provider_oauth_complete_callback(&payload.callback_url) {
|
||||
Ok(callback) => callback,
|
||||
Err(response) => return Ok(response),
|
||||
};
|
||||
|
||||
let state_data = match state
|
||||
.consume_provider_oauth_state(&callback.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 state
|
||||
.exchange_admin_provider_oauth_code(
|
||||
template,
|
||||
&callback.code,
|
||||
&callback.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) = state.encrypt_catalog_secret_with_fallbacks(&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) =
|
||||
state.encrypt_catalog_secret_with_fallbacks(&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())
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
mod key;
|
||||
mod provider;
|
||||
mod shared;
|
||||
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
response::Response,
|
||||
};
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_complete_key(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
key::handle_admin_provider_oauth_complete_key(state, request_context, request_body).await
|
||||
}
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_complete_provider(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
provider::handle_admin_provider_oauth_complete_provider(state, request_context, request_body)
|
||||
.await
|
||||
}
|
||||
@@ -0,0 +1,234 @@
|
||||
use super::super::super::duplicates::find_duplicate_provider_oauth_key;
|
||||
use super::super::super::errors::build_internal_control_error_response;
|
||||
use super::super::super::provisioning::{
|
||||
build_provider_oauth_auth_config_from_token_payload, create_provider_oauth_catalog_key,
|
||||
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
|
||||
update_existing_provider_oauth_catalog_key,
|
||||
};
|
||||
use super::super::super::runtime::refresh_provider_oauth_account_state_after_update;
|
||||
use super::super::super::state::{
|
||||
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
|
||||
is_fixed_provider_type_for_provider_oauth,
|
||||
};
|
||||
use super::shared::{
|
||||
parse_admin_provider_oauth_complete_callback, parse_admin_provider_oauth_complete_request_body,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_complete_provider_id;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::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_provider(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let Some(provider_id) = admin_provider_oauth_complete_provider_id(request_context.path())
|
||||
else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Provider 不存在",
|
||||
));
|
||||
};
|
||||
let payload = match parse_admin_provider_oauth_complete_request_body(request_body) {
|
||||
Ok(payload) => payload,
|
||||
Err(response) => return Ok(response),
|
||||
};
|
||||
if !state.has_provider_catalog_data_reader() {
|
||||
return Ok(build_admin_provider_oauth_backend_unavailable_response());
|
||||
}
|
||||
let callback = match parse_admin_provider_oauth_complete_callback(&payload.callback_url) {
|
||||
Ok(callback) => callback,
|
||||
Err(response) => return Ok(response),
|
||||
};
|
||||
|
||||
let state_data = match state
|
||||
.consume_provider_oauth_state(&callback.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 state
|
||||
.exchange_admin_provider_oauth_code(
|
||||
template,
|
||||
&callback.code,
|
||||
&callback.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(payload.proxy_node_id.as_deref());
|
||||
let duplicate = match state
|
||||
.find_duplicate_provider_oauth_key(&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 state
|
||||
.update_existing_provider_oauth_catalog_key(
|
||||
&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 = payload
|
||||
.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 state
|
||||
.create_provider_oauth_catalog_key(
|
||||
&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 _ = state
|
||||
.refresh_provider_oauth_account_state_after_update(&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,116 @@
|
||||
use super::super::super::errors::build_internal_control_error_response;
|
||||
use super::super::super::state::parse_provider_oauth_callback_params;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
response::Response,
|
||||
};
|
||||
|
||||
pub(super) struct AdminProviderOAuthCompleteRequest {
|
||||
pub(super) callback_url: String,
|
||||
pub(super) name: Option<String>,
|
||||
pub(super) proxy_node_id: Option<String>,
|
||||
}
|
||||
|
||||
pub(super) struct AdminProviderOAuthCompleteCallback {
|
||||
pub(super) code: String,
|
||||
pub(super) state_nonce: String,
|
||||
}
|
||||
|
||||
pub(super) fn parse_admin_provider_oauth_callback_url(
|
||||
raw_payload: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Result<String, Response<Body>> {
|
||||
raw_payload
|
||||
.get("callback_url")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.ok_or_else(|| {
|
||||
build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"callback_url 缺少 code/state",
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn extract_admin_provider_oauth_code(
|
||||
params: &std::collections::BTreeMap<String, String>,
|
||||
) -> Result<String, Response<Body>> {
|
||||
params
|
||||
.get("code")
|
||||
.map(String::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.ok_or_else(|| {
|
||||
build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"callback_url 缺少 code/state",
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn extract_admin_provider_oauth_state(
|
||||
params: &std::collections::BTreeMap<String, String>,
|
||||
) -> Result<String, Response<Body>> {
|
||||
params
|
||||
.get("state")
|
||||
.map(String::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.ok_or_else(|| {
|
||||
build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"callback_url 缺少 code/state",
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn parse_admin_provider_oauth_complete_request_body(
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<AdminProviderOAuthCompleteRequest, Response<Body>> {
|
||||
let Some(request_body) = request_body else {
|
||||
return Err(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 Err(build_internal_control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"请求体必须是合法的 JSON 对象",
|
||||
));
|
||||
}
|
||||
};
|
||||
let callback_url = parse_admin_provider_oauth_callback_url(&raw_payload)?;
|
||||
|
||||
Ok(AdminProviderOAuthCompleteRequest {
|
||||
callback_url,
|
||||
name: raw_payload
|
||||
.get("name")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned),
|
||||
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),
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn parse_admin_provider_oauth_complete_callback(
|
||||
callback_url: &str,
|
||||
) -> Result<AdminProviderOAuthCompleteCallback, Response<Body>> {
|
||||
let params = parse_provider_oauth_callback_params(callback_url);
|
||||
let code = extract_admin_provider_oauth_code(¶ms)?;
|
||||
let state_nonce = extract_admin_provider_oauth_state(¶ms)?;
|
||||
|
||||
Ok(AdminProviderOAuthCompleteCallback { code, state_nonce })
|
||||
}
|
||||
@@ -1,550 +0,0 @@
|
||||
use super::super::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::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::paths::{
|
||||
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,194 @@
|
||||
use super::session::AdminProviderOAuthDeviceAuthorizePayload;
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
build_admin_provider_oauth_backend_unavailable_response, current_unix_secs,
|
||||
default_kiro_device_start_url, generate_provider_oauth_nonce, json_non_empty_string,
|
||||
json_u64_value, normalize_kiro_device_region,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_device_authorize_provider_id;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::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_json::json;
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_device_authorize(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
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.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 state
|
||||
.register_admin_kiro_device_oidc_client(®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 state
|
||||
.start_admin_kiro_device_authorization(®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) = state
|
||||
.save_provider_oauth_device_session(
|
||||
&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())
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
mod authorize;
|
||||
mod poll;
|
||||
mod session;
|
||||
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
response::Response,
|
||||
};
|
||||
|
||||
#[cfg(any())]
|
||||
pub(super) use self::authorize::handle_admin_provider_oauth_device_authorize;
|
||||
#[cfg(any())]
|
||||
pub(super) use self::poll::handle_admin_provider_oauth_device_poll;
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_device_authorize(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
authorize::handle_admin_provider_oauth_device_authorize(state, request_context, request_body)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_device_poll(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
poll::handle_admin_provider_oauth_device_poll(state, request_context, request_body).await
|
||||
}
|
||||
@@ -0,0 +1,329 @@
|
||||
use super::session::{
|
||||
attach_admin_provider_oauth_device_poll_terminal_response, AdminProviderOAuthDevicePollPayload,
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::duplicates::find_duplicate_provider_oauth_key;
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::provisioning::{
|
||||
create_provider_oauth_catalog_key, provider_oauth_active_api_formats,
|
||||
provider_oauth_key_proxy_value, update_existing_provider_oauth_catalog_key,
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::runtime::refresh_provider_oauth_account_state_after_update;
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
build_admin_provider_oauth_backend_unavailable_response, build_kiro_device_key_name,
|
||||
current_unix_secs, decode_jwt_claims, json_non_empty_string, json_u64_value,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_device_poll_provider_id;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_device_poll(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
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.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) = state.read_provider_oauth_device_session(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 _ = state
|
||||
.save_provider_oauth_device_session(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 state
|
||||
.poll_admin_kiro_device_token(
|
||||
&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 _ = state
|
||||
.save_provider_oauth_device_session(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 _ = state
|
||||
.save_provider_oauth_device_session(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 state
|
||||
.find_duplicate_provider_oauth_key(&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 = provider_oauth_active_api_formats(
|
||||
&state
|
||||
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider_id))
|
||||
.await?,
|
||||
);
|
||||
let mut replaced = false;
|
||||
let persisted_key = if let Some(existing_key) = duplicate {
|
||||
replaced = true;
|
||||
match state
|
||||
.update_existing_provider_oauth_catalog_key(
|
||||
&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 state
|
||||
.create_provider_oauth_catalog_key(
|
||||
&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 _ = state
|
||||
.refresh_provider_oauth_account_state_after_update(&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 _ = state
|
||||
.save_provider_oauth_device_session(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(),
|
||||
))
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
default_kiro_device_region, default_kiro_device_start_url,
|
||||
};
|
||||
use crate::handlers::admin::shared::attach_admin_audit_response;
|
||||
use axum::{body::Body, response::Response};
|
||||
use serde::Deserialize;
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub(super) struct AdminProviderOAuthDeviceAuthorizePayload {
|
||||
#[serde(default = "default_kiro_device_start_url")]
|
||||
pub(super) start_url: String,
|
||||
#[serde(default = "default_kiro_device_region")]
|
||||
pub(super) region: String,
|
||||
pub(super) proxy_node_id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub(super) struct AdminProviderOAuthDevicePollPayload {
|
||||
pub(super) session_id: String,
|
||||
}
|
||||
|
||||
pub(super) 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,
|
||||
}
|
||||
}
|
||||
@@ -1,16 +1,18 @@
|
||||
use super::super::refresh::{
|
||||
build_internal_control_error_response, build_provider_oauth_auth_config_from_token_payload,
|
||||
create_provider_oauth_catalog_key, find_duplicate_provider_oauth_key,
|
||||
use super::super::duplicates::find_duplicate_provider_oauth_key;
|
||||
use super::super::errors::build_internal_control_error_response;
|
||||
use super::super::provisioning::{
|
||||
build_provider_oauth_auth_config_from_token_payload, create_provider_oauth_catalog_key,
|
||||
provider_oauth_active_api_formats, provider_oauth_key_proxy_value,
|
||||
refresh_provider_oauth_account_state_after_update, update_existing_provider_oauth_catalog_key,
|
||||
update_existing_provider_oauth_catalog_key,
|
||||
};
|
||||
use super::super::runtime::refresh_provider_oauth_account_state_after_update;
|
||||
use super::super::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::paths::admin_provider_oauth_import_provider_id;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
@@ -21,15 +23,14 @@ 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,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
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 {
|
||||
let Some(provider_id) = admin_provider_oauth_import_provider_id(request_context.path()) else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Provider 不存在",
|
||||
@@ -96,13 +97,13 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
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 token_payload = match state
|
||||
.exchange_admin_provider_oauth_refresh_token(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);
|
||||
@@ -124,28 +125,30 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
.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 duplicate = match state
|
||||
.find_duplicate_provider_oauth_key(&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?
|
||||
match state
|
||||
.update_existing_provider_oauth_catalog_key(
|
||||
&existing_key,
|
||||
&access_token,
|
||||
&auth_config,
|
||||
key_proxy.clone(),
|
||||
expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => key,
|
||||
None => {
|
||||
@@ -175,17 +178,17 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
.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?
|
||||
match state
|
||||
.create_provider_oauth_catalog_key(
|
||||
&provider_id,
|
||||
&name,
|
||||
&access_token,
|
||||
&auth_config,
|
||||
&api_formats,
|
||||
key_proxy.clone(),
|
||||
expires_at,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(key) => key,
|
||||
None => {
|
||||
@@ -197,7 +200,8 @@ pub(super) async fn handle_admin_provider_oauth_import_refresh_token(
|
||||
}
|
||||
};
|
||||
|
||||
let _ = refresh_provider_oauth_account_state_after_update(state, &provider, &persisted_key.id)
|
||||
let _ = state
|
||||
.refresh_provider_oauth_account_state_after_update(&provider, &persisted_key.id)
|
||||
.await;
|
||||
|
||||
Ok(Json(json!({
|
||||
|
||||
@@ -2,7 +2,6 @@ use super::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::paths::{
|
||||
admin_provider_oauth_batch_import_provider_id,
|
||||
admin_provider_oauth_batch_import_task_provider_id, admin_provider_oauth_complete_key_id,
|
||||
@@ -10,7 +9,8 @@ use crate::handlers::admin::provider::shared::paths::{
|
||||
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::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::{Body, Bytes},
|
||||
http,
|
||||
@@ -28,11 +28,11 @@ mod start;
|
||||
mod tasks;
|
||||
|
||||
pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.control_decision.as_ref() else {
|
||||
let Some(decision) = request_context.decision() else {
|
||||
return Ok(None);
|
||||
};
|
||||
if decision.route_family.as_deref() != Some("provider_oauth_manage") {
|
||||
@@ -40,11 +40,11 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
}
|
||||
|
||||
let route_kind = decision.route_kind.as_deref();
|
||||
let method = &request_context.request_method;
|
||||
let method = &request_context.method();
|
||||
|
||||
if route_kind == Some("supported_types")
|
||||
&& *method == http::Method::GET
|
||||
&& request_context.request_path == "/api/admin/provider-oauth/supported-types"
|
||||
&& request_context.path() == "/api/admin/provider-oauth/supported-types"
|
||||
{
|
||||
return Ok(Some(
|
||||
Json(build_admin_provider_oauth_supported_types_payload()).into_response(),
|
||||
@@ -58,7 +58,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
"admin_provider_oauth_authorization_started",
|
||||
"start_provider_oauth_for_key",
|
||||
"provider_key",
|
||||
admin_provider_oauth_start_key_id(&request_context.request_path),
|
||||
admin_provider_oauth_start_key_id(request_context.path()),
|
||||
)));
|
||||
}
|
||||
|
||||
@@ -70,7 +70,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
"admin_provider_oauth_authorization_started",
|
||||
"start_provider_oauth_for_provider",
|
||||
"provider",
|
||||
admin_provider_oauth_start_provider_id(&request_context.request_path),
|
||||
admin_provider_oauth_start_provider_id(request_context.path()),
|
||||
)));
|
||||
}
|
||||
|
||||
@@ -93,7 +93,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
"admin_provider_oauth_completed",
|
||||
"complete_provider_oauth_for_key",
|
||||
"provider_key",
|
||||
admin_provider_oauth_complete_key_id(&request_context.request_path),
|
||||
admin_provider_oauth_complete_key_id(request_context.path()),
|
||||
)));
|
||||
}
|
||||
|
||||
@@ -105,7 +105,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
"admin_provider_oauth_refreshed",
|
||||
"refresh_provider_oauth_for_key",
|
||||
"provider_key",
|
||||
admin_provider_oauth_refresh_key_id(&request_context.request_path),
|
||||
admin_provider_oauth_refresh_key_id(request_context.path()),
|
||||
)));
|
||||
}
|
||||
|
||||
@@ -121,7 +121,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
"admin_provider_oauth_completed",
|
||||
"complete_provider_oauth_for_provider",
|
||||
"provider",
|
||||
admin_provider_oauth_complete_provider_id(&request_context.request_path),
|
||||
admin_provider_oauth_complete_provider_id(request_context.path()),
|
||||
)));
|
||||
}
|
||||
|
||||
@@ -137,7 +137,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
"admin_provider_oauth_refresh_token_imported",
|
||||
"import_provider_oauth_refresh_token",
|
||||
"provider",
|
||||
admin_provider_oauth_import_provider_id(&request_context.request_path),
|
||||
admin_provider_oauth_import_provider_id(request_context.path()),
|
||||
)));
|
||||
}
|
||||
|
||||
@@ -150,7 +150,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
"admin_provider_oauth_batch_import_completed",
|
||||
"batch_import_provider_oauth",
|
||||
"provider",
|
||||
admin_provider_oauth_batch_import_provider_id(&request_context.request_path),
|
||||
admin_provider_oauth_batch_import_provider_id(request_context.path()),
|
||||
)));
|
||||
}
|
||||
|
||||
@@ -166,7 +166,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_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),
|
||||
admin_provider_oauth_batch_import_task_provider_id(request_context.path()),
|
||||
)));
|
||||
}
|
||||
|
||||
@@ -182,7 +182,7 @@ pub(crate) async fn maybe_build_local_admin_provider_oauth_response(
|
||||
"admin_provider_oauth_device_authorization_started",
|
||||
"start_provider_oauth_device_authorization",
|
||||
"provider",
|
||||
admin_provider_oauth_device_authorize_provider_id(&request_context.request_path),
|
||||
admin_provider_oauth_device_authorize_provider_id(request_context.path()),
|
||||
)));
|
||||
}
|
||||
|
||||
|
||||
@@ -1,227 +1,28 @@
|
||||
use super::super::quota::shared::persist_provider_quota_refresh_state;
|
||||
use super::super::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::state::is_fixed_provider_type_for_provider_oauth;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_refresh_key_id;
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
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};
|
||||
mod execution;
|
||||
mod helpers;
|
||||
mod request;
|
||||
mod response;
|
||||
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{body::Body, response::Response};
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_refresh_key(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
) -> 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 request =
|
||||
match request::parse_admin_provider_oauth_refresh_request(state, request_context).await? {
|
||||
helpers::RefreshDispatch::Continue(request) => request,
|
||||
helpers::RefreshDispatch::Respond(response) => return Ok(response),
|
||||
};
|
||||
|
||||
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",
|
||||
));
|
||||
let refreshed = match execution::execute_admin_provider_oauth_refresh(state, request).await? {
|
||||
helpers::RefreshDispatch::Continue(refreshed) => refreshed,
|
||||
helpers::RefreshDispatch::Respond(response) => return Ok(response),
|
||||
};
|
||||
|
||||
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())
|
||||
Ok(response::admin_provider_oauth_refresh_success_response(
|
||||
refreshed,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
use super::super::super::errors::{
|
||||
merge_provider_oauth_refresh_failure_reason, normalize_provider_oauth_refresh_error_message,
|
||||
};
|
||||
use super::super::super::quota::shared::persist_provider_quota_refresh_state;
|
||||
use super::super::super::runtime::refresh_provider_oauth_account_state_after_update;
|
||||
use super::helpers::{self, RefreshDispatch, RefreshRequestContext, RefreshSuccessContext};
|
||||
use super::response;
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
|
||||
};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminLocalOAuthRefreshError};
|
||||
use crate::GatewayError;
|
||||
use axum::http;
|
||||
|
||||
pub(super) async fn execute_admin_provider_oauth_refresh(
|
||||
state: &AdminAppState<'_>,
|
||||
request: RefreshRequestContext,
|
||||
) -> Result<RefreshDispatch<RefreshSuccessContext>, GatewayError> {
|
||||
let RefreshRequestContext {
|
||||
key_id,
|
||||
key,
|
||||
provider,
|
||||
provider_type,
|
||||
transport,
|
||||
} = request;
|
||||
|
||||
match state.force_local_oauth_refresh_entry(&transport).await {
|
||||
Ok(Some(_)) => {}
|
||||
Ok(None) => {
|
||||
return Ok(RefreshDispatch::Respond(response::control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"缺少 refresh_token,需要重新授权",
|
||||
)));
|
||||
}
|
||||
Err(AdminLocalOAuthRefreshError::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 failure_reason = format!(
|
||||
"{OAUTH_REFRESH_FAILED_PREFIX}Token 续期失败 ({status_code}): {error_reason}"
|
||||
);
|
||||
let merged_reason = merge_provider_oauth_refresh_failure_reason(
|
||||
key.oauth_invalid_reason.as_deref(),
|
||||
&failure_reason,
|
||||
);
|
||||
if let Some(merged_reason) = merged_reason {
|
||||
let _ = persist_provider_quota_refresh_state(
|
||||
state,
|
||||
&key_id,
|
||||
None,
|
||||
Some(helpers::unix_now_secs()),
|
||||
Some(merged_reason),
|
||||
None,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
response::oauth_refresh_failed_bad_request_response(&error_reason),
|
||||
));
|
||||
}
|
||||
Err(AdminLocalOAuthRefreshError::Transport { source, .. }) => {
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
response::oauth_refresh_failed_service_unavailable_response(source.to_string()),
|
||||
));
|
||||
}
|
||||
Err(AdminLocalOAuthRefreshError::InvalidResponse { message, .. }) => {
|
||||
return Ok(RefreshDispatch::Respond(
|
||||
response::oauth_refresh_failed_bad_request_response(&message),
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
if !helpers::key_is_account_blocked(&key, 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 = helpers::refreshed_auth_config_object(
|
||||
state,
|
||||
refreshed_key.encrypted_auth_config.as_deref(),
|
||||
);
|
||||
let (account_state_recheck_attempted, account_state_recheck_error) = state
|
||||
.refresh_provider_oauth_account_state_after_update(&provider, &key_id)
|
||||
.await?;
|
||||
|
||||
Ok(RefreshDispatch::Continue(RefreshSuccessContext {
|
||||
provider_type,
|
||||
refreshed_auth_config,
|
||||
account_state_recheck_attempted,
|
||||
account_state_recheck_error,
|
||||
}))
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::handlers::admin::shared::decrypt_catalog_secret_with_fallbacks;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use axum::{body::Body, response::Response};
|
||||
use serde_json::{Map, Value};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(super) enum RefreshDispatch<T> {
|
||||
Continue(T),
|
||||
Respond(Response<Body>),
|
||||
}
|
||||
|
||||
pub(super) struct RefreshRequestContext {
|
||||
pub(super) key_id: String,
|
||||
pub(super) key: StoredProviderCatalogKey,
|
||||
pub(super) provider: StoredProviderCatalogProvider,
|
||||
pub(super) provider_type: String,
|
||||
pub(super) transport: AdminGatewayProviderTransportSnapshot,
|
||||
}
|
||||
|
||||
pub(super) struct RefreshSuccessContext {
|
||||
pub(super) provider_type: String,
|
||||
pub(super) refreshed_auth_config: Map<String, Value>,
|
||||
pub(super) account_state_recheck_attempted: bool,
|
||||
pub(super) account_state_recheck_error: Option<String>,
|
||||
}
|
||||
|
||||
pub(super) fn decrypt_auth_config(
|
||||
state: &AdminAppState<'_>,
|
||||
encrypted_auth_config: &str,
|
||||
) -> Option<String> {
|
||||
state.decrypt_catalog_secret_with_fallbacks(encrypted_auth_config)
|
||||
}
|
||||
|
||||
pub(super) fn parse_auth_config_object(plaintext: &str) -> Map<String, Value> {
|
||||
serde_json::from_str::<Value>(plaintext)
|
||||
.ok()
|
||||
.and_then(|value| value.as_object().cloned())
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
pub(super) fn refreshed_auth_config_object(
|
||||
state: &AdminAppState<'_>,
|
||||
encrypted_auth_config: Option<&str>,
|
||||
) -> Map<String, Value> {
|
||||
encrypted_auth_config
|
||||
.and_then(|ciphertext| decrypt_auth_config(state, ciphertext))
|
||||
.map(|plaintext| parse_auth_config_object(&plaintext))
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
pub(super) fn auth_config_has_refresh_token(auth_config: &Map<String, Value>) -> bool {
|
||||
auth_config
|
||||
.get("refresh_token")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
pub(super) fn key_is_account_blocked(key: &StoredProviderCatalogKey, block_prefix: &str) -> bool {
|
||||
key.oauth_invalid_reason
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| value.starts_with(block_prefix))
|
||||
}
|
||||
|
||||
pub(super) fn unix_now_secs() -> u64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0)
|
||||
}
|
||||
@@ -0,0 +1,106 @@
|
||||
use super::super::super::runtime::provider_oauth_runtime_endpoint_for_provider;
|
||||
use super::super::super::state::is_fixed_provider_type_for_provider_oauth;
|
||||
use super::helpers::{self, RefreshDispatch, RefreshRequestContext};
|
||||
use super::response;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_refresh_key_id;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::http;
|
||||
|
||||
pub(super) async fn parse_admin_provider_oauth_refresh_request(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
) -> Result<RefreshDispatch<RefreshRequestContext>, GatewayError> {
|
||||
let Some(key_id) = admin_provider_oauth_refresh_key_id(request_context.path()) else {
|
||||
return Ok(RefreshDispatch::Respond(response::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(RefreshDispatch::Respond(response::control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Key 不存在",
|
||||
)));
|
||||
};
|
||||
if !key.auth_type.eq_ignore_ascii_case("oauth") {
|
||||
return Ok(RefreshDispatch::Respond(response::control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"该 Key 不是 oauth 认证类型",
|
||||
)));
|
||||
}
|
||||
|
||||
let Some(encrypted_auth_config) = key.encrypted_auth_config.as_deref() else {
|
||||
return Ok(RefreshDispatch::Respond(response::control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"缺少 auth_config,无法 refresh",
|
||||
)));
|
||||
};
|
||||
let Some(decrypted_auth_config) = helpers::decrypt_auth_config(state, encrypted_auth_config)
|
||||
else {
|
||||
return Ok(RefreshDispatch::Respond(response::control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
"provider oauth encryption unavailable",
|
||||
)));
|
||||
};
|
||||
let parsed_auth_config = helpers::parse_auth_config_object(&decrypted_auth_config);
|
||||
if !helpers::auth_config_has_refresh_token(&parsed_auth_config) {
|
||||
return Ok(RefreshDispatch::Respond(response::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(RefreshDispatch::Respond(response::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(RefreshDispatch::Respond(response::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(RefreshDispatch::Respond(response::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(RefreshDispatch::Respond(response::control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
"Provider transport snapshot unavailable",
|
||||
)));
|
||||
};
|
||||
|
||||
Ok(RefreshDispatch::Continue(RefreshRequestContext {
|
||||
key_id,
|
||||
key,
|
||||
provider,
|
||||
provider_type,
|
||||
transport,
|
||||
}))
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
use super::super::super::errors::build_internal_control_error_response;
|
||||
use super::helpers::RefreshSuccessContext;
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::{json, Value};
|
||||
|
||||
pub(super) fn control_error_response(
|
||||
status: http::StatusCode,
|
||||
message: impl Into<String>,
|
||||
) -> Response<Body> {
|
||||
build_internal_control_error_response(status, message)
|
||||
}
|
||||
|
||||
pub(super) fn oauth_refresh_failed_bad_request_response(
|
||||
error_reason: impl AsRef<str>,
|
||||
) -> Response<Body> {
|
||||
control_error_response(
|
||||
http::StatusCode::BAD_REQUEST,
|
||||
format!("Token 刷新失败:{}", error_reason.as_ref()),
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) fn oauth_refresh_failed_service_unavailable_response(
|
||||
error_reason: impl Into<String>,
|
||||
) -> Response<Body> {
|
||||
control_error_response(
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
format!("Token 刷新失败:{}", error_reason.into()),
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_oauth_refresh_success_response(
|
||||
success: RefreshSuccessContext,
|
||||
) -> Response<Body> {
|
||||
Json(json!({
|
||||
"provider_type": success.provider_type,
|
||||
"expires_at": success
|
||||
.refreshed_auth_config
|
||||
.get("expires_at")
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null),
|
||||
"has_refresh_token": success
|
||||
.refreshed_auth_config
|
||||
.get("refresh_token")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty()),
|
||||
"email": success
|
||||
.refreshed_auth_config
|
||||
.get("email")
|
||||
.cloned()
|
||||
.unwrap_or(Value::Null),
|
||||
"account_state_recheck_attempted": success.account_state_recheck_attempted,
|
||||
"account_state_recheck_error": success.account_state_recheck_error,
|
||||
}))
|
||||
.into_response()
|
||||
}
|
||||
@@ -1,14 +1,14 @@
|
||||
use super::super::refresh::build_internal_control_error_response;
|
||||
use super::super::errors::build_internal_control_error_response;
|
||||
use super::super::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,
|
||||
provider_oauth_pkce_s256,
|
||||
};
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::{
|
||||
admin_provider_oauth_start_key_id, admin_provider_oauth_start_provider_id,
|
||||
};
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
@@ -17,10 +17,10 @@ use axum::{
|
||||
};
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_start_key(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let Some(key_id) = admin_provider_oauth_start_key_id(&request_context.request_path) else {
|
||||
let Some(key_id) = admin_provider_oauth_start_key_id(request_context.path()) else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Key 不存在",
|
||||
@@ -74,14 +74,14 @@ pub(super) async fn handle_admin_provider_oauth_start_key(
|
||||
.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
|
||||
let nonce = match state
|
||||
.save_provider_oauth_state(
|
||||
&key_id,
|
||||
&provider_id,
|
||||
&provider_type,
|
||||
pkce_verifier.as_deref(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(nonce) => nonce,
|
||||
Err(_) => {
|
||||
@@ -101,11 +101,10 @@ pub(super) async fn handle_admin_provider_oauth_start_key(
|
||||
}
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_start_provider(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let Some(provider_id) = admin_provider_oauth_start_provider_id(&request_context.request_path)
|
||||
else {
|
||||
let Some(provider_id) = admin_provider_oauth_start_provider_id(request_context.path()) else {
|
||||
return Ok(build_internal_control_error_response(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
"Provider 不存在",
|
||||
@@ -146,14 +145,9 @@ pub(super) async fn handle_admin_provider_oauth_start_provider(
|
||||
.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
|
||||
let nonce = match state
|
||||
.save_provider_oauth_state("", &provider_id, &provider_type, pkce_verifier.as_deref())
|
||||
.await
|
||||
{
|
||||
Ok(nonce) => nonce,
|
||||
Err(_) => {
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
use super::super::refresh::build_internal_control_error_response;
|
||||
use super::super::state::read_provider_oauth_batch_task_payload;
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use super::super::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_task_path;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::attach_admin_audit_response;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
@@ -12,18 +11,20 @@ use axum::{
|
||||
};
|
||||
|
||||
pub(super) async fn handle_admin_provider_oauth_batch_import_task_status(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
let Some((provider_id, task_id)) =
|
||||
admin_provider_oauth_batch_import_task_path(&request_context.request_path)
|
||||
admin_provider_oauth_batch_import_task_path(request_context.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
|
||||
let payload = match state
|
||||
.read_provider_oauth_batch_task_payload(&provider_id, &task_id)
|
||||
.await
|
||||
{
|
||||
Ok(Some(payload)) => payload,
|
||||
Ok(None) => {
|
||||
|
||||
@@ -0,0 +1,223 @@
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||
|
||||
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: &AdminAppState<'_>,
|
||||
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) = state.parse_catalog_auth_config_json(&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)
|
||||
}
|
||||
@@ -0,0 +1,145 @@
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX,
|
||||
};
|
||||
use axum::{
|
||||
body::Body,
|
||||
http,
|
||||
response::{IntoResponse, Response},
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
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())
|
||||
}
|
||||
@@ -1,6 +1,9 @@
|
||||
mod dispatch;
|
||||
pub(crate) mod duplicates;
|
||||
pub(crate) mod errors;
|
||||
pub(crate) mod provisioning;
|
||||
pub(crate) mod quota;
|
||||
pub(crate) mod refresh;
|
||||
pub(crate) mod runtime;
|
||||
pub(crate) mod state;
|
||||
|
||||
pub(crate) use self::dispatch::maybe_build_local_admin_provider_oauth_response;
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
use super::state::{
|
||||
enrich_admin_provider_oauth_auth_config, json_non_empty_string, json_u64_value,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::collections::BTreeSet;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use uuid::Uuid;
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
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: &AdminAppState<'_>,
|
||||
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) = state.encrypt_catalog_secret_with_fallbacks(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) =
|
||||
state.encrypt_catalog_secret_with_fallbacks(&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: &AdminAppState<'_>,
|
||||
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) = state.encrypt_catalog_secret_with_fallbacks(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) =
|
||||
state.encrypt_catalog_secret_with_fallbacks(&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
|
||||
}
|
||||
@@ -4,86 +4,25 @@ use super::shared::{
|
||||
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::ANTIGRAVITY_FETCH_AVAILABLE_MODELS_PATH;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::quota::parse_antigravity_usage_response;
|
||||
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};
|
||||
|
||||
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,
|
||||
state: &AdminAppState<'_>,
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
authorization: (String, String),
|
||||
project_id: &str,
|
||||
auth: &crate::provider_transport::antigravity::AntigravityRequestAuthSupport,
|
||||
mut identity_headers: BTreeMap<String, String>,
|
||||
) -> 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,
|
||||
);
|
||||
let mut headers = std::mem::take(&mut identity_headers);
|
||||
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());
|
||||
@@ -117,28 +56,27 @@ async fn execute_antigravity_quota_plan(
|
||||
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 {
|
||||
proxy: state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
|
||||
.await,
|
||||
tls_profile: state.resolve_transport_tls_profile(transport),
|
||||
timeouts: state
|
||||
.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,
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
keys: Vec<StoredProviderCatalogKey>,
|
||||
@@ -165,11 +103,8 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
|
||||
}
|
||||
};
|
||||
|
||||
let authorization = match state.resolve_local_oauth_request_auth(&transport).await? {
|
||||
Some(crate::provider_transport::LocalResolvedOAuthRequestAuth::Header {
|
||||
name,
|
||||
value,
|
||||
}) => (name, value),
|
||||
let authorization = match state.resolve_local_oauth_header_auth(&transport).await? {
|
||||
Some(auth) => auth,
|
||||
_ => {
|
||||
failed_count += 1;
|
||||
results.push(json!({
|
||||
@@ -182,26 +117,17 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
|
||||
}
|
||||
};
|
||||
|
||||
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 Some((project_id, identity_headers)) =
|
||||
state.resolve_local_antigravity_identity_headers(&transport)
|
||||
else {
|
||||
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(
|
||||
@@ -209,7 +135,7 @@ pub(crate) async fn refresh_antigravity_provider_quota_locally(
|
||||
&transport,
|
||||
authorization,
|
||||
&project_id,
|
||||
&antigravity_auth,
|
||||
identity_headers,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
|
||||
@@ -1,696 +0,0 @@
|
||||
use super::shared::{
|
||||
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::payloads::{
|
||||
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,32 @@
|
||||
use aether_admin::provider::quota as admin_provider_quota_pure;
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||
|
||||
pub(super) fn codex_build_invalid_state(
|
||||
key: &StoredProviderCatalogKey,
|
||||
candidate_reason: String,
|
||||
now_unix_secs: u64,
|
||||
) -> (Option<u64>, Option<String>) {
|
||||
admin_provider_quota_pure::codex_build_invalid_state(key, candidate_reason, now_unix_secs)
|
||||
}
|
||||
|
||||
pub(super) fn codex_looks_like_token_invalidated(message: Option<&str>) -> bool {
|
||||
admin_provider_quota_pure::codex_looks_like_token_invalidated(message)
|
||||
}
|
||||
|
||||
pub(super) fn codex_looks_like_workspace_deactivated(message: Option<&str>) -> bool {
|
||||
admin_provider_quota_pure::codex_looks_like_workspace_deactivated(message)
|
||||
}
|
||||
|
||||
pub(super) fn codex_structured_invalid_reason(
|
||||
status_code: u16,
|
||||
upstream_message: Option<&str>,
|
||||
) -> String {
|
||||
admin_provider_quota_pure::codex_structured_invalid_reason(status_code, upstream_message)
|
||||
}
|
||||
|
||||
pub(super) fn codex_soft_request_failure_reason(
|
||||
status_code: u16,
|
||||
upstream_message: Option<&str>,
|
||||
) -> String {
|
||||
admin_provider_quota_pure::codex_soft_request_failure_reason(status_code, upstream_message)
|
||||
}
|
||||
@@ -0,0 +1,292 @@
|
||||
mod invalid;
|
||||
mod parse;
|
||||
mod plan;
|
||||
|
||||
use self::invalid::{
|
||||
codex_build_invalid_state, codex_looks_like_token_invalidated,
|
||||
codex_looks_like_workspace_deactivated, codex_soft_request_failure_reason,
|
||||
codex_structured_invalid_reason,
|
||||
};
|
||||
use self::parse::{
|
||||
build_codex_quota_exhausted_fallback_metadata, parse_codex_usage_headers,
|
||||
parse_codex_wham_usage_response,
|
||||
};
|
||||
use self::plan::{build_codex_refresh_headers, execute_codex_quota_plan};
|
||||
use super::shared::{
|
||||
extract_execution_error_message, 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::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(crate) async fn refresh_codex_provider_quota_locally(
|
||||
state: &AdminAppState<'_>,
|
||||
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") {
|
||||
state.resolve_local_oauth_header_auth(&transport).await?
|
||||
} 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,30 @@
|
||||
use aether_admin::provider::quota as admin_provider_quota_pure;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
pub(super) fn normalize_codex_plan_type(value: Option<&str>) -> Option<String> {
|
||||
admin_provider_quota_pure::normalize_codex_plan_type(value)
|
||||
}
|
||||
|
||||
pub(super) fn build_codex_quota_exhausted_fallback_metadata(
|
||||
plan_type: Option<&str>,
|
||||
updated_at_unix_secs: u64,
|
||||
) -> serde_json::Value {
|
||||
admin_provider_quota_pure::build_codex_quota_exhausted_fallback_metadata(
|
||||
plan_type,
|
||||
updated_at_unix_secs,
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) fn parse_codex_wham_usage_response(
|
||||
value: &serde_json::Value,
|
||||
updated_at_unix_secs: u64,
|
||||
) -> Option<serde_json::Value> {
|
||||
admin_provider_quota_pure::parse_codex_wham_usage_response(value, updated_at_unix_secs)
|
||||
}
|
||||
|
||||
pub(super) fn parse_codex_usage_headers(
|
||||
headers: &BTreeMap<String, String>,
|
||||
updated_at_unix_secs: u64,
|
||||
) -> Option<serde_json::Value> {
|
||||
admin_provider_quota_pure::parse_codex_usage_headers(headers, updated_at_unix_secs)
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
use super::super::shared::{execute_provider_quota_plan, ProviderQuotaExecutionOutcome};
|
||||
use super::parse::normalize_codex_plan_type;
|
||||
use crate::handlers::admin::provider::shared::payloads::CODEX_WHAM_USAGE_URL;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
pub(super) fn build_codex_refresh_headers(
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
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)
|
||||
}
|
||||
|
||||
pub(super) async fn execute_codex_quota_plan(
|
||||
state: &AdminAppState<'_>,
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
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: state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
|
||||
.await,
|
||||
tls_profile: state.resolve_transport_tls_profile(transport),
|
||||
timeouts: state
|
||||
.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
|
||||
}
|
||||
@@ -1,452 +0,0 @@
|
||||
use super::shared::{
|
||||
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::payloads::{
|
||||
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,199 @@
|
||||
mod parse;
|
||||
mod plan;
|
||||
|
||||
use self::parse::parse_kiro_usage_response;
|
||||
use self::plan::execute_kiro_quota_plan;
|
||||
use super::shared::{
|
||||
extract_execution_error_message, persist_provider_quota_refresh_state,
|
||||
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(crate) async fn refresh_kiro_provider_quota_locally(
|
||||
state: &AdminAppState<'_>,
|
||||
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) = state
|
||||
.resolve_local_oauth_kiro_request_auth(&transport)
|
||||
.await?
|
||||
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) =
|
||||
state.encrypt_catalog_secret_with_fallbacks(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,8 @@
|
||||
use aether_admin::provider::quota as admin_provider_quota_pure;
|
||||
|
||||
pub(super) fn parse_kiro_usage_response(
|
||||
value: &serde_json::Value,
|
||||
updated_at_unix_secs: u64,
|
||||
) -> Option<serde_json::Value> {
|
||||
admin_provider_quota_pure::parse_kiro_usage_response(value, updated_at_unix_secs)
|
||||
}
|
||||
@@ -0,0 +1,107 @@
|
||||
use super::super::shared::{execute_provider_quota_plan, ProviderQuotaExecutionOutcome};
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
KIRO_USAGE_LIMITS_PATH, KIRO_USAGE_SDK_VERSION,
|
||||
};
|
||||
use crate::handlers::admin::request::{
|
||||
AdminAppState, AdminGatewayProviderTransportSnapshot, AdminKiroRequestAuth,
|
||||
};
|
||||
use crate::GatewayError;
|
||||
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
|
||||
use std::collections::BTreeMap;
|
||||
use url::form_urlencoded;
|
||||
use uuid::Uuid;
|
||||
|
||||
fn build_kiro_usage_headers(auth: &AdminKiroRequestAuth) -> 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: &AdminKiroRequestAuth) -> 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()
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) async fn execute_kiro_quota_plan(
|
||||
state: &AdminAppState<'_>,
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
auth: &AdminKiroRequestAuth,
|
||||
) -> 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: state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(transport)
|
||||
.await,
|
||||
tls_profile: state.resolve_transport_tls_profile(transport),
|
||||
timeouts: state
|
||||
.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
|
||||
}
|
||||
@@ -1,10 +1,11 @@
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
|
||||
};
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminGatewayProviderTransportSnapshot};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::quota as admin_provider_quota_pure;
|
||||
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;
|
||||
|
||||
@@ -14,59 +15,27 @@ pub(super) enum ProviderQuotaExecutionOutcome {
|
||||
}
|
||||
|
||||
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)
|
||||
admin_provider_quota_pure::provider_auto_remove_banned_keys(config)
|
||||
}
|
||||
|
||||
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))
|
||||
admin_provider_quota_pure::should_auto_remove_structured_reason(reason)
|
||||
}
|
||||
|
||||
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)
|
||||
admin_provider_quota_pure::normalize_string_id_list(values)
|
||||
}
|
||||
|
||||
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,
|
||||
}
|
||||
admin_provider_quota_pure::coerce_json_u64(value)
|
||||
}
|
||||
|
||||
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,
|
||||
}
|
||||
admin_provider_quota_pure::coerce_json_f64(value)
|
||||
}
|
||||
|
||||
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,
|
||||
}
|
||||
admin_provider_quota_pure::coerce_json_bool(value)
|
||||
}
|
||||
|
||||
fn merge_upstream_metadata(
|
||||
@@ -86,65 +55,21 @@ fn merge_upstream_metadata(
|
||||
}
|
||||
|
||||
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())
|
||||
admin_provider_quota_pure::extract_execution_error_message(result)
|
||||
}
|
||||
|
||||
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)
|
||||
admin_provider_quota_pure::quota_refresh_success_invalid_state(key)
|
||||
}
|
||||
|
||||
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)
|
||||
admin_provider_quota_pure::coerce_json_string(value)
|
||||
}
|
||||
|
||||
pub(crate) async fn persist_provider_quota_refresh_state(
|
||||
state: &AppState,
|
||||
state: &AdminAppState<'_>,
|
||||
key_id: &str,
|
||||
metadata_update: Option<&serde_json::Value>,
|
||||
oauth_invalid_at_unix_secs: Option<u64>,
|
||||
@@ -182,12 +107,12 @@ pub(crate) async fn persist_provider_quota_refresh_state(
|
||||
}
|
||||
|
||||
pub(super) async fn execute_provider_quota_plan(
|
||||
state: &AppState,
|
||||
transport: &crate::provider_transport::GatewayProviderTransportSnapshot,
|
||||
state: &AdminAppState<'_>,
|
||||
transport: &AdminGatewayProviderTransportSnapshot,
|
||||
plan: ExecutionPlan,
|
||||
quota_kind: &str,
|
||||
) -> Result<ProviderQuotaExecutionOutcome, GatewayError> {
|
||||
match crate::execution_runtime::execute_execution_runtime_sync_plan(state, None, &plan).await {
|
||||
match state.execute_execution_runtime_sync_plan(None, &plan).await {
|
||||
Ok(result) => Ok(ProviderQuotaExecutionOutcome::Response(result)),
|
||||
Err(err) => {
|
||||
let error = match err {
|
||||
|
||||
@@ -1,647 +0,0 @@
|
||||
use super::quota::antigravity::refresh_antigravity_provider_quota_locally;
|
||||
use super::quota::codex::refresh_codex_provider_quota_locally;
|
||||
use super::quota::kiro::refresh_kiro_provider_quota_locally;
|
||||
use super::quota::shared::persist_provider_quota_refresh_state;
|
||||
use super::state::{
|
||||
enrich_admin_provider_oauth_auth_config, json_non_empty_string, json_u64_value,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::payloads::{
|
||||
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,107 @@
|
||||
use super::quota::antigravity::refresh_antigravity_provider_quota_locally;
|
||||
use super::quota::codex::refresh_codex_provider_quota_locally;
|
||||
use super::quota::kiro::refresh_kiro_provider_quota_locally;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
|
||||
};
|
||||
|
||||
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: &AdminAppState<'_>,
|
||||
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))
|
||||
}
|
||||
@@ -1,996 +0,0 @@
|
||||
use super::refresh::{
|
||||
build_internal_control_error_response, normalize_provider_oauth_refresh_error_message,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::support::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,
|
||||
)))
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
pub(crate) use aether_admin::provider::state::{
|
||||
decode_jwt_claims, enrich_admin_provider_oauth_auth_config, json_non_empty_string,
|
||||
json_u64_value,
|
||||
};
|
||||
|
||||
#[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")));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,185 @@
|
||||
use super::super::errors::{
|
||||
build_internal_control_error_response, normalize_provider_oauth_refresh_error_message,
|
||||
};
|
||||
use super::json_non_empty_string;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminProviderOAuthTemplate};
|
||||
use axum::{body::Body, http, response::Response};
|
||||
|
||||
pub(crate) async fn exchange_admin_provider_oauth_code(
|
||||
state: &AdminAppState<'_>,
|
||||
template: AdminProviderOAuthTemplate,
|
||||
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.http_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: &AdminAppState<'_>,
|
||||
template: AdminProviderOAuthTemplate,
|
||||
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.http_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)
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
mod auth_config;
|
||||
mod exchange;
|
||||
mod storage;
|
||||
mod template;
|
||||
|
||||
pub(crate) use self::auth_config::enrich_admin_provider_oauth_auth_config;
|
||||
pub(crate) use self::exchange::{
|
||||
exchange_admin_provider_oauth_code, exchange_admin_provider_oauth_refresh_token,
|
||||
};
|
||||
pub(crate) use self::storage::build_provider_oauth_start_response;
|
||||
pub(crate) use self::template::{
|
||||
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
|
||||
build_admin_provider_oauth_supported_types_payload, is_fixed_provider_type_for_provider_oauth,
|
||||
};
|
||||
pub(crate) use aether_admin::provider::state::{
|
||||
build_kiro_device_key_name, current_unix_secs, decode_jwt_claims, default_kiro_device_region,
|
||||
default_kiro_device_start_url, generate_provider_oauth_nonce,
|
||||
generate_provider_oauth_pkce_verifier, json_non_empty_string, json_u64_value,
|
||||
normalize_kiro_device_region, parse_provider_oauth_callback_params, provider_oauth_pkce_s256,
|
||||
};
|
||||
@@ -0,0 +1,34 @@
|
||||
use crate::handlers::admin::request::AdminProviderOAuthTemplate;
|
||||
use serde_json::json;
|
||||
use url::form_urlencoded;
|
||||
|
||||
pub(crate) fn build_provider_oauth_start_response(
|
||||
template: AdminProviderOAuthTemplate,
|
||||
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",
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
use super::super::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::shared::support::ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL;
|
||||
use crate::handlers::admin::request::{
|
||||
admin_provider_oauth_template as request_admin_provider_oauth_template,
|
||||
admin_provider_oauth_template_types,
|
||||
is_fixed_provider_type_for_admin_oauth as request_is_fixed_provider_type_for_admin_oauth,
|
||||
AdminProviderOAuthTemplate,
|
||||
};
|
||||
use axum::{body::Body, http, response::Response};
|
||||
use serde_json::json;
|
||||
|
||||
pub(crate) fn is_fixed_provider_type_for_provider_oauth(provider_type: &str) -> bool {
|
||||
request_is_fixed_provider_type_for_admin_oauth(provider_type)
|
||||
}
|
||||
|
||||
pub(crate) fn admin_provider_oauth_template(
|
||||
provider_type: &str,
|
||||
) -> Option<AdminProviderOAuthTemplate> {
|
||||
request_admin_provider_oauth_template(provider_type)
|
||||
}
|
||||
|
||||
pub(crate) fn build_admin_provider_oauth_supported_types_payload() -> Vec<serde_json::Value> {
|
||||
admin_provider_oauth_template_types()
|
||||
.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,
|
||||
)
|
||||
}
|
||||
File diff suppressed because one or more lines are too long
@@ -1,7 +1,7 @@
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::handlers::admin::provider::shared::paths::{
|
||||
admin_provider_ops_architecture_id_from_path, is_admin_provider_ops_architectures_root,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminRequestContext;
|
||||
use crate::GatewayError;
|
||||
use axum::{
|
||||
body::Body,
|
||||
@@ -35,16 +35,16 @@ fn admin_provider_ops_architecture_payload(architecture_id: &str) -> Option<Valu
|
||||
}
|
||||
|
||||
pub(super) async fn maybe_build_local_admin_provider_ops_architectures_response(
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
let Some(decision) = request_context.control_decision.as_ref() else {
|
||||
let Some(decision) = request_context.decision() 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)
|
||||
&& request_context.method() == http::Method::GET
|
||||
&& is_admin_provider_ops_architectures_root(request_context.path())
|
||||
{
|
||||
return Ok(Some(
|
||||
Json(admin_provider_ops_architectures_list_payload()).into_response(),
|
||||
@@ -53,10 +53,10 @@ pub(super) async fn maybe_build_local_admin_provider_ops_architectures_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
|
||||
&& request_context.method() == http::Method::GET
|
||||
{
|
||||
let Some(architecture_id) =
|
||||
admin_provider_ops_architecture_id_from_path(&request_context.request_path)
|
||||
admin_provider_ops_architecture_id_from_path(request_context.path())
|
||||
else {
|
||||
return Ok(Some(
|
||||
(
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use crate::control::GatewayPublicRequestContext;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::GatewayError;
|
||||
use axum::body::{Body, Bytes};
|
||||
use axum::http::Response;
|
||||
|
||||
@@ -7,8 +7,8 @@ mod architectures;
|
||||
pub(crate) mod providers;
|
||||
|
||||
pub(crate) async fn maybe_build_local_admin_provider_ops_response(
|
||||
state: &AppState,
|
||||
request_context: &GatewayPublicRequestContext,
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
request_body: Option<&Bytes>,
|
||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||
if let Some(response) =
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,6 @@
|
||||
mod probe;
|
||||
mod run;
|
||||
mod shared;
|
||||
|
||||
pub(super) use probe::admin_provider_ops_probe_new_api_checkin;
|
||||
pub(super) use run::admin_provider_ops_run_checkin_action;
|
||||
+110
@@ -0,0 +1,110 @@
|
||||
use super::super::super::support::AdminProviderOpsCheckinOutcome;
|
||||
use super::super::support::{admin_provider_ops_json_object_map, admin_provider_ops_request_url};
|
||||
use super::shared::{
|
||||
admin_provider_ops_checkin_already_done, admin_provider_ops_checkin_auth_failure,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use serde_json::json;
|
||||
|
||||
pub(in super::super) async fn admin_provider_ops_probe_new_api_checkin(
|
||||
state: &AdminAppState<'_>,
|
||||
base_url: &str,
|
||||
action_config: &serde_json::Map<String, serde_json::Value>,
|
||||
headers: &reqwest::header::HeaderMap,
|
||||
has_cookie: bool,
|
||||
) -> Option<AdminProviderOpsCheckinOutcome> {
|
||||
let endpoint = action_config
|
||||
.get("checkin_endpoint")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or("/api/user/checkin");
|
||||
let url = admin_provider_ops_request_url(
|
||||
base_url,
|
||||
&admin_provider_ops_json_object_map(json!({ "endpoint": endpoint })),
|
||||
endpoint,
|
||||
);
|
||||
let response = match state
|
||||
.http_client()
|
||||
.request(reqwest::Method::POST, url)
|
||||
.headers(headers.clone())
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
Ok(response) => response,
|
||||
Err(_) => return None,
|
||||
};
|
||||
|
||||
if response.status() == http::StatusCode::NOT_FOUND {
|
||||
return None;
|
||||
}
|
||||
if matches!(
|
||||
response.status(),
|
||||
http::StatusCode::UNAUTHORIZED | http::StatusCode::FORBIDDEN
|
||||
) {
|
||||
return has_cookie.then(|| AdminProviderOpsCheckinOutcome {
|
||||
success: None,
|
||||
message: "Cookie 已失效".to_string(),
|
||||
cookie_expired: true,
|
||||
});
|
||||
}
|
||||
|
||||
let response_json = match response.bytes().await {
|
||||
Ok(bytes) => {
|
||||
serde_json::from_slice::<serde_json::Value>(&bytes).unwrap_or_else(|_| json!({}))
|
||||
}
|
||||
Err(_) => json!({}),
|
||||
};
|
||||
let message = response_json
|
||||
.get("message")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or_default()
|
||||
.trim()
|
||||
.to_string();
|
||||
if response_json
|
||||
.get("success")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
== Some(true)
|
||||
{
|
||||
return Some(AdminProviderOpsCheckinOutcome {
|
||||
success: Some(true),
|
||||
message: if message.is_empty() {
|
||||
"签到成功".to_string()
|
||||
} else {
|
||||
message
|
||||
},
|
||||
cookie_expired: false,
|
||||
});
|
||||
}
|
||||
if admin_provider_ops_checkin_already_done(&message) {
|
||||
return Some(AdminProviderOpsCheckinOutcome {
|
||||
success: None,
|
||||
message: if message.is_empty() {
|
||||
"今日已签到".to_string()
|
||||
} else {
|
||||
message
|
||||
},
|
||||
cookie_expired: false,
|
||||
});
|
||||
}
|
||||
if admin_provider_ops_checkin_auth_failure(&message) {
|
||||
return has_cookie.then(|| AdminProviderOpsCheckinOutcome {
|
||||
success: None,
|
||||
message: if message.is_empty() {
|
||||
"Cookie 已失效".to_string()
|
||||
} else {
|
||||
message
|
||||
},
|
||||
cookie_expired: true,
|
||||
});
|
||||
}
|
||||
Some(AdminProviderOpsCheckinOutcome {
|
||||
success: Some(false),
|
||||
message: if message.is_empty() {
|
||||
"签到失败".to_string()
|
||||
} else {
|
||||
message
|
||||
},
|
||||
cookie_expired: false,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
use super::super::super::support::ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE;
|
||||
use super::super::responses::{
|
||||
admin_provider_ops_action_error, admin_provider_ops_action_not_supported,
|
||||
admin_provider_ops_action_response,
|
||||
};
|
||||
use super::super::support::{admin_provider_ops_request_method, admin_provider_ops_request_url};
|
||||
use super::shared::{
|
||||
admin_provider_ops_checkin_already_done, admin_provider_ops_checkin_auth_failure,
|
||||
admin_provider_ops_checkin_payload,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
|
||||
pub(in super::super) async fn admin_provider_ops_run_checkin_action(
|
||||
state: &AdminAppState<'_>,
|
||||
base_url: &str,
|
||||
architecture_id: &str,
|
||||
action_config: &serde_json::Map<String, serde_json::Value>,
|
||||
headers: &reqwest::header::HeaderMap,
|
||||
has_cookie: bool,
|
||||
) -> serde_json::Value {
|
||||
let start = std::time::Instant::now();
|
||||
if !matches!(architecture_id, "generic_api" | "new_api") {
|
||||
return admin_provider_ops_action_not_supported(
|
||||
"checkin",
|
||||
ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE,
|
||||
);
|
||||
}
|
||||
|
||||
let url = admin_provider_ops_request_url(base_url, action_config, "/api/user/checkin");
|
||||
let method = admin_provider_ops_request_method(action_config, "POST");
|
||||
let response = match state
|
||||
.http_client()
|
||||
.request(method, url)
|
||||
.headers(headers.clone())
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
Ok(response) => response,
|
||||
Err(err) if err.is_timeout() => {
|
||||
return admin_provider_ops_action_error("network_error", "checkin", "请求超时", None);
|
||||
}
|
||||
Err(err) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"checkin",
|
||||
format!("网络错误: {err}"),
|
||||
None,
|
||||
);
|
||||
}
|
||||
};
|
||||
let response_time_ms = Some(start.elapsed().as_millis() as u64);
|
||||
let status = response.status();
|
||||
let response_json = match response.bytes().await {
|
||||
Ok(bytes) => match serde_json::from_slice::<serde_json::Value>(&bytes) {
|
||||
Ok(value) => value,
|
||||
Err(_) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"parse_error",
|
||||
"checkin",
|
||||
"响应不是有效的 JSON",
|
||||
response_time_ms,
|
||||
);
|
||||
}
|
||||
},
|
||||
Err(err) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"checkin",
|
||||
format!("网络错误: {err}"),
|
||||
response_time_ms,
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
if status == http::StatusCode::NOT_FOUND {
|
||||
return admin_provider_ops_action_error(
|
||||
"not_supported",
|
||||
"checkin",
|
||||
"功能未开放",
|
||||
response_time_ms,
|
||||
);
|
||||
}
|
||||
if status == http::StatusCode::TOO_MANY_REQUESTS {
|
||||
return admin_provider_ops_action_error(
|
||||
"rate_limited",
|
||||
"checkin",
|
||||
"请求频率限制",
|
||||
response_time_ms,
|
||||
);
|
||||
}
|
||||
if status == http::StatusCode::UNAUTHORIZED {
|
||||
return admin_provider_ops_action_error(
|
||||
if has_cookie {
|
||||
"auth_expired"
|
||||
} else {
|
||||
"auth_failed"
|
||||
},
|
||||
"checkin",
|
||||
if has_cookie {
|
||||
"Cookie 已失效,请重新配置"
|
||||
} else {
|
||||
"认证失败"
|
||||
},
|
||||
response_time_ms,
|
||||
);
|
||||
}
|
||||
if status == http::StatusCode::FORBIDDEN {
|
||||
return admin_provider_ops_action_error(
|
||||
if has_cookie {
|
||||
"auth_expired"
|
||||
} else {
|
||||
"auth_failed"
|
||||
},
|
||||
"checkin",
|
||||
if has_cookie {
|
||||
"Cookie 已失效或无权限"
|
||||
} else {
|
||||
"无权限访问"
|
||||
},
|
||||
response_time_ms,
|
||||
);
|
||||
}
|
||||
if status != http::StatusCode::OK {
|
||||
return admin_provider_ops_action_error(
|
||||
"unknown_error",
|
||||
"checkin",
|
||||
format!(
|
||||
"HTTP {}: {}",
|
||||
status.as_u16(),
|
||||
status.canonical_reason().unwrap_or("Unknown")
|
||||
),
|
||||
response_time_ms,
|
||||
);
|
||||
}
|
||||
|
||||
let message = response_json
|
||||
.get("message")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or_default()
|
||||
.trim()
|
||||
.to_string();
|
||||
if response_json
|
||||
.get("success")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
== Some(true)
|
||||
{
|
||||
return admin_provider_ops_action_response(
|
||||
"success",
|
||||
"checkin",
|
||||
admin_provider_ops_checkin_payload(&response_json, Some(message)),
|
||||
None,
|
||||
response_time_ms,
|
||||
3600,
|
||||
);
|
||||
}
|
||||
if admin_provider_ops_checkin_already_done(&message) {
|
||||
return admin_provider_ops_action_response(
|
||||
"already_done",
|
||||
"checkin",
|
||||
admin_provider_ops_checkin_payload(&response_json, Some(message)),
|
||||
None,
|
||||
response_time_ms,
|
||||
3600,
|
||||
);
|
||||
}
|
||||
if admin_provider_ops_checkin_auth_failure(&message) {
|
||||
return admin_provider_ops_action_error(
|
||||
if has_cookie {
|
||||
"auth_expired"
|
||||
} else {
|
||||
"auth_failed"
|
||||
},
|
||||
"checkin",
|
||||
if message.is_empty() {
|
||||
if has_cookie {
|
||||
"Cookie 已失效"
|
||||
} else {
|
||||
"认证失败"
|
||||
}
|
||||
} else {
|
||||
message.as_str()
|
||||
},
|
||||
response_time_ms,
|
||||
);
|
||||
}
|
||||
admin_provider_ops_action_error(
|
||||
"unknown_error",
|
||||
"checkin",
|
||||
if message.is_empty() {
|
||||
"签到失败"
|
||||
} else {
|
||||
message.as_str()
|
||||
},
|
||||
response_time_ms,
|
||||
)
|
||||
}
|
||||
+83
@@ -0,0 +1,83 @@
|
||||
use super::super::super::verify::admin_provider_ops_value_as_f64;
|
||||
use super::super::support::admin_provider_ops_checkin_data;
|
||||
|
||||
fn admin_provider_ops_message_contains_any(message: &str, indicators: &[&str]) -> bool {
|
||||
let normalized = message.trim().to_ascii_lowercase();
|
||||
indicators
|
||||
.iter()
|
||||
.any(|indicator| normalized.contains(&indicator.to_ascii_lowercase()))
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_ops_checkin_already_done(message: &str) -> bool {
|
||||
admin_provider_ops_message_contains_any(
|
||||
message,
|
||||
&["already", "已签到", "已经签到", "今日已签", "重复签到"],
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_ops_checkin_auth_failure(message: &str) -> bool {
|
||||
admin_provider_ops_message_contains_any(
|
||||
message,
|
||||
&[
|
||||
"未登录",
|
||||
"请登录",
|
||||
"login",
|
||||
"unauthorized",
|
||||
"无权限",
|
||||
"权限不足",
|
||||
"turnstile",
|
||||
"captcha",
|
||||
"验证码",
|
||||
],
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_ops_checkin_payload(
|
||||
response_json: &serde_json::Value,
|
||||
fallback_message: Option<String>,
|
||||
) -> serde_json::Value {
|
||||
let details = response_json
|
||||
.get("data")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.or_else(|| response_json.as_object());
|
||||
let reward = details.and_then(|value| {
|
||||
admin_provider_ops_value_as_f64(
|
||||
value
|
||||
.get("reward")
|
||||
.or_else(|| value.get("quota"))
|
||||
.or_else(|| value.get("amount")),
|
||||
)
|
||||
});
|
||||
let streak_days = details
|
||||
.and_then(|value| value.get("streak_days").or_else(|| value.get("streak")))
|
||||
.and_then(serde_json::Value::as_i64);
|
||||
let next_reward = details.and_then(|value| {
|
||||
admin_provider_ops_value_as_f64(value.get("next_reward").or_else(|| value.get("next")))
|
||||
});
|
||||
let message = fallback_message.or_else(|| {
|
||||
response_json
|
||||
.get("message")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(ToOwned::to_owned)
|
||||
});
|
||||
let mut extra = serde_json::Map::new();
|
||||
if let Some(details) = details {
|
||||
for (key, value) in details {
|
||||
if matches!(
|
||||
key.as_str(),
|
||||
"reward"
|
||||
| "quota"
|
||||
| "amount"
|
||||
| "streak_days"
|
||||
| "streak"
|
||||
| "next_reward"
|
||||
| "next"
|
||||
| "message"
|
||||
) {
|
||||
continue;
|
||||
}
|
||||
extra.insert(key.clone(), value.clone());
|
||||
}
|
||||
}
|
||||
admin_provider_ops_checkin_data(reward, streak_days, next_reward, message, extra)
|
||||
}
|
||||
@@ -0,0 +1,130 @@
|
||||
mod checkin;
|
||||
mod query_balance;
|
||||
mod responses;
|
||||
mod support;
|
||||
|
||||
use super::config::{
|
||||
admin_provider_ops_config_object, admin_provider_ops_connector_object,
|
||||
admin_provider_ops_decrypted_credentials, resolve_admin_provider_ops_base_url,
|
||||
};
|
||||
use super::support::ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE;
|
||||
use super::verify::admin_provider_ops_verify_headers;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
|
||||
};
|
||||
|
||||
pub(super) fn admin_provider_ops_is_valid_action_type(action_type: &str) -> bool {
|
||||
matches!(
|
||||
action_type,
|
||||
"query_balance"
|
||||
| "checkin"
|
||||
| "claim_quota"
|
||||
| "refresh_token"
|
||||
| "get_usage"
|
||||
| "get_models"
|
||||
| "custom"
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) async fn admin_provider_ops_local_action_response(
|
||||
state: &AdminAppState<'_>,
|
||||
_provider_id: &str,
|
||||
provider: Option<&StoredProviderCatalogProvider>,
|
||||
endpoints: &[StoredProviderCatalogEndpoint],
|
||||
action_type: &str,
|
||||
request_config: Option<&serde_json::Map<String, serde_json::Value>>,
|
||||
) -> serde_json::Value {
|
||||
let Some(provider) = provider else {
|
||||
return responses::admin_provider_ops_action_not_configured(action_type, "未配置操作设置");
|
||||
};
|
||||
let Some(provider_ops_config) = admin_provider_ops_config_object(provider) else {
|
||||
return responses::admin_provider_ops_action_not_configured(action_type, "未配置操作设置");
|
||||
};
|
||||
let architecture_id = provider_ops_config
|
||||
.get("architecture_id")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or("generic_api");
|
||||
let connector_config = admin_provider_ops_connector_object(provider_ops_config)
|
||||
.and_then(|connector| connector.get("config"))
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
if support::admin_provider_ops_should_use_rust_only_action_stub(
|
||||
architecture_id,
|
||||
&connector_config,
|
||||
) {
|
||||
return responses::admin_provider_ops_action_not_supported(
|
||||
action_type,
|
||||
ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE,
|
||||
);
|
||||
}
|
||||
|
||||
let Some(base_url) =
|
||||
resolve_admin_provider_ops_base_url(provider, endpoints, Some(provider_ops_config))
|
||||
else {
|
||||
return responses::admin_provider_ops_action_not_configured(
|
||||
action_type,
|
||||
"Provider 未配置 base_url",
|
||||
);
|
||||
};
|
||||
let 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")),
|
||||
);
|
||||
let headers =
|
||||
match admin_provider_ops_verify_headers(architecture_id, &connector_config, &credentials) {
|
||||
Ok(headers) => headers,
|
||||
Err(message) => {
|
||||
return responses::admin_provider_ops_action_not_configured(action_type, message);
|
||||
}
|
||||
};
|
||||
let Some(action_config) = support::admin_provider_ops_resolved_action_config(
|
||||
architecture_id,
|
||||
provider_ops_config,
|
||||
action_type,
|
||||
request_config,
|
||||
) else {
|
||||
return responses::admin_provider_ops_action_not_supported(
|
||||
action_type,
|
||||
ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE,
|
||||
);
|
||||
};
|
||||
|
||||
match action_type {
|
||||
"query_balance" => {
|
||||
query_balance::admin_provider_ops_run_query_balance_action(
|
||||
state,
|
||||
&base_url,
|
||||
architecture_id,
|
||||
&action_config,
|
||||
&headers,
|
||||
&credentials,
|
||||
)
|
||||
.await
|
||||
}
|
||||
"checkin" => {
|
||||
let has_cookie = credentials
|
||||
.get("cookie")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.is_some_and(|value| !value.trim().is_empty());
|
||||
checkin::admin_provider_ops_run_checkin_action(
|
||||
state,
|
||||
&base_url,
|
||||
architecture_id,
|
||||
&action_config,
|
||||
&headers,
|
||||
has_cookie,
|
||||
)
|
||||
.await
|
||||
}
|
||||
_ => responses::admin_provider_ops_action_not_supported(
|
||||
action_type,
|
||||
ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE,
|
||||
),
|
||||
}
|
||||
}
|
||||
+212
@@ -0,0 +1,212 @@
|
||||
mod parsers;
|
||||
mod yescode;
|
||||
|
||||
use super::super::support::{
|
||||
AdminProviderOpsCheckinOutcome, ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE,
|
||||
};
|
||||
use super::checkin::admin_provider_ops_probe_new_api_checkin;
|
||||
use super::responses::{
|
||||
admin_provider_ops_action_error, admin_provider_ops_action_not_supported,
|
||||
admin_provider_ops_action_response,
|
||||
};
|
||||
use super::support::{
|
||||
admin_provider_ops_is_cookie_auth_architecture, admin_provider_ops_request_method,
|
||||
admin_provider_ops_request_url,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
|
||||
pub(super) async fn admin_provider_ops_run_query_balance_action(
|
||||
state: &AdminAppState<'_>,
|
||||
base_url: &str,
|
||||
architecture_id: &str,
|
||||
action_config: &serde_json::Map<String, serde_json::Value>,
|
||||
headers: &reqwest::header::HeaderMap,
|
||||
credentials: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> serde_json::Value {
|
||||
if architecture_id == "yescode" {
|
||||
return yescode::admin_provider_ops_yescode_balance_payload(
|
||||
state,
|
||||
base_url,
|
||||
headers,
|
||||
action_config,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
let mut balance_checkin = None::<AdminProviderOpsCheckinOutcome>;
|
||||
if matches!(architecture_id, "generic_api" | "new_api") {
|
||||
let has_cookie = credentials
|
||||
.get("cookie")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.is_some_and(|value| !value.trim().is_empty());
|
||||
balance_checkin = admin_provider_ops_probe_new_api_checkin(
|
||||
state,
|
||||
base_url,
|
||||
action_config,
|
||||
headers,
|
||||
has_cookie,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
let start = std::time::Instant::now();
|
||||
let url = admin_provider_ops_request_url(base_url, action_config, "/api/user/balance");
|
||||
let method = admin_provider_ops_request_method(action_config, "GET");
|
||||
let response = match state
|
||||
.http_client()
|
||||
.request(method, url)
|
||||
.headers(headers.clone())
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
Ok(response) => response,
|
||||
Err(err) if err.is_timeout() => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"query_balance",
|
||||
"请求超时",
|
||||
None,
|
||||
);
|
||||
}
|
||||
Err(err) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"query_balance",
|
||||
format!("网络错误: {err}"),
|
||||
None,
|
||||
);
|
||||
}
|
||||
};
|
||||
let response_time_ms = Some(start.elapsed().as_millis() as u64);
|
||||
let status = response.status();
|
||||
let response_json = match response.bytes().await {
|
||||
Ok(bytes) => match serde_json::from_slice::<serde_json::Value>(&bytes) {
|
||||
Ok(value) => value,
|
||||
Err(_) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"parse_error",
|
||||
"query_balance",
|
||||
"响应不是有效的 JSON",
|
||||
response_time_ms,
|
||||
);
|
||||
}
|
||||
},
|
||||
Err(err) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"network_error",
|
||||
"query_balance",
|
||||
format!("网络错误: {err}"),
|
||||
response_time_ms,
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
if status != http::StatusCode::OK {
|
||||
let cookie_auth = admin_provider_ops_is_cookie_auth_architecture(architecture_id);
|
||||
return match status {
|
||||
http::StatusCode::UNAUTHORIZED => admin_provider_ops_action_error(
|
||||
"auth_failed",
|
||||
"query_balance",
|
||||
if cookie_auth {
|
||||
"Cookie 已失效,请重新配置"
|
||||
} else {
|
||||
"认证失败"
|
||||
},
|
||||
response_time_ms,
|
||||
),
|
||||
http::StatusCode::FORBIDDEN => admin_provider_ops_action_error(
|
||||
"auth_failed",
|
||||
"query_balance",
|
||||
if cookie_auth {
|
||||
"Cookie 已失效或无权限"
|
||||
} else {
|
||||
"无权限访问"
|
||||
},
|
||||
response_time_ms,
|
||||
),
|
||||
http::StatusCode::NOT_FOUND => admin_provider_ops_action_error(
|
||||
"not_supported",
|
||||
"query_balance",
|
||||
"功能未开放",
|
||||
response_time_ms,
|
||||
),
|
||||
http::StatusCode::TOO_MANY_REQUESTS => admin_provider_ops_action_error(
|
||||
"rate_limited",
|
||||
"query_balance",
|
||||
"请求频率限制",
|
||||
response_time_ms,
|
||||
),
|
||||
_ => admin_provider_ops_action_error(
|
||||
"unknown_error",
|
||||
"query_balance",
|
||||
format!(
|
||||
"HTTP {}: {}",
|
||||
status.as_u16(),
|
||||
status.canonical_reason().unwrap_or("Unknown")
|
||||
),
|
||||
response_time_ms,
|
||||
),
|
||||
};
|
||||
}
|
||||
|
||||
let data = match architecture_id {
|
||||
"generic_api" | "new_api" => {
|
||||
match parsers::admin_provider_ops_new_api_balance_payload(action_config, &response_json)
|
||||
{
|
||||
Ok(data) => data,
|
||||
Err(message) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"unknown_error",
|
||||
"query_balance",
|
||||
message,
|
||||
response_time_ms,
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
"cubence" => {
|
||||
match parsers::admin_provider_ops_cubence_balance_payload(action_config, &response_json)
|
||||
{
|
||||
Ok(data) => data,
|
||||
Err(message) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"parse_error",
|
||||
"query_balance",
|
||||
message,
|
||||
response_time_ms,
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
"nekocode" => match parsers::admin_provider_ops_nekocode_balance_payload(&response_json) {
|
||||
Ok(data) => data,
|
||||
Err(message) => {
|
||||
return admin_provider_ops_action_error(
|
||||
"parse_error",
|
||||
"query_balance",
|
||||
message,
|
||||
response_time_ms,
|
||||
);
|
||||
}
|
||||
},
|
||||
_ => {
|
||||
return admin_provider_ops_action_not_supported(
|
||||
"query_balance",
|
||||
ADMIN_PROVIDER_OPS_ACTION_RUST_ONLY_MESSAGE,
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
let mut payload = admin_provider_ops_action_response(
|
||||
"success",
|
||||
"query_balance",
|
||||
data,
|
||||
None,
|
||||
response_time_ms,
|
||||
86400,
|
||||
);
|
||||
if let Some(outcome) = balance_checkin.as_ref() {
|
||||
parsers::admin_provider_ops_attach_balance_checkin_outcome(&mut payload, outcome);
|
||||
}
|
||||
payload
|
||||
}
|
||||
+312
@@ -0,0 +1,312 @@
|
||||
use super::super::super::support::AdminProviderOpsCheckinOutcome;
|
||||
use super::super::super::verify::admin_provider_ops_value_as_f64;
|
||||
use super::super::support::{
|
||||
admin_provider_ops_balance_data, admin_provider_ops_parse_rfc3339_unix_secs,
|
||||
admin_provider_ops_quota_divisor,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
pub(super) fn admin_provider_ops_new_api_balance_payload(
|
||||
action_config: &serde_json::Map<String, serde_json::Value>,
|
||||
response_json: &serde_json::Value,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
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 Err(response_json
|
||||
.get("message")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("业务状态码表示失败")
|
||||
.to_string());
|
||||
} else {
|
||||
Some(response_json)
|
||||
};
|
||||
let Some(user_data) = user_data.and_then(serde_json::Value::as_object) else {
|
||||
return Err("响应格式无效".to_string());
|
||||
};
|
||||
let quota_divisor = admin_provider_ops_quota_divisor(action_config);
|
||||
let total_available =
|
||||
admin_provider_ops_value_as_f64(user_data.get("quota")).map(|value| value / quota_divisor);
|
||||
let total_used = admin_provider_ops_value_as_f64(user_data.get("used_quota"))
|
||||
.map(|value| value / quota_divisor);
|
||||
Ok(admin_provider_ops_balance_data(
|
||||
None,
|
||||
total_used,
|
||||
total_available,
|
||||
action_config
|
||||
.get("currency")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("USD"),
|
||||
serde_json::Map::new(),
|
||||
))
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_ops_cubence_balance_payload(
|
||||
action_config: &serde_json::Map<String, serde_json::Value>,
|
||||
response_json: &serde_json::Value,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let response_data = response_json
|
||||
.get("data")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.ok_or_else(|| "响应格式无效".to_string())?;
|
||||
let balance_data = response_data
|
||||
.get("balance")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
let subscription_limits = response_data
|
||||
.get("subscription_limits")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
let mut extra = serde_json::Map::new();
|
||||
if let Some(five_hour) = subscription_limits
|
||||
.get("five_hour")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
{
|
||||
extra.insert(
|
||||
"five_hour_limit".to_string(),
|
||||
json!({
|
||||
"limit": five_hour.get("limit"),
|
||||
"used": five_hour.get("used"),
|
||||
"remaining": five_hour.get("remaining"),
|
||||
"resets_at": five_hour.get("resets_at"),
|
||||
}),
|
||||
);
|
||||
}
|
||||
if let Some(weekly) = subscription_limits
|
||||
.get("weekly")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
{
|
||||
extra.insert(
|
||||
"weekly_limit".to_string(),
|
||||
json!({
|
||||
"limit": weekly.get("limit"),
|
||||
"used": weekly.get("used"),
|
||||
"remaining": weekly.get("remaining"),
|
||||
"resets_at": weekly.get("resets_at"),
|
||||
}),
|
||||
);
|
||||
}
|
||||
for key in [
|
||||
"normal_balance_dollar",
|
||||
"subscription_balance_dollar",
|
||||
"charity_balance_dollar",
|
||||
] {
|
||||
if let Some(value) = balance_data.get(key) {
|
||||
extra.insert(
|
||||
key.trim_end_matches("_dollar").replace("_dollar", ""),
|
||||
value.clone(),
|
||||
);
|
||||
}
|
||||
}
|
||||
if let Some(value) = balance_data.get("normal_balance_dollar") {
|
||||
extra.insert("normal_balance".to_string(), value.clone());
|
||||
}
|
||||
if let Some(value) = balance_data.get("subscription_balance_dollar") {
|
||||
extra.insert("subscription_balance".to_string(), value.clone());
|
||||
}
|
||||
if let Some(value) = balance_data.get("charity_balance_dollar") {
|
||||
extra.insert("charity_balance".to_string(), value.clone());
|
||||
}
|
||||
Ok(admin_provider_ops_balance_data(
|
||||
None,
|
||||
None,
|
||||
admin_provider_ops_value_as_f64(balance_data.get("total_balance_dollar")),
|
||||
action_config
|
||||
.get("currency")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("USD"),
|
||||
extra,
|
||||
))
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_ops_nekocode_balance_payload(
|
||||
response_json: &serde_json::Value,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let response_data = response_json
|
||||
.get("data")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.ok_or_else(|| "响应格式无效".to_string())?;
|
||||
let subscription = response_data
|
||||
.get("subscription")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
let balance = admin_provider_ops_value_as_f64(response_data.get("balance"));
|
||||
let daily_quota_limit = admin_provider_ops_value_as_f64(subscription.get("daily_quota_limit"));
|
||||
let daily_remaining_quota =
|
||||
admin_provider_ops_value_as_f64(subscription.get("daily_remaining_quota"));
|
||||
let daily_used = match (daily_quota_limit, daily_remaining_quota) {
|
||||
(Some(limit), Some(remaining)) => Some(limit - remaining),
|
||||
_ => None,
|
||||
};
|
||||
let mut extra = serde_json::Map::new();
|
||||
for key in [
|
||||
"plan_name",
|
||||
"status",
|
||||
"daily_quota_limit",
|
||||
"daily_remaining_quota",
|
||||
"effective_start_date",
|
||||
"effective_end_date",
|
||||
] {
|
||||
if let Some(value) = subscription.get(key) {
|
||||
extra.insert(
|
||||
match key {
|
||||
"status" => "subscription_status",
|
||||
other => other,
|
||||
}
|
||||
.to_string(),
|
||||
value.clone(),
|
||||
);
|
||||
}
|
||||
}
|
||||
if let Some(value) = daily_used {
|
||||
extra.insert("daily_used_quota".to_string(), json!(value));
|
||||
}
|
||||
if let Some(month_data) = response_data
|
||||
.get("month")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
{
|
||||
extra.insert(
|
||||
"month_stats".to_string(),
|
||||
json!({
|
||||
"total_input_tokens": month_data.get("total_input_tokens"),
|
||||
"total_output_tokens": month_data.get("total_output_tokens"),
|
||||
"total_quota": month_data.get("total_quota"),
|
||||
"total_requests": month_data.get("total_requests"),
|
||||
}),
|
||||
);
|
||||
}
|
||||
if let Some(today_data) = response_data
|
||||
.get("today")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
{
|
||||
if let Some(stats) = today_data.get("stats") {
|
||||
extra.insert("today_stats".to_string(), stats.clone());
|
||||
}
|
||||
}
|
||||
Ok(admin_provider_ops_balance_data(
|
||||
daily_quota_limit,
|
||||
daily_used,
|
||||
balance,
|
||||
"USD",
|
||||
extra,
|
||||
))
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_ops_yescode_balance_extra(
|
||||
combined_data: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> serde_json::Map<String, serde_json::Value> {
|
||||
let pay_as_you_go =
|
||||
admin_provider_ops_value_as_f64(combined_data.get("pay_as_you_go_balance")).unwrap_or(0.0);
|
||||
let subscription =
|
||||
admin_provider_ops_value_as_f64(combined_data.get("subscription_balance")).unwrap_or(0.0);
|
||||
let plan = combined_data
|
||||
.get("subscription_plan")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.cloned()
|
||||
.unwrap_or_default();
|
||||
let daily_balance =
|
||||
admin_provider_ops_value_as_f64(plan.get("daily_balance")).unwrap_or(subscription);
|
||||
let weekly_limit = admin_provider_ops_value_as_f64(
|
||||
combined_data
|
||||
.get("weekly_limit")
|
||||
.or_else(|| plan.get("weekly_limit")),
|
||||
);
|
||||
let weekly_spent =
|
||||
admin_provider_ops_value_as_f64(combined_data.get("weekly_spent_balance")).unwrap_or(0.0);
|
||||
let subscription_available = weekly_limit
|
||||
.map(|limit| (limit - weekly_spent).max(0.0).min(subscription))
|
||||
.unwrap_or(subscription);
|
||||
|
||||
let mut extra = serde_json::Map::new();
|
||||
extra.insert("pay_as_you_go_balance".to_string(), json!(pay_as_you_go));
|
||||
extra.insert("daily_limit".to_string(), json!(daily_balance));
|
||||
if let Some(limit) = weekly_limit {
|
||||
extra.insert("weekly_limit".to_string(), json!(limit));
|
||||
}
|
||||
extra.insert("weekly_spent".to_string(), json!(weekly_spent));
|
||||
if let Some(last_week_reset) =
|
||||
admin_provider_ops_parse_rfc3339_unix_secs(combined_data.get("last_week_reset"))
|
||||
{
|
||||
extra.insert(
|
||||
"weekly_resets_at".to_string(),
|
||||
json!(last_week_reset + 7 * 24 * 3600),
|
||||
);
|
||||
}
|
||||
if let Some(last_daily_add) =
|
||||
admin_provider_ops_parse_rfc3339_unix_secs(combined_data.get("last_daily_balance_add"))
|
||||
{
|
||||
extra.insert(
|
||||
"daily_resets_at".to_string(),
|
||||
json!(last_daily_add + 24 * 3600),
|
||||
);
|
||||
}
|
||||
let daily_spent = if let Some(limit) = weekly_limit {
|
||||
daily_balance - daily_balance.min(subscription_available.min(limit.max(0.0)))
|
||||
} else {
|
||||
(daily_balance - subscription).max(0.0)
|
||||
};
|
||||
extra.insert("daily_spent".to_string(), json!(daily_spent));
|
||||
extra.insert(
|
||||
"_subscription_available".to_string(),
|
||||
json!(subscription_available),
|
||||
);
|
||||
extra.insert(
|
||||
"_total_available".to_string(),
|
||||
json!(pay_as_you_go + subscription_available),
|
||||
);
|
||||
extra
|
||||
}
|
||||
|
||||
pub(super) fn admin_provider_ops_attach_balance_checkin_outcome(
|
||||
action_payload: &mut serde_json::Value,
|
||||
outcome: &AdminProviderOpsCheckinOutcome,
|
||||
) {
|
||||
if let Some(data) = action_payload
|
||||
.get_mut("data")
|
||||
.and_then(serde_json::Value::as_object_mut)
|
||||
{
|
||||
let extra = data
|
||||
.entry("extra".to_string())
|
||||
.or_insert_with(|| serde_json::Value::Object(serde_json::Map::new()));
|
||||
if let Some(extra) = extra.as_object_mut() {
|
||||
if outcome.cookie_expired {
|
||||
extra.insert("cookie_expired".to_string(), serde_json::Value::Bool(true));
|
||||
extra.insert(
|
||||
"cookie_expired_message".to_string(),
|
||||
serde_json::Value::String(outcome.message.clone()),
|
||||
);
|
||||
} else {
|
||||
extra.insert(
|
||||
"checkin_success".to_string(),
|
||||
outcome
|
||||
.success
|
||||
.map(serde_json::Value::Bool)
|
||||
.unwrap_or(serde_json::Value::Null),
|
||||
);
|
||||
extra.insert(
|
||||
"checkin_message".to_string(),
|
||||
serde_json::Value::String(outcome.message.clone()),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
if outcome.cookie_expired {
|
||||
if let Some(object) = action_payload.as_object_mut() {
|
||||
object.insert("status".to_string(), json!("auth_expired"));
|
||||
}
|
||||
}
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user