mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 18:37:46 +08:00
refactor: 大规模模块拆分与代码精简,新增 ai-pipeline/data-contracts 独立 crate
- 新增 aether-ai-pipeline 和 aether-data-contracts crate,将 pipeline 逻辑与数据契约从 gateway 中解耦 - 重构 admin handlers:拆分单体模块为 auth/billing/endpoint/features/model/observability/provider/system 等独立子模块 - 合并 chat/cli 重复代码路径:精简 conversion、finalize、planner 中的 sync/chat/cli 分支 - 重构 scheduler/executor/data 层,引入 facade 模式降低模块间耦合 - 移除冗余的 intent 模块,将 plan_fallback/policy/stream_path/sync_path 迁移至 executor - 前端适配:调整 admin API 调用和 provider 模型测试对话框
This commit is contained in:
@@ -0,0 +1,495 @@
|
||||
use super::{
|
||||
normalize_auth_type, normalize_json_object, normalize_string_list, validate_vertex_api_formats,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::{
|
||||
AdminProviderKeyCreateRequest, AdminProviderKeyUpdateRequest,
|
||||
};
|
||||
use crate::handlers::admin::shared::{
|
||||
build_admin_provider_key_response, decrypt_catalog_secret_with_fallbacks,
|
||||
encrypt_catalog_secret_with_fallbacks, json_string_list, parse_catalog_auth_config_json,
|
||||
};
|
||||
use crate::AppState;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use uuid::Uuid;
|
||||
|
||||
pub(crate) async fn build_admin_create_provider_key_record(
|
||||
state: &AppState,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
payload: AdminProviderKeyCreateRequest,
|
||||
) -> Result<StoredProviderCatalogKey, String> {
|
||||
let name = payload.name.trim();
|
||||
if name.is_empty() {
|
||||
return Err("name 为必填字段".to_string());
|
||||
}
|
||||
|
||||
let api_formats = normalize_string_list(payload.api_formats)
|
||||
.ok_or_else(|| "api_formats 为必填字段".to_string())?;
|
||||
let auth_type = normalize_auth_type(payload.auth_type.as_deref())?;
|
||||
validate_vertex_api_formats(&provider.provider_type, &auth_type, &api_formats)?;
|
||||
|
||||
let api_key = payload.api_key.unwrap_or_default().trim().to_string();
|
||||
let auth_config = normalize_json_object(payload.auth_config, "auth_config")?;
|
||||
let auth_config_object = auth_config
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.cloned();
|
||||
|
||||
match auth_type.as_str() {
|
||||
"api_key" => {
|
||||
if api_key.is_empty() {
|
||||
return Err("API Key 认证模式下 api_key 为必填字段".to_string());
|
||||
}
|
||||
}
|
||||
"service_account" => {
|
||||
if auth_config_object.is_none() {
|
||||
return Err("Service Account 认证模式下 auth_config 为必填字段".to_string());
|
||||
}
|
||||
}
|
||||
"oauth" => {
|
||||
if !api_key.is_empty() {
|
||||
return Err("OAuth 认证模式下不允许直接填写 api_key".to_string());
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
let existing_keys = state
|
||||
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?;
|
||||
|
||||
if auth_type == "api_key" {
|
||||
for existing in existing_keys
|
||||
.iter()
|
||||
.filter(|existing| existing.auth_type.trim().eq_ignore_ascii_case("api_key"))
|
||||
{
|
||||
let Some(decrypted) = decrypt_catalog_secret_with_fallbacks(
|
||||
state.encryption_key(),
|
||||
&existing.encrypted_api_key,
|
||||
) else {
|
||||
continue;
|
||||
};
|
||||
if decrypted != "__placeholder__" && decrypted == api_key {
|
||||
return Err(format!(
|
||||
"该 API Key 已存在于当前 Provider 中(名称: {})",
|
||||
existing.name
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if auth_type == "service_account" {
|
||||
let new_client_email = auth_config_object
|
||||
.as_ref()
|
||||
.and_then(|config| config.get("client_email"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
if let Some(new_client_email) = new_client_email {
|
||||
for existing in existing_keys.iter().filter(|existing| {
|
||||
matches!(
|
||||
existing.auth_type.trim().to_ascii_lowercase().as_str(),
|
||||
"service_account" | "vertex_ai"
|
||||
)
|
||||
}) {
|
||||
let Some(existing_config) = parse_catalog_auth_config_json(state, existing) else {
|
||||
continue;
|
||||
};
|
||||
let Some(existing_email) = existing_config
|
||||
.get("client_email")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
if existing_email == new_client_email {
|
||||
return Err(format!(
|
||||
"该 Service Account ({new_client_email}) 已存在于当前 Provider 中(名称: {})",
|
||||
existing.name
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let encrypted_api_key = match auth_type.as_str() {
|
||||
"api_key" => encrypt_catalog_secret_with_fallbacks(state, &api_key),
|
||||
_ => encrypt_catalog_secret_with_fallbacks(state, "__placeholder__"),
|
||||
}
|
||||
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?;
|
||||
|
||||
let encrypted_auth_config = auth_config
|
||||
.as_ref()
|
||||
.map(serde_json::to_string)
|
||||
.transpose()
|
||||
.map_err(|err| err.to_string())?
|
||||
.and_then(|plaintext| encrypt_catalog_secret_with_fallbacks(state, &plaintext));
|
||||
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
let mut key = StoredProviderCatalogKey::new(
|
||||
Uuid::new_v4().to_string(),
|
||||
provider.id.clone(),
|
||||
name.to_string(),
|
||||
auth_type,
|
||||
normalize_json_object(payload.capabilities, "capabilities")?,
|
||||
true,
|
||||
)
|
||||
.map_err(|err| err.to_string())?
|
||||
.with_transport_fields(
|
||||
Some(json!(api_formats)),
|
||||
encrypted_api_key,
|
||||
encrypted_auth_config,
|
||||
normalize_json_object(payload.rate_multipliers, "rate_multipliers")?,
|
||||
None,
|
||||
normalize_string_list(payload.allowed_models).map(|value| json!(value)),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.map_err(|err| err.to_string())?;
|
||||
key.note = payload
|
||||
.note
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty());
|
||||
key.internal_priority = payload.internal_priority.unwrap_or(50);
|
||||
key.rpm_limit = payload.rpm_limit;
|
||||
key.cache_ttl_minutes = payload.cache_ttl_minutes.unwrap_or(5);
|
||||
key.max_probe_interval_minutes = payload.max_probe_interval_minutes.unwrap_or(32);
|
||||
key.request_count = Some(0);
|
||||
key.success_count = Some(0);
|
||||
key.error_count = Some(0);
|
||||
key.total_response_time_ms = Some(0);
|
||||
key.auto_fetch_models = payload.auto_fetch_models.unwrap_or(false);
|
||||
key.locked_models = normalize_string_list(payload.locked_models).map(|value| json!(value));
|
||||
key.model_include_patterns =
|
||||
normalize_string_list(payload.model_include_patterns).map(|value| json!(value));
|
||||
key.model_exclude_patterns =
|
||||
normalize_string_list(payload.model_exclude_patterns).map(|value| json!(value));
|
||||
key.health_by_format = Some(json!({}));
|
||||
key.circuit_breaker_by_format = Some(json!({}));
|
||||
key.created_at_unix_secs = Some(now_unix_secs);
|
||||
key.updated_at_unix_secs = Some(now_unix_secs);
|
||||
Ok(key)
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_update_provider_key_record(
|
||||
state: &AppState,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
existing: &StoredProviderCatalogKey,
|
||||
raw_payload: &serde_json::Map<String, serde_json::Value>,
|
||||
payload: AdminProviderKeyUpdateRequest,
|
||||
) -> Result<StoredProviderCatalogKey, String> {
|
||||
let mut updated = existing.clone();
|
||||
let current_auth_type = normalize_auth_type(Some(&existing.auth_type))?;
|
||||
let target_auth_type = payload
|
||||
.auth_type
|
||||
.as_deref()
|
||||
.map(|value| normalize_auth_type(Some(value)))
|
||||
.transpose()?
|
||||
.unwrap_or_else(|| current_auth_type.clone());
|
||||
let auth_type_switch = payload
|
||||
.auth_type
|
||||
.as_deref()
|
||||
.is_some_and(|_| target_auth_type != current_auth_type);
|
||||
|
||||
let api_key_present = raw_payload.contains_key("api_key");
|
||||
let api_key_value = payload
|
||||
.api_key
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.map(ToOwned::to_owned);
|
||||
if api_key_present && api_key_value.as_deref() == Some("") {
|
||||
return Err("api_key 不能为空".to_string());
|
||||
}
|
||||
|
||||
let auth_config_present = raw_payload.contains_key("auth_config");
|
||||
let auth_config = normalize_json_object(payload.auth_config, "auth_config")?;
|
||||
let auth_config_object = auth_config
|
||||
.as_ref()
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.cloned();
|
||||
|
||||
let existing_keys = state
|
||||
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?;
|
||||
|
||||
match target_auth_type.as_str() {
|
||||
"api_key" => {
|
||||
if auth_type_switch
|
||||
&& matches!(
|
||||
api_key_value.as_deref(),
|
||||
None | Some("") | Some("__placeholder__")
|
||||
)
|
||||
{
|
||||
return Err("切换到 API Key 认证模式时,必须提供新的 API Key".to_string());
|
||||
}
|
||||
if api_key_present
|
||||
&& matches!(
|
||||
api_key_value.as_deref(),
|
||||
None | Some("") | Some("__placeholder__")
|
||||
)
|
||||
{
|
||||
return Err("API Key 认证模式下 api_key 不能为空".to_string());
|
||||
}
|
||||
if let Some(api_key) = api_key_value.as_deref() {
|
||||
for existing_key in existing_keys.iter().filter(|key| {
|
||||
key.id != existing.id && key.auth_type.trim().eq_ignore_ascii_case("api_key")
|
||||
}) {
|
||||
let Some(decrypted) = decrypt_catalog_secret_with_fallbacks(
|
||||
state.encryption_key(),
|
||||
&existing_key.encrypted_api_key,
|
||||
) else {
|
||||
continue;
|
||||
};
|
||||
if decrypted != "__placeholder__" && decrypted == api_key {
|
||||
return Err(format!(
|
||||
"该 API Key 已存在于当前 Provider 中(名称: {})",
|
||||
existing_key.name
|
||||
));
|
||||
}
|
||||
}
|
||||
updated.encrypted_api_key =
|
||||
encrypt_catalog_secret_with_fallbacks(state, api_key)
|
||||
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?;
|
||||
}
|
||||
updated.encrypted_auth_config = None;
|
||||
}
|
||||
"service_account" => {
|
||||
if auth_type_switch && auth_config_object.is_none() {
|
||||
return Err(
|
||||
"切换到 Service Account 认证模式时,必须提供 Service Account JSON".to_string(),
|
||||
);
|
||||
}
|
||||
if api_key_present
|
||||
&& !matches!(
|
||||
api_key_value.as_deref(),
|
||||
None | Some("") | Some("__placeholder__")
|
||||
)
|
||||
{
|
||||
return Err("Service Account 认证模式下不允许直接填写 api_key".to_string());
|
||||
}
|
||||
if auth_type_switch || api_key_present {
|
||||
updated.encrypted_api_key =
|
||||
encrypt_catalog_secret_with_fallbacks(state, "__placeholder__")
|
||||
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?;
|
||||
}
|
||||
if let Some(client_email) = auth_config_object
|
||||
.as_ref()
|
||||
.and_then(|config| config.get("client_email"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
for existing_key in existing_keys.iter().filter(|key| {
|
||||
key.id != existing.id
|
||||
&& matches!(
|
||||
key.auth_type.trim().to_ascii_lowercase().as_str(),
|
||||
"service_account" | "vertex_ai"
|
||||
)
|
||||
}) {
|
||||
let Some(existing_config) = parse_catalog_auth_config_json(state, existing_key)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
let Some(existing_email) = existing_config
|
||||
.get("client_email")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
if existing_email == client_email {
|
||||
return Err(format!(
|
||||
"该 Service Account ({client_email}) 已存在于当前 Provider 中(名称: {})",
|
||||
existing_key.name
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
if auth_config_present {
|
||||
updated.encrypted_auth_config = auth_config
|
||||
.as_ref()
|
||||
.map(serde_json::to_string)
|
||||
.transpose()
|
||||
.map_err(|err| err.to_string())?
|
||||
.map(|plaintext| {
|
||||
encrypt_catalog_secret_with_fallbacks(state, &plaintext)
|
||||
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())
|
||||
})
|
||||
.transpose()?;
|
||||
}
|
||||
}
|
||||
"oauth" => {
|
||||
if api_key_present
|
||||
&& !matches!(
|
||||
api_key_value.as_deref(),
|
||||
None | Some("") | Some("__placeholder__")
|
||||
)
|
||||
{
|
||||
return Err("OAuth 认证模式下不允许直接填写 api_key".to_string());
|
||||
}
|
||||
if auth_type_switch {
|
||||
updated.encrypted_api_key =
|
||||
encrypt_catalog_secret_with_fallbacks(state, "__placeholder__")
|
||||
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())?;
|
||||
updated.encrypted_auth_config = None;
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
if raw_payload.contains_key("api_formats") {
|
||||
let api_formats = normalize_string_list(payload.api_formats)
|
||||
.ok_or_else(|| "api_formats 为必填字段".to_string())?;
|
||||
validate_vertex_api_formats(&provider.provider_type, &target_auth_type, &api_formats)?;
|
||||
updated.api_formats = Some(json!(api_formats));
|
||||
} else if payload.auth_type.is_some() {
|
||||
let api_formats = json_string_list(existing.api_formats.as_ref());
|
||||
validate_vertex_api_formats(&provider.provider_type, &target_auth_type, &api_formats)?;
|
||||
}
|
||||
|
||||
updated.auth_type = target_auth_type;
|
||||
|
||||
if let Some(name) = payload.name {
|
||||
let trimmed = name.trim();
|
||||
if trimmed.is_empty() {
|
||||
return Err("name 为必填字段".to_string());
|
||||
}
|
||||
updated.name = trimmed.to_string();
|
||||
}
|
||||
if raw_payload.contains_key("rate_multipliers") {
|
||||
updated.rate_multipliers =
|
||||
normalize_json_object(payload.rate_multipliers, "rate_multipliers")?;
|
||||
}
|
||||
if let Some(internal_priority) = payload.internal_priority {
|
||||
updated.internal_priority = internal_priority;
|
||||
}
|
||||
if raw_payload.contains_key("global_priority_by_format") {
|
||||
updated.global_priority_by_format = normalize_json_object(
|
||||
payload.global_priority_by_format,
|
||||
"global_priority_by_format",
|
||||
)?;
|
||||
}
|
||||
if raw_payload.contains_key("rpm_limit") {
|
||||
updated.rpm_limit = payload.rpm_limit;
|
||||
if payload.rpm_limit.is_none() {
|
||||
updated.learned_rpm_limit = None;
|
||||
}
|
||||
}
|
||||
if raw_payload.contains_key("allowed_models") {
|
||||
updated.allowed_models =
|
||||
normalize_string_list(payload.allowed_models).map(|value| json!(value));
|
||||
}
|
||||
if raw_payload.contains_key("capabilities") {
|
||||
updated.capabilities = normalize_json_object(payload.capabilities, "capabilities")?;
|
||||
}
|
||||
if let Some(cache_ttl_minutes) = payload.cache_ttl_minutes {
|
||||
updated.cache_ttl_minutes = cache_ttl_minutes;
|
||||
}
|
||||
if let Some(max_probe_interval_minutes) = payload.max_probe_interval_minutes {
|
||||
updated.max_probe_interval_minutes = max_probe_interval_minutes;
|
||||
}
|
||||
if let Some(is_active) = payload.is_active {
|
||||
updated.is_active = is_active;
|
||||
}
|
||||
if raw_payload.contains_key("note") {
|
||||
updated.note = payload
|
||||
.note
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty());
|
||||
}
|
||||
if let Some(auto_fetch_models) = payload.auto_fetch_models {
|
||||
updated.auto_fetch_models = auto_fetch_models;
|
||||
}
|
||||
if raw_payload.contains_key("locked_models") {
|
||||
updated.locked_models =
|
||||
normalize_string_list(payload.locked_models).map(|value| json!(value));
|
||||
}
|
||||
if raw_payload.contains_key("model_include_patterns") {
|
||||
updated.model_include_patterns =
|
||||
normalize_string_list(payload.model_include_patterns).map(|value| json!(value));
|
||||
}
|
||||
if raw_payload.contains_key("model_exclude_patterns") {
|
||||
updated.model_exclude_patterns =
|
||||
normalize_string_list(payload.model_exclude_patterns).map(|value| json!(value));
|
||||
}
|
||||
if raw_payload.contains_key("proxy") {
|
||||
updated.proxy = normalize_json_object(payload.proxy, "proxy")?;
|
||||
}
|
||||
if raw_payload.contains_key("fingerprint") {
|
||||
updated.fingerprint = normalize_json_object(payload.fingerprint, "fingerprint")?;
|
||||
}
|
||||
if auth_config_present && !auth_type_switch && updated.auth_type != "api_key" {
|
||||
updated.encrypted_auth_config = auth_config
|
||||
.as_ref()
|
||||
.map(serde_json::to_string)
|
||||
.transpose()
|
||||
.map_err(|err| err.to_string())?
|
||||
.map(|plaintext| {
|
||||
encrypt_catalog_secret_with_fallbacks(state, &plaintext)
|
||||
.ok_or_else(|| "gateway 未配置 provider key 加密密钥".to_string())
|
||||
})
|
||||
.transpose()?;
|
||||
}
|
||||
|
||||
updated.updated_at_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs());
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_provider_keys_payload(
|
||||
state: &AppState,
|
||||
provider_id: &str,
|
||||
skip: usize,
|
||||
limit: usize,
|
||||
) -> Option<serde_json::Value> {
|
||||
if !state.has_provider_catalog_data_reader() {
|
||||
return None;
|
||||
}
|
||||
let provider = state
|
||||
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
|
||||
.await
|
||||
.ok()
|
||||
.and_then(|mut providers| providers.drain(..).next())?;
|
||||
let mut keys = state
|
||||
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
|
||||
.await
|
||||
.ok()
|
||||
.unwrap_or_default();
|
||||
keys.sort_by(|left, right| {
|
||||
left.internal_priority
|
||||
.cmp(&right.internal_priority)
|
||||
.then_with(|| {
|
||||
left.created_at_unix_secs
|
||||
.unwrap_or_default()
|
||||
.cmp(&right.created_at_unix_secs.unwrap_or_default())
|
||||
})
|
||||
.then_with(|| left.id.cmp(&right.id))
|
||||
});
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
Some(serde_json::Value::Array(
|
||||
keys.into_iter()
|
||||
.skip(skip)
|
||||
.take(limit)
|
||||
.map(|key| build_admin_provider_key_response(state, &key, now_unix_secs))
|
||||
.collect(),
|
||||
))
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
use crate::handlers::admin::shared::{normalize_json_object, normalize_string_list};
|
||||
|
||||
pub(crate) fn normalize_provider_type_input(value: &str) -> Result<String, String> {
|
||||
let normalized = value.trim().to_ascii_lowercase();
|
||||
match normalized.as_str() {
|
||||
"custom" | "claude_code" | "kiro" | "codex" | "gemini_cli" | "antigravity"
|
||||
| "vertex_ai" => Ok(normalized),
|
||||
_ => Err(
|
||||
"provider_type 仅支持 custom / claude_code / kiro / codex / gemini_cli / antigravity / vertex_ai"
|
||||
.to_string(),
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn normalize_provider_billing_type(value: &str) -> Result<String, String> {
|
||||
let normalized = value.trim().to_ascii_lowercase();
|
||||
match normalized.as_str() {
|
||||
"monthly_quota" | "pay_as_you_go" | "free_tier" => Ok(normalized),
|
||||
_ => Err("billing_type 仅支持 monthly_quota / pay_as_you_go / free_tier".to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn parse_optional_rfc3339_unix_secs(
|
||||
value: &str,
|
||||
field_name: &str,
|
||||
) -> Result<u64, String> {
|
||||
let trimmed = value.trim();
|
||||
if trimmed.is_empty() {
|
||||
return Err(format!("{field_name} 不能为空"));
|
||||
}
|
||||
let parsed = chrono::DateTime::parse_from_rfc3339(trimmed)
|
||||
.map_err(|_| format!("{field_name} 必须是合法的 RFC3339 时间"))?;
|
||||
u64::try_from(parsed.timestamp()).map_err(|_| format!("{field_name} 超出有效时间范围"))
|
||||
}
|
||||
|
||||
pub(crate) fn normalize_auth_type(value: Option<&str>) -> Result<String, String> {
|
||||
let auth_type = value.unwrap_or("api_key").trim().to_ascii_lowercase();
|
||||
match auth_type.as_str() {
|
||||
"api_key" | "service_account" | "oauth" => Ok(auth_type),
|
||||
_ => Err("auth_type 仅支持 api_key / service_account / oauth".to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn validate_vertex_api_formats(
|
||||
provider_type: &str,
|
||||
auth_type: &str,
|
||||
api_formats: &[String],
|
||||
) -> Result<(), String> {
|
||||
if !provider_type.trim().eq_ignore_ascii_case("vertex_ai") {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let allowed = match auth_type {
|
||||
"api_key" => &["gemini:chat"][..],
|
||||
"service_account" | "vertex_ai" => &["claude:chat", "gemini:chat"][..],
|
||||
_ => return Ok(()),
|
||||
};
|
||||
let invalid = api_formats
|
||||
.iter()
|
||||
.filter(|value| !allowed.contains(&value.as_str()))
|
||||
.cloned()
|
||||
.collect::<Vec<_>>();
|
||||
if invalid.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
Err(format!(
|
||||
"Vertex {auth_type} 不支持以下 API 格式: {};允许: {}",
|
||||
invalid.join(", "),
|
||||
allowed.join(", ")
|
||||
))
|
||||
}
|
||||
|
||||
mod keys;
|
||||
mod provider;
|
||||
mod reveal;
|
||||
|
||||
pub(crate) use self::keys::{
|
||||
build_admin_create_provider_key_record, build_admin_provider_keys_payload,
|
||||
build_admin_update_provider_key_record,
|
||||
};
|
||||
pub(crate) use self::provider::{
|
||||
build_admin_create_provider_record, build_admin_fixed_provider_endpoint_record,
|
||||
build_admin_update_provider_record,
|
||||
};
|
||||
pub(crate) use self::reveal::{build_admin_export_key_payload, build_admin_reveal_key_payload};
|
||||
@@ -0,0 +1,510 @@
|
||||
use super::{
|
||||
normalize_json_object, normalize_provider_billing_type, normalize_provider_type_input,
|
||||
parse_optional_rfc3339_unix_secs,
|
||||
};
|
||||
use crate::api::ai::{admin_default_body_rules_for_signature, admin_endpoint_signature_parts};
|
||||
use crate::handlers::admin::provider::shared::{
|
||||
AdminProviderCreateRequest, AdminProviderUpdateRequest,
|
||||
};
|
||||
use crate::handlers::public::normalize_admin_base_url;
|
||||
use crate::provider_transport::provider_types::provider_type_enables_format_conversion_by_default;
|
||||
use crate::AppState;
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
|
||||
};
|
||||
use serde_json::json;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
use uuid::Uuid;
|
||||
|
||||
pub(crate) async fn build_admin_update_provider_record(
|
||||
state: &AppState,
|
||||
existing: &StoredProviderCatalogProvider,
|
||||
raw_payload: &serde_json::Map<String, serde_json::Value>,
|
||||
payload: AdminProviderUpdateRequest,
|
||||
) -> Result<StoredProviderCatalogProvider, String> {
|
||||
let mut updated = existing.clone();
|
||||
|
||||
if let Some(value) = raw_payload.get("name") {
|
||||
let Some(name) = payload.name.as_deref() else {
|
||||
return Err(if value.is_null() {
|
||||
"name 不能为空".to_string()
|
||||
} else {
|
||||
"name 必须是字符串".to_string()
|
||||
});
|
||||
};
|
||||
let trimmed = name.trim();
|
||||
if trimmed.is_empty() {
|
||||
return Err("name 不能为空".to_string());
|
||||
}
|
||||
let duplicate = state
|
||||
.list_provider_catalog_providers(false)
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?
|
||||
.into_iter()
|
||||
.any(|provider| provider.id != existing.id && provider.name == trimmed);
|
||||
if duplicate {
|
||||
return Err(format!("提供商名称 '{trimmed}' 已存在"));
|
||||
}
|
||||
updated.name = trimmed.to_string();
|
||||
}
|
||||
|
||||
let target_provider_type = if let Some(value) = raw_payload.get("provider_type") {
|
||||
let Some(provider_type) = payload.provider_type.as_deref() else {
|
||||
return Err(if value.is_null() {
|
||||
"provider_type 不能为空".to_string()
|
||||
} else {
|
||||
"provider_type 必须是字符串".to_string()
|
||||
});
|
||||
};
|
||||
let normalized = normalize_provider_type_input(provider_type)?;
|
||||
updated.provider_type = normalized.clone();
|
||||
normalized
|
||||
} else {
|
||||
updated.provider_type.clone()
|
||||
};
|
||||
|
||||
if raw_payload.contains_key("description") {
|
||||
updated.description = payload
|
||||
.description
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty());
|
||||
}
|
||||
|
||||
if let Some(value) = raw_payload.get("website") {
|
||||
updated.website = match payload.website {
|
||||
None => {
|
||||
if value.is_null() {
|
||||
None
|
||||
} else {
|
||||
return Err("website 必须是字符串".to_string());
|
||||
}
|
||||
}
|
||||
Some(website) => {
|
||||
let trimmed = website.trim();
|
||||
if trimmed.is_empty() {
|
||||
None
|
||||
} else if !trimmed.starts_with("http://") && !trimmed.starts_with("https://") {
|
||||
return Err("website 必须以 http:// 或 https:// 开头".to_string());
|
||||
} else {
|
||||
Some(trimmed.to_string())
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
if let Some(value) = raw_payload.get("billing_type") {
|
||||
let Some(billing_type) = payload.billing_type.as_deref() else {
|
||||
return Err(if value.is_null() {
|
||||
"billing_type 不能为空".to_string()
|
||||
} else {
|
||||
"billing_type 必须是字符串".to_string()
|
||||
});
|
||||
};
|
||||
updated.billing_type = Some(normalize_provider_billing_type(billing_type)?);
|
||||
}
|
||||
|
||||
if let Some(value) = raw_payload.get("monthly_quota_usd") {
|
||||
if value.is_null() {
|
||||
updated.monthly_quota_usd = None;
|
||||
} else {
|
||||
let Some(monthly_quota_usd) = payload.monthly_quota_usd else {
|
||||
return Err("monthly_quota_usd 必须是非负数".to_string());
|
||||
};
|
||||
if !monthly_quota_usd.is_finite() || monthly_quota_usd < 0.0 {
|
||||
return Err("monthly_quota_usd 必须是非负数".to_string());
|
||||
}
|
||||
updated.monthly_quota_usd = Some(monthly_quota_usd);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(value) = raw_payload.get("quota_reset_day") {
|
||||
if value.is_null() {
|
||||
updated.quota_reset_day = None;
|
||||
} else {
|
||||
let Some(quota_reset_day) = payload.quota_reset_day else {
|
||||
return Err("quota_reset_day 必须是 1 到 365 之间的整数".to_string());
|
||||
};
|
||||
if !(1..=365).contains("a_reset_day) {
|
||||
return Err("quota_reset_day 必须是 1 到 365 之间的整数".to_string());
|
||||
}
|
||||
updated.quota_reset_day = Some(quota_reset_day);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(value) = raw_payload.get("quota_last_reset_at") {
|
||||
if value.is_null() {
|
||||
updated.quota_last_reset_at_unix_secs = None;
|
||||
} else {
|
||||
let Some(raw) = payload.quota_last_reset_at.as_deref() else {
|
||||
return Err("quota_last_reset_at 必须是字符串".to_string());
|
||||
};
|
||||
updated.quota_last_reset_at_unix_secs = Some(parse_optional_rfc3339_unix_secs(
|
||||
raw,
|
||||
"quota_last_reset_at",
|
||||
)?);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(value) = raw_payload.get("quota_expires_at") {
|
||||
if value.is_null() {
|
||||
updated.quota_expires_at_unix_secs = None;
|
||||
} else {
|
||||
let Some(raw) = payload.quota_expires_at.as_deref() else {
|
||||
return Err("quota_expires_at 必须是字符串".to_string());
|
||||
};
|
||||
updated.quota_expires_at_unix_secs =
|
||||
Some(parse_optional_rfc3339_unix_secs(raw, "quota_expires_at")?);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(value) = raw_payload.get("provider_priority") {
|
||||
let Some(provider_priority) = payload.provider_priority else {
|
||||
return Err(if value.is_null() {
|
||||
"provider_priority 不能为空".to_string()
|
||||
} else {
|
||||
"provider_priority 必须是整数".to_string()
|
||||
});
|
||||
};
|
||||
if !(0..=10_000).contains(&provider_priority) {
|
||||
return Err("provider_priority 必须在 0 到 10000 之间".to_string());
|
||||
}
|
||||
updated.provider_priority = provider_priority;
|
||||
}
|
||||
|
||||
if let Some(_value) = raw_payload.get("keep_priority_on_conversion") {
|
||||
let Some(keep_priority_on_conversion) = payload.keep_priority_on_conversion else {
|
||||
return Err("keep_priority_on_conversion 必须是布尔值".to_string());
|
||||
};
|
||||
updated.keep_priority_on_conversion = keep_priority_on_conversion;
|
||||
}
|
||||
|
||||
if let Some(_value) = raw_payload.get("is_active") {
|
||||
let Some(is_active) = payload.is_active else {
|
||||
return Err("is_active 必须是布尔值".to_string());
|
||||
};
|
||||
updated.is_active = is_active;
|
||||
}
|
||||
|
||||
if raw_payload.contains_key("concurrent_limit") {
|
||||
updated.concurrent_limit = match payload.concurrent_limit {
|
||||
Some(value) if value >= 0 => Some(value),
|
||||
Some(_) => return Err("concurrent_limit 必须是非负整数".to_string()),
|
||||
None => None,
|
||||
};
|
||||
}
|
||||
|
||||
if raw_payload.contains_key("max_retries") {
|
||||
updated.max_retries = match payload.max_retries {
|
||||
Some(value) if (0..=999).contains(&value) => Some(value),
|
||||
Some(_) => return Err("max_retries 必须是 0 到 999 之间的整数".to_string()),
|
||||
None => None,
|
||||
};
|
||||
}
|
||||
|
||||
if raw_payload.contains_key("proxy") {
|
||||
updated.proxy = normalize_json_object(payload.proxy, "proxy")?;
|
||||
}
|
||||
|
||||
if raw_payload.contains_key("stream_first_byte_timeout") {
|
||||
updated.stream_first_byte_timeout_secs = match payload.stream_first_byte_timeout {
|
||||
Some(value) if (1.0..=300.0).contains(&value) => Some(value),
|
||||
Some(_) => {
|
||||
return Err("stream_first_byte_timeout 必须是 1 到 300 之间的数字".to_string())
|
||||
}
|
||||
None => None,
|
||||
};
|
||||
}
|
||||
|
||||
if raw_payload.contains_key("request_timeout") {
|
||||
updated.request_timeout_secs = match payload.request_timeout {
|
||||
Some(value) if (1.0..=600.0).contains(&value) => Some(value),
|
||||
Some(_) => return Err("request_timeout 必须是 1 到 600 之间的数字".to_string()),
|
||||
None => None,
|
||||
};
|
||||
}
|
||||
|
||||
if let Some(_value) = raw_payload.get("enable_format_conversion") {
|
||||
let Some(enable_format_conversion) = payload.enable_format_conversion else {
|
||||
return Err("enable_format_conversion 必须是布尔值".to_string());
|
||||
};
|
||||
updated.enable_format_conversion = enable_format_conversion;
|
||||
}
|
||||
|
||||
let config_seed = if raw_payload.contains_key("config") {
|
||||
normalize_json_object(payload.config, "config")?
|
||||
} else {
|
||||
updated.config.clone()
|
||||
};
|
||||
let mut config_map = config_seed
|
||||
.and_then(|value| value.as_object().cloned())
|
||||
.unwrap_or_default();
|
||||
|
||||
if raw_payload.contains_key("claude_code_advanced") {
|
||||
if raw_payload
|
||||
.get("claude_code_advanced")
|
||||
.is_some_and(serde_json::Value::is_null)
|
||||
{
|
||||
config_map.remove("claude_code_advanced");
|
||||
} else {
|
||||
if target_provider_type != "claude_code" {
|
||||
return Err("claude_code_advanced 仅适用于 provider_type=claude_code".to_string());
|
||||
}
|
||||
let value =
|
||||
normalize_json_object(payload.claude_code_advanced, "claude_code_advanced")?
|
||||
.ok_or_else(|| "claude_code_advanced 必须是 JSON 对象".to_string())?;
|
||||
config_map.insert("claude_code_advanced".to_string(), value);
|
||||
}
|
||||
} else if target_provider_type != "claude_code" {
|
||||
config_map.remove("claude_code_advanced");
|
||||
}
|
||||
|
||||
if raw_payload.contains_key("pool_advanced") {
|
||||
if raw_payload
|
||||
.get("pool_advanced")
|
||||
.is_some_and(serde_json::Value::is_null)
|
||||
{
|
||||
config_map.remove("pool_advanced");
|
||||
} else {
|
||||
let value = normalize_json_object(payload.pool_advanced, "pool_advanced")?
|
||||
.ok_or_else(|| "pool_advanced 必须是 JSON 对象".to_string())?;
|
||||
config_map.insert("pool_advanced".to_string(), value);
|
||||
}
|
||||
}
|
||||
|
||||
if raw_payload.contains_key("failover_rules") {
|
||||
if raw_payload
|
||||
.get("failover_rules")
|
||||
.is_some_and(serde_json::Value::is_null)
|
||||
{
|
||||
config_map.remove("failover_rules");
|
||||
} else {
|
||||
let value = normalize_json_object(payload.failover_rules, "failover_rules")?
|
||||
.ok_or_else(|| "failover_rules 必须是 JSON 对象".to_string())?;
|
||||
config_map.insert("failover_rules".to_string(), value);
|
||||
}
|
||||
}
|
||||
|
||||
updated.config = (!config_map.is_empty()).then_some(serde_json::Value::Object(config_map));
|
||||
updated.updated_at_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs());
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_create_provider_record(
|
||||
state: &AppState,
|
||||
payload: AdminProviderCreateRequest,
|
||||
) -> Result<(StoredProviderCatalogProvider, Option<i32>), String> {
|
||||
let name = payload.name.trim();
|
||||
if name.is_empty() {
|
||||
return Err("name 为必填字段".to_string());
|
||||
}
|
||||
|
||||
let existing_providers = state
|
||||
.list_provider_catalog_providers(false)
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?;
|
||||
if existing_providers
|
||||
.iter()
|
||||
.any(|provider| provider.name == name)
|
||||
{
|
||||
return Err(format!("提供商名称 '{name}' 已存在"));
|
||||
}
|
||||
|
||||
let provider_type =
|
||||
normalize_provider_type_input(payload.provider_type.as_deref().unwrap_or("custom"))?;
|
||||
let billing_type = normalize_provider_billing_type(
|
||||
payload.billing_type.as_deref().unwrap_or("pay_as_you_go"),
|
||||
)?;
|
||||
|
||||
let website = payload.website.and_then(|value| {
|
||||
let trimmed = value.trim().to_string();
|
||||
(!trimmed.is_empty()).then_some(trimmed)
|
||||
});
|
||||
let website = website.map(|value| {
|
||||
if value.starts_with("http://") || value.starts_with("https://") {
|
||||
value
|
||||
} else {
|
||||
format!("https://{value}")
|
||||
}
|
||||
});
|
||||
|
||||
let monthly_quota_usd = match payload.monthly_quota_usd {
|
||||
Some(value) if value.is_finite() && value >= 0.0 => Some(value),
|
||||
Some(_) => return Err("monthly_quota_usd 必须是非负数".to_string()),
|
||||
None => None,
|
||||
};
|
||||
let quota_reset_day = match payload.quota_reset_day {
|
||||
Some(value) if (1..=365).contains(&value) => Some(value),
|
||||
Some(_) => return Err("quota_reset_day 必须是 1 到 365 之间的整数".to_string()),
|
||||
None => Some(30),
|
||||
};
|
||||
let quota_last_reset_at_unix_secs = payload
|
||||
.quota_last_reset_at
|
||||
.as_deref()
|
||||
.map(|value| parse_optional_rfc3339_unix_secs(value, "quota_last_reset_at"))
|
||||
.transpose()?;
|
||||
let quota_expires_at_unix_secs = payload
|
||||
.quota_expires_at
|
||||
.as_deref()
|
||||
.map(|value| parse_optional_rfc3339_unix_secs(value, "quota_expires_at"))
|
||||
.transpose()?;
|
||||
let provider_priority = match payload.provider_priority {
|
||||
Some(value) if (0..=10_000).contains(&value) => value,
|
||||
Some(_) => return Err("provider_priority 必须在 0 到 10000 之间".to_string()),
|
||||
None => {
|
||||
let current_min_priority = existing_providers
|
||||
.iter()
|
||||
.map(|provider| provider.provider_priority)
|
||||
.min();
|
||||
match current_min_priority {
|
||||
Some(value) if value <= 0 => 0,
|
||||
Some(value) => value - 1,
|
||||
None => 100,
|
||||
}
|
||||
}
|
||||
};
|
||||
let shift_existing_priorities_from = match payload.provider_priority {
|
||||
Some(_) => Some(provider_priority),
|
||||
None => existing_providers
|
||||
.iter()
|
||||
.map(|provider| provider.provider_priority)
|
||||
.min()
|
||||
.filter(|value| *value <= 0)
|
||||
.map(|_| 0),
|
||||
};
|
||||
|
||||
let is_active = payload.is_active.unwrap_or(true);
|
||||
let concurrent_limit = match payload.concurrent_limit {
|
||||
Some(value) if value >= 0 => Some(value),
|
||||
Some(_) => return Err("concurrent_limit 必须是非负整数".to_string()),
|
||||
None => None,
|
||||
};
|
||||
let max_retries = match payload.max_retries {
|
||||
Some(value) if (0..=999).contains(&value) => Some(value),
|
||||
Some(_) => return Err("max_retries 必须是 0 到 999 之间的整数".to_string()),
|
||||
None => Some(2),
|
||||
};
|
||||
let proxy = normalize_json_object(payload.proxy, "proxy")?;
|
||||
let stream_first_byte_timeout_secs = match payload.stream_first_byte_timeout {
|
||||
Some(value) if (1.0..=300.0).contains(&value) => Some(value),
|
||||
Some(_) => return Err("stream_first_byte_timeout 必须是 1 到 300 之间的数字".to_string()),
|
||||
None => None,
|
||||
};
|
||||
let request_timeout_secs = match payload.request_timeout {
|
||||
Some(value) if (1.0..=600.0).contains(&value) => Some(value),
|
||||
Some(_) => return Err("request_timeout 必须是 1 到 600 之间的数字".to_string()),
|
||||
None => None,
|
||||
};
|
||||
|
||||
let mut config_map = normalize_json_object(payload.config, "config")?
|
||||
.and_then(|value| value.as_object().cloned())
|
||||
.unwrap_or_default();
|
||||
if let Some(value) = normalize_json_object(payload.pool_advanced, "pool_advanced")? {
|
||||
config_map.insert("pool_advanced".to_string(), value);
|
||||
}
|
||||
if let Some(value) = normalize_json_object(payload.failover_rules, "failover_rules")? {
|
||||
config_map.insert("failover_rules".to_string(), value);
|
||||
}
|
||||
if let Some(value) =
|
||||
normalize_json_object(payload.claude_code_advanced, "claude_code_advanced")?
|
||||
{
|
||||
if provider_type != "claude_code" {
|
||||
return Err("claude_code_advanced 仅适用于 provider_type=claude_code".to_string());
|
||||
}
|
||||
config_map.insert("claude_code_advanced".to_string(), value);
|
||||
}
|
||||
let config = (!config_map.is_empty()).then_some(serde_json::Value::Object(config_map));
|
||||
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
|
||||
let record = StoredProviderCatalogProvider::new(
|
||||
Uuid::new_v4().to_string(),
|
||||
name.to_string(),
|
||||
website,
|
||||
provider_type.clone(),
|
||||
)
|
||||
.map_err(|err| err.to_string())?
|
||||
.with_description(
|
||||
payload
|
||||
.description
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty()),
|
||||
)
|
||||
.with_billing_fields(
|
||||
Some(billing_type),
|
||||
monthly_quota_usd,
|
||||
None,
|
||||
quota_reset_day,
|
||||
quota_last_reset_at_unix_secs,
|
||||
quota_expires_at_unix_secs,
|
||||
)
|
||||
.with_routing_fields(provider_priority)
|
||||
.with_transport_fields(
|
||||
is_active,
|
||||
payload.keep_priority_on_conversion.unwrap_or(false),
|
||||
provider_type_enables_format_conversion_by_default(&provider_type),
|
||||
concurrent_limit,
|
||||
max_retries,
|
||||
proxy,
|
||||
request_timeout_secs,
|
||||
stream_first_byte_timeout_secs,
|
||||
config,
|
||||
)
|
||||
.with_timestamps(Some(now_unix_secs), Some(now_unix_secs));
|
||||
|
||||
Ok((record, shift_existing_priorities_from))
|
||||
}
|
||||
|
||||
pub(crate) fn build_admin_fixed_provider_endpoint_record(
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
api_format: &str,
|
||||
base_url: &str,
|
||||
) -> Result<StoredProviderCatalogEndpoint, String> {
|
||||
let (normalized_api_format, api_family, endpoint_kind) =
|
||||
admin_endpoint_signature_parts(api_format)
|
||||
.ok_or_else(|| format!("无效的 api_format: {api_format}"))?;
|
||||
let body_rules = admin_default_body_rules_for_signature(
|
||||
normalized_api_format,
|
||||
Some(provider.provider_type.as_str()),
|
||||
)
|
||||
.and_then(|(_, rules)| (!rules.is_empty()).then_some(serde_json::Value::Array(rules)));
|
||||
let endpoint_config =
|
||||
if provider.provider_type == "codex" && normalized_api_format == "openai:cli" {
|
||||
Some(json!({ "upstream_stream_policy": "force_stream" }))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
|
||||
StoredProviderCatalogEndpoint::new(
|
||||
Uuid::new_v4().to_string(),
|
||||
provider.id.clone(),
|
||||
normalized_api_format.to_string(),
|
||||
Some(api_family.to_string()),
|
||||
Some(endpoint_kind.to_string()),
|
||||
true,
|
||||
)
|
||||
.map_err(|err| err.to_string())?
|
||||
.with_timestamps(Some(now_unix_secs), Some(now_unix_secs))
|
||||
.with_transport_fields(
|
||||
normalize_admin_base_url(base_url)?,
|
||||
None,
|
||||
body_rules,
|
||||
Some(provider.max_retries.unwrap_or(2)),
|
||||
None,
|
||||
endpoint_config,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.map_err(|err| err.to_string())
|
||||
}
|
||||
@@ -0,0 +1,151 @@
|
||||
fn normalize_reveal_auth_type(value: &str) -> &str {
|
||||
match value.trim().to_ascii_lowercase().as_str() {
|
||||
"service_account" | "vertex_ai" => "service_account",
|
||||
"oauth" => "oauth",
|
||||
_ => "api_key",
|
||||
}
|
||||
}
|
||||
|
||||
use crate::handlers::admin::shared::{
|
||||
decrypt_catalog_secret_with_fallbacks, parse_catalog_auth_config_json,
|
||||
};
|
||||
use crate::AppState;
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||
use chrono::{SecondsFormat, Utc};
|
||||
use serde_json::json;
|
||||
|
||||
pub(crate) fn build_admin_reveal_key_payload(
|
||||
state: &AppState,
|
||||
key: &StoredProviderCatalogKey,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let auth_type = normalize_reveal_auth_type(&key.auth_type);
|
||||
if matches!(auth_type, "service_account") {
|
||||
if let Some(auth_config) = parse_catalog_auth_config_json(state, key) {
|
||||
return Ok(json!({
|
||||
"auth_type": auth_type,
|
||||
"auth_config": auth_config,
|
||||
}));
|
||||
}
|
||||
let decrypted =
|
||||
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), &key.encrypted_api_key)
|
||||
.ok_or_else(|| {
|
||||
"无法解密认证配置,可能是加密密钥已更改。请重新添加该密钥。".to_string()
|
||||
})?;
|
||||
if decrypted == "__placeholder__" {
|
||||
return Err("认证配置丢失,请重新添加该密钥。".to_string());
|
||||
}
|
||||
return Ok(json!({
|
||||
"auth_type": auth_type,
|
||||
"auth_config": decrypted,
|
||||
}));
|
||||
}
|
||||
|
||||
let decrypted =
|
||||
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), &key.encrypted_api_key)
|
||||
.ok_or_else(|| {
|
||||
"无法解密 API Key,可能是加密密钥已更改。请重新添加该密钥。".to_string()
|
||||
})?;
|
||||
Ok(json!({
|
||||
"auth_type": auth_type,
|
||||
"api_key": decrypted,
|
||||
}))
|
||||
}
|
||||
|
||||
fn provider_oauth_export_payload(
|
||||
provider_type: &str,
|
||||
auth_config: &serde_json::Map<String, serde_json::Value>,
|
||||
upstream_metadata: Option<&serde_json::Value>,
|
||||
) -> serde_json::Map<String, serde_json::Value> {
|
||||
let normalized_provider_type = provider_type.trim().to_ascii_lowercase();
|
||||
let skip_keys: &[&str] = match normalized_provider_type.as_str() {
|
||||
"kiro" => &["access_token", "expires_at", "updated_at"],
|
||||
_ => &[
|
||||
"access_token",
|
||||
"expires_at",
|
||||
"updated_at",
|
||||
"token_type",
|
||||
"scope",
|
||||
],
|
||||
};
|
||||
let mut payload = serde_json::Map::new();
|
||||
for (key, value) in auth_config {
|
||||
if skip_keys.contains(&key.as_str()) {
|
||||
continue;
|
||||
}
|
||||
if value.is_null() || value.as_str().is_some_and(str::is_empty) {
|
||||
continue;
|
||||
}
|
||||
payload.insert(key.clone(), value.clone());
|
||||
}
|
||||
if normalized_provider_type == "kiro" && !payload.contains_key("email") {
|
||||
if let Some(email) = upstream_metadata
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|meta| meta.get("kiro"))
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|meta| meta.get("email"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
{
|
||||
payload.insert("email".to_string(), json!(email));
|
||||
}
|
||||
}
|
||||
payload
|
||||
}
|
||||
|
||||
pub(crate) async fn build_admin_export_key_payload(
|
||||
state: &AppState,
|
||||
key: &StoredProviderCatalogKey,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let auth_type = normalize_reveal_auth_type(&key.auth_type);
|
||||
if auth_type != "oauth" {
|
||||
return Err("仅 OAuth 类型的 Key 支持导出".to_string());
|
||||
}
|
||||
|
||||
let ciphertext = key
|
||||
.encrypted_auth_config
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.ok_or_else(|| "缺少认证配置,无法导出".to_string())?;
|
||||
let plaintext = decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)
|
||||
.ok_or_else(|| "无法解密认证配置".to_string())?;
|
||||
let auth_config = serde_json::from_str::<serde_json::Value>(&plaintext)
|
||||
.ok()
|
||||
.and_then(|value| value.as_object().cloned())
|
||||
.ok_or_else(|| "无法解密认证配置".to_string())?;
|
||||
if !auth_config
|
||||
.get("refresh_token")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.is_some_and(|value| !value.trim().is_empty())
|
||||
{
|
||||
return Err("缺少 refresh_token,无法导出".to_string());
|
||||
}
|
||||
|
||||
let provider_type_from_config = auth_config
|
||||
.get("provider_type")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
let provider_type = if let Some(provider_type) = provider_type_from_config {
|
||||
provider_type
|
||||
} else {
|
||||
state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&key.provider_id))
|
||||
.await
|
||||
.map_err(|err| format!("{err:?}"))?
|
||||
.into_iter()
|
||||
.next()
|
||||
.map(|provider| provider.provider_type)
|
||||
.unwrap_or_default()
|
||||
};
|
||||
|
||||
let mut payload =
|
||||
provider_oauth_export_payload(&provider_type, &auth_config, key.upstream_metadata.as_ref());
|
||||
payload.insert("name".to_string(), json!(key.name));
|
||||
payload.insert(
|
||||
"exported_at".to_string(),
|
||||
json!(Utc::now().to_rfc3339_opts(SecondsFormat::Secs, true)),
|
||||
);
|
||||
Ok(serde_json::Value::Object(payload))
|
||||
}
|
||||
Reference in New Issue
Block a user