fix(admin): 优先级脏检查、base_url 校验、usage detail 延迟加载及导入数据验证

- 前端优先级管理: 保存时对比原始快照,仅提交实际变更的 provider/key 优先级,
  并限制并发请求数(SAVE_CONCURRENCY=6),避免无效 API 调用
- handler_adapter_base: _normalize_test_base_url 改为 _validate_test_base_url,
  移除对 dict 类型 base_url 的兼容,严格要求字符串输入
- provider_query: 新增 _require_test_endpoint_base_url,在测试链路提前校验
  endpoint.base_url 类型和非空
- system.py: 导入 endpoint 时通过 ProviderEndpointCreate 模型校验数据,
  拒绝非法 base_url 类型(如 dict)
- usage detail: 使用 defer() 延迟加载 body 列,通过 SQL CASE 表达式在
  数据库端计算 has_*_body 标记,减少不必要的大字段传输
- provider routes: 新建 provider 时 priority=0 边界处理,clamp 并 shift
This commit is contained in:
fawney19
2026-03-12 17:11:14 +08:00
parent 8d69f72e2a
commit e9678ea899
11 changed files with 553 additions and 145 deletions

View File

@@ -537,6 +537,12 @@ const dragOverKey = ref<Record<string, string | null>>({})
const loadingKeys = ref(false) const loadingKeys = ref(false)
const saving = ref(false) const saving = ref(false)
const SAVE_CONCURRENCY = 6
let originalProviderPriorityById = new Map<string, number>()
let originalPoolPriorityByProviderId = new Map<string, number | null>()
let originalKeyPriorityById = new Map<string, Record<string, number>>()
// Key 优先级编辑状态 // Key 优先级编辑状态
const editingKeyPriority = ref<Record<string, string | null>>({}) // format -> keyId const editingKeyPriority = ref<Record<string, string | null>>({}) // format -> keyId
@@ -642,6 +648,77 @@ function normalizePriorityMap(
return normalized return normalized
} }
function normalizeOptionalPriority(value: unknown): number | null {
const num = Number(value)
if (!Number.isFinite(num)) return null
return Math.trunc(num)
}
function normalizeRequiredPriority(value: unknown, fallback: number): number {
return normalizeOptionalPriority(value) ?? fallback
}
function normalizeProvidersForEditing(
providers: ProviderWithEndpointsSummary[]
): ProviderWithEndpointsSummary[] {
const normalized = providers.map((provider, index) => ({
...provider,
provider_priority: normalizeRequiredPriority(provider.provider_priority, 100 + index),
pool_advanced: provider.pool_advanced
? {
...provider.pool_advanced,
global_priority: normalizeOptionalPriority(provider.pool_advanced.global_priority),
}
: provider.pool_advanced,
}))
const minProviderPriority = normalized.reduce(
(min, provider) => Math.min(min, provider.provider_priority),
Number.POSITIVE_INFINITY,
)
const providerOffset = Number.isFinite(minProviderPriority) && minProviderPriority < 1
? 1 - minProviderPriority
: 0
const explicitPoolPriorities = normalized
.map((provider) => normalizeOptionalPriority(provider.pool_advanced?.global_priority))
.filter((priority): priority is number => priority != null)
const minPoolPriority = explicitPoolPriorities.length > 0
? Math.min(...explicitPoolPriorities)
: null
const poolOffset = minPoolPriority != null && minPoolPriority < 1
? 1 - minPoolPriority
: 0
return normalized.map((provider) => ({
...provider,
provider_priority: provider.provider_priority + providerOffset,
pool_advanced: provider.pool_advanced
? {
...provider.pool_advanced,
global_priority: provider.pool_advanced.global_priority == null
? null
: provider.pool_advanced.global_priority + poolOffset,
}
: provider.pool_advanced,
}))
}
function snapshotProviderBaseline(providers: ProviderWithEndpointsSummary[]) {
originalProviderPriorityById = new Map(
providers.map((provider, index) => [
provider.id,
normalizeRequiredPriority(provider.provider_priority, 100 + index),
])
)
originalPoolPriorityByProviderId = new Map(
providers.map((provider) => [
provider.id,
normalizeOptionalPriority(provider.pool_advanced?.global_priority),
])
)
}
function normalizeRateMultipliers( function normalizeRateMultipliers(
value: Record<string, unknown> | null | undefined value: Record<string, unknown> | null | undefined
): Record<string, number> | null { ): Record<string, number> | null {
@@ -657,6 +734,69 @@ function normalizeRateMultipliers(
return Object.keys(normalized).length > 0 ? normalized : null return Object.keys(normalized).length > 0 ? normalized : null
} }
function buildEditableKeyPriorityMap(): Map<string, Record<string, number>> {
const priorityMapByKeyId = new Map<string, Record<string, number>>()
for (const format of Object.keys(keysByFormat.value)) {
const normalizedFormat = normalizeApiFormatKey(format)
if (!normalizedFormat) continue
const keys = keysByFormat.value[format].filter((key) => !isPoolManagedKey(key))
keys.forEach((key) => {
const existing = priorityMapByKeyId.get(key.id) || normalizePriorityMap(key.global_priority_by_format)
existing[normalizedFormat] = Math.max(0, Math.trunc(key.priority))
priorityMapByKeyId.set(key.id, existing)
})
}
return priorityMapByKeyId
}
function snapshotKeyBaseline() {
originalKeyPriorityById = new Map(
Array.from(buildEditableKeyPriorityMap().entries()).map(([keyId, priorityMap]) => [
keyId,
{ ...priorityMap },
])
)
}
function arePriorityMapsEqual(
left: Record<string, number> | undefined,
right: Record<string, number> | undefined,
): boolean {
const leftEntries = Object.entries(left || {}).sort(([a], [b]) => a.localeCompare(b))
const rightEntries = Object.entries(right || {}).sort(([a], [b]) => a.localeCompare(b))
if (leftEntries.length !== rightEntries.length) return false
return leftEntries.every(([format, priority], index) => {
const [rightFormat, rightPriority] = rightEntries[index] || []
return format === rightFormat && priority === rightPriority
})
}
async function runTasksWithConcurrency(
tasks: Array<() => Promise<unknown>>,
concurrency: number = SAVE_CONCURRENCY,
) {
if (tasks.length === 0) return
let cursor = 0
const runNext = async (): Promise<void> => {
while (cursor < tasks.length) {
const index = cursor++
await tasks[index]()
}
}
const workers = Array.from(
{ length: Math.min(concurrency, tasks.length) },
() => runNext(),
)
await Promise.all(workers)
}
const providerById = computed(() => { const providerById = computed(() => {
const map = new Map<string, ProviderWithEndpointsSummary>() const map = new Map<string, ProviderWithEndpointsSummary>()
sortedProviders.value.forEach((provider) => { sortedProviders.value.forEach((provider) => {
@@ -857,8 +997,13 @@ watch(internalOpen, async (open) => {
async function loadAllProviders() { async function loadAllProviders() {
try { try {
const response = await getProvidersSummary({ page: 1, page_size: 9999 }) const response = await getProvidersSummary({ page: 1, page_size: 9999 })
sortedProviders.value = sortProvidersByActiveAndPriority(response.items) snapshotProviderBaseline(response.items)
sortedProviders.value = sortProvidersByActiveAndPriority(
normalizeProvidersForEditing(response.items)
)
} catch { } catch {
originalProviderPriorityById = new Map()
originalPoolPriorityByProviderId = new Map()
sortedProviders.value = [] sortedProviders.value = []
} }
} }
@@ -995,12 +1140,14 @@ async function loadKeysByFormat() {
data[format] = sortKeysByActiveAndPriority(data[format]) data[format] = sortKeysByActiveAndPriority(data[format])
} }
keysByFormat.value = data keysByFormat.value = data
snapshotKeyBaseline()
const formats = sortApiFormats(Object.keys(data)) const formats = sortApiFormats(Object.keys(data))
if (formats.length > 0 && !formats.includes(activeFormatTab.value)) { if (formats.length > 0 && !formats.includes(activeFormatTab.value)) {
activeFormatTab.value = formats[0] activeFormatTab.value = formats[0]
} }
} catch (err: unknown) { } catch (err: unknown) {
originalKeyPriorityById = new Map()
showError(parseApiError(err, '加载 Key 列表失败'), '错误') showError(parseApiError(err, '加载 Key 列表失败'), '错误')
} finally { } finally {
loadingKeys.value = false loadingKeys.value = false
@@ -1307,35 +1454,43 @@ async function save() {
// 第一步:先保存所有 Provider 和 Key 的优先级数据 // 第一步:先保存所有 Provider 和 Key 的优先级数据
// 确保优先级数据全部到位后,再切换调度模式,避免瞬态不一致 // 确保优先级数据全部到位后,再切换调度模式,避免瞬态不一致
const providerUpdates = sortedProviders.value.map((provider) => { const providerTasks: Array<() => Promise<unknown>> = []
const payload: Parameters<typeof updateProvider>[1] = { sortedProviders.value.forEach((provider) => {
provider_priority: provider.provider_priority, const payload: Parameters<typeof updateProvider>[1] = {}
const currentProviderPriority = Math.max(0, Math.trunc(provider.provider_priority))
const originalProviderPriority = originalProviderPriorityById.get(provider.id)
if (originalProviderPriority == null || originalProviderPriority !== currentProviderPriority) {
payload.provider_priority = currentProviderPriority
} }
// 号池 provider 同时保存 pool_advanced含 global_priority
if (provider.pool_advanced) { const currentPoolPriority = normalizeOptionalPriority(provider.pool_advanced?.global_priority)
const originalPoolPriority = originalPoolPriorityByProviderId.get(provider.id) ?? null
if (currentPoolPriority !== originalPoolPriority) {
payload.pool_advanced = provider.pool_advanced payload.pool_advanced = provider.pool_advanced
? {
...provider.pool_advanced,
global_priority: currentPoolPriority,
}
: null
}
if (Object.keys(payload).length > 0) {
providerTasks.push(() => updateProvider(provider.id, payload))
} }
return updateProvider(provider.id, payload)
}) })
// 收集每个 Key 的按格式优先级(保留原有其他格式的配置) // 收集每个 Key 的按格式优先级(保留原有其他格式的配置)
const keyPriorityByFormatMap = new Map<string, Record<string, number>>() const keyPriorityByFormatMap = buildEditableKeyPriorityMap()
for (const format of Object.keys(keysByFormat.value)) { const keyTasks = Array.from(keyPriorityByFormatMap.entries())
const keys = keysByFormat.value[format].filter((key) => !isPoolManagedKey(key)) .filter(([keyId, priorityByFormat]) => !arePriorityMapsEqual(
keys.forEach((key) => { originalKeyPriorityById.get(keyId),
// 合并原有配置,避免丢失未显示格式的优先级 priorityByFormat,
const existing = keyPriorityByFormatMap.get(key.id) ))
|| normalizePriorityMap(key.global_priority_by_format) .map(([keyId, priorityByFormat]) => () =>
existing[normalizeApiFormatKey(format)] = key.priority
keyPriorityByFormatMap.set(key.id, existing)
})
}
const keyUpdates = Array.from(keyPriorityByFormatMap.entries()).map(([keyId, priorityByFormat]) =>
updateProviderKey(keyId, { global_priority_by_format: priorityByFormat }) updateProviderKey(keyId, { global_priority_by_format: priorityByFormat })
) )
await Promise.all([...providerUpdates, ...keyUpdates]) await runTasksWithConcurrency([...providerTasks, ...keyTasks])
// 第二步:优先级数据全部就绪后,顺序保存调度配置 // 第二步:优先级数据全部就绪后,顺序保存调度配置
// 先保存优先级模式,再保存调度模式,确保 Scheduler 状态完整切换 // 先保存优先级模式,再保存调度模式,确保 Scheduler 状态完整切换
@@ -1350,6 +1505,7 @@ async function save() {
'调度模式cache_affinity(缓存亲和模式) 或 load_balance(负载均衡模式) 或 fixed_order(固定顺序模式)' '调度模式cache_affinity(缓存亲和模式) 或 load_balance(负载均衡模式) 或 fixed_order(固定顺序模式)'
) )
await loadAllProviders()
await loadKeysByFormat() await loadKeysByFormat()
success('优先级已保存') success('优先级已保存')

View File

@@ -83,6 +83,25 @@ def _get_adapter_for_format(api_format: str) -> Any:
return get_adapter_class(api_format) or get_cli_adapter_class(api_format) return get_adapter_class(api_format) or get_cli_adapter_class(api_format)
def _require_test_endpoint_base_url(endpoint: Any) -> str:
"""校验测试链路里的 endpoint.base_url。"""
base_url = getattr(endpoint, "base_url", None)
if not isinstance(base_url, str):
endpoint_id = str(getattr(endpoint, "id", "") or "unknown")
api_format = str(getattr(endpoint, "api_format", "") or "unknown")
raise ValueError(
f"Endpoint {endpoint_id} ({api_format}) has invalid base_url type: "
f"expected str, got {type(base_url).__name__}"
)
normalized = base_url.strip()
if not normalized:
endpoint_id = str(getattr(endpoint, "id", "") or "unknown")
api_format = str(getattr(endpoint, "api_format", "") or "unknown")
raise ValueError(f"Endpoint {endpoint_id} ({api_format}) has empty base_url")
return normalized
def _antigravity_sort_keys(api_keys: list[Any]) -> list[Any]: def _antigravity_sort_keys(api_keys: list[Any]) -> list[Any]:
"""按 tier/可用性对 Antigravity Key 降序排列。 """按 tier/可用性对 Antigravity Key 降序排列。
@@ -917,7 +936,7 @@ async def test_model(
endpoint_config = { endpoint_config = {
"api_key": api_key_value, "api_key": api_key_value,
"api_key_id": api_key.id, # 添加API Key ID用于用量记录 "api_key_id": api_key.id, # 添加API Key ID用于用量记录
"base_url": endpoint.base_url, "base_url": _require_test_endpoint_base_url(endpoint),
"api_format": endpoint.api_format, "api_format": endpoint.api_format,
"extra_headers": extra_headers if extra_headers else None, "extra_headers": extra_headers if extra_headers else None,
"timeout": TimeoutDefaults.HTTP_REQUEST, "timeout": TimeoutDefaults.HTTP_REQUEST,
@@ -1601,7 +1620,7 @@ async def _execute_test_check(
async def _run_check(stream: bool) -> dict[str, Any]: async def _run_check(stream: bool) -> dict[str, Any]:
return await adapter_class.check_endpoint( return await adapter_class.check_endpoint(
None, None,
endpoint.base_url, _require_test_endpoint_base_url(endpoint),
api_key_value, api_key_value,
{ {
**request_payload, **request_payload,

View File

@@ -76,13 +76,17 @@ def _resolve_new_provider_priority(
"""Resolve insertion priority for a newly created provider. """Resolve insertion priority for a newly created provider.
Returns ``(priority, needs_shift)``. When the caller explicitly specifies Returns ``(priority, needs_shift)``. When the caller explicitly specifies
a priority we need to shift existing rows; when auto-topping we simply pick a priority we need to shift existing rows. For auto-top insertion we prefer
``min - 1`` so no shift is required. ``min - 1`` when that still stays non-negative; otherwise we clamp to ``0``
and shift existing rows down to preserve ordering.
""" """
if requested_priority is not None: if requested_priority is not None:
return int(requested_priority), True return int(requested_priority), True
if current_min_priority is not None: if current_min_priority is not None:
return int(current_min_priority) - 1, False current_min = int(current_min_priority)
if current_min <= 0:
return 0, True
return current_min - 1, False
return 100, False return 100, False

View File

@@ -1255,6 +1255,41 @@ class AdminImportConfigAdapter(AdminApiAdapter):
return sorted(endpoint_formats) return sorted(endpoint_formats)
return [] return []
@staticmethod
def _normalize_import_endpoint_payload(
provider_id: str,
ep_data: dict[str, Any],
existing_ep: Any | None = None,
) -> dict[str, Any]:
"""校验并规范化导入的 Endpoint 数据。"""
from src.models.endpoint_models import ProviderEndpointCreate
payload = {
"provider_id": provider_id,
"api_format": ep_data.get("api_format", getattr(existing_ep, "api_format", None)),
"base_url": ep_data.get("base_url", getattr(existing_ep, "base_url", None)),
"custom_path": ep_data.get("custom_path", getattr(existing_ep, "custom_path", None)),
"header_rules": ep_data.get("header_rules", getattr(existing_ep, "header_rules", None)),
"body_rules": ep_data.get("body_rules", getattr(existing_ep, "body_rules", None)),
"max_retries": ep_data.get("max_retries", getattr(existing_ep, "max_retries", 2)),
"config": ep_data.get("config", getattr(existing_ep, "config", None)),
"proxy": ep_data.get("proxy", getattr(existing_ep, "proxy", None)),
"format_acceptance_config": ep_data.get(
"format_acceptance_config",
getattr(existing_ep, "format_acceptance_config", None),
),
}
try:
validated = ProviderEndpointCreate.model_validate(payload)
except Exception as exc:
api_format = payload.get("api_format") or "unknown"
raise InvalidRequestException(
f"导入 Endpoint 失败: provider_id={provider_id}, api_format={api_format}, error={exc}"
) from exc
return validated.model_dump(mode="python")
def _encrypt_provider_config(self, config: dict, crypto_service: Any) -> dict: def _encrypt_provider_config(self, config: dict, crypto_service: Any) -> dict:
"""加密 Provider config 中的 provider_ops credentials""" """加密 Provider config 中的 provider_ops credentials"""
if not config: if not config:
@@ -1579,18 +1614,23 @@ class AdminImportConfigAdapter(AdminApiAdapter):
f"Endpoint '{ep_format}' 已存在于 Provider '{prov_data['name']}'" f"Endpoint '{ep_format}' 已存在于 Provider '{prov_data['name']}'"
) )
elif merge_mode == "overwrite": elif merge_mode == "overwrite":
existing_ep.base_url = ep_data.get("base_url", existing_ep.base_url) normalized_ep = self._normalize_import_endpoint_payload(
existing_ep.header_rules = ep_data.get("header_rules") provider_id,
existing_ep.body_rules = ep_data.get("body_rules") {**ep_data, "api_format": ep_format},
existing_ep.max_retries = ep_data.get("max_retries", 2) existing_ep=existing_ep,
)
existing_ep.base_url = normalized_ep["base_url"]
existing_ep.header_rules = normalized_ep.get("header_rules")
existing_ep.body_rules = normalized_ep.get("body_rules")
existing_ep.max_retries = normalized_ep.get("max_retries", 2)
existing_ep.is_active = ep_data.get("is_active", True) existing_ep.is_active = ep_data.get("is_active", True)
existing_ep.custom_path = ep_data.get("custom_path") existing_ep.custom_path = normalized_ep.get("custom_path")
existing_ep.config = ep_data.get("config") existing_ep.config = normalized_ep.get("config")
existing_ep.format_acceptance_config = ep_data.get( existing_ep.format_acceptance_config = normalized_ep.get(
"format_acceptance_config" "format_acceptance_config"
) )
existing_ep.proxy = self._remap_proxy_node_id( existing_ep.proxy = self._remap_proxy_node_id(
ep_data.get("proxy"), proxy_node_id_map normalized_ep.get("proxy"), proxy_node_id_map
) )
sig = parse_signature_key(ep_format) sig = parse_signature_key(ep_format)
existing_ep.api_format = sig.key # 使用归一化后的格式 existing_ep.api_format = sig.key # 使用归一化后的格式
@@ -1599,6 +1639,10 @@ class AdminImportConfigAdapter(AdminApiAdapter):
existing_ep.updated_at = datetime.now(timezone.utc) existing_ep.updated_at = datetime.now(timezone.utc)
stats["endpoints"]["updated"] += 1 stats["endpoints"]["updated"] += 1
else: else:
normalized_ep = self._normalize_import_endpoint_payload(
provider_id,
{**ep_data, "api_format": ep_format},
)
sig = parse_signature_key(ep_format) sig = parse_signature_key(ep_format)
api_family = sig.api_family.value api_family = sig.api_family.value
endpoint_kind = sig.endpoint_kind.value endpoint_kind = sig.endpoint_kind.value
@@ -1608,16 +1652,16 @@ class AdminImportConfigAdapter(AdminApiAdapter):
api_format=sig.key, # 使用归一化后的格式 api_format=sig.key, # 使用归一化后的格式
api_family=api_family, api_family=api_family,
endpoint_kind=endpoint_kind, endpoint_kind=endpoint_kind,
base_url=ep_data["base_url"], base_url=normalized_ep["base_url"],
header_rules=ep_data.get("header_rules"), header_rules=normalized_ep.get("header_rules"),
body_rules=ep_data.get("body_rules"), body_rules=normalized_ep.get("body_rules"),
max_retries=ep_data.get("max_retries", 2), max_retries=normalized_ep.get("max_retries", 2),
is_active=ep_data.get("is_active", True), is_active=ep_data.get("is_active", True),
custom_path=ep_data.get("custom_path"), custom_path=normalized_ep.get("custom_path"),
config=ep_data.get("config"), config=normalized_ep.get("config"),
format_acceptance_config=ep_data.get("format_acceptance_config"), format_acceptance_config=normalized_ep.get("format_acceptance_config"),
proxy=self._remap_proxy_node_id( proxy=self._remap_proxy_node_id(
ep_data.get("proxy"), proxy_node_id_map normalized_ep.get("proxy"), proxy_node_id_map
), ),
) )
db.add(new_ep) db.add(new_ep)

View File

@@ -9,7 +9,7 @@ from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Query, Request from fastapi import APIRouter, Depends, HTTPException, Query, Request
from sqlalchemy import case, func from sqlalchemy import case, func
from sqlalchemy.orm import Session from sqlalchemy.orm import Session, defer
from src.api.base.admin_adapter import AdminApiAdapter from src.api.base.admin_adapter import AdminApiAdapter
from src.api.base.context import ApiRequestContext from src.api.base.context import ApiRequestContext
@@ -886,8 +886,8 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
count_query = count_query.outerjoin(ApiKey, Usage.api_key_id == ApiKey.id) count_query = count_query.outerjoin(ApiKey, Usage.api_key_id == ApiKey.id)
# -- 构建数据查询(完整 JOIN -- # -- 构建数据查询(完整 JOIN --
usage_model_version = Usage.request_metadata["model_version"].as_string().label( usage_model_version = (
"model_version" Usage.request_metadata["model_version"].as_string().label("model_version")
) )
query = ( query = (
@@ -1246,16 +1246,77 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
usage_id: str usage_id: str
include_bodies: bool = True include_bodies: bool = True
def _build_usage_detail_query(self, db: Session) -> Any:
query = db.query(
Usage,
case(
(
(Usage.request_body.isnot(None)) | (Usage.request_body_compressed.isnot(None)),
True,
),
else_=False,
).label("has_request_body"),
case(
(
(Usage.provider_request_body.isnot(None))
| (Usage.provider_request_body_compressed.isnot(None)),
True,
),
else_=False,
).label("has_provider_request_body"),
case(
(
(Usage.response_body.isnot(None))
| (Usage.response_body_compressed.isnot(None)),
True,
),
else_=False,
).label("has_response_body"),
case(
(
(Usage.client_response_body.isnot(None))
| (Usage.client_response_body_compressed.isnot(None)),
True,
),
else_=False,
).label("has_client_response_body"),
)
if not self.include_bodies:
query = query.options(
defer(Usage.request_body),
defer(Usage.provider_request_body),
defer(Usage.response_body),
defer(Usage.client_response_body),
defer(Usage.request_body_compressed),
defer(Usage.provider_request_body_compressed),
defer(Usage.response_body_compressed),
defer(Usage.client_response_body_compressed),
)
return query
def _load_usage_detail_row(self, db: Session) -> Any:
usage_row = self._build_usage_detail_query(db).filter(Usage.id == self.usage_id).first()
if usage_row:
return usage_row
return self._build_usage_detail_query(db).filter(Usage.request_id == self.usage_id).first()
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override] async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
db = context.db db = context.db
# 先通过主键 id 查找,如果找不到再尝试通过 request_id 查找 # 先通过主键 id 查找,如果找不到再尝试通过 request_id 查找
usage_record = db.query(Usage).filter(Usage.id == self.usage_id).first() usage_row = self._load_usage_detail_row(db)
if not usage_record: if not usage_row:
# 兼容通过 request_id 查找(用于异步任务等场景)
usage_record = db.query(Usage).filter(Usage.request_id == self.usage_id).first()
if not usage_record:
raise HTTPException(status_code=404, detail="Usage record not found") raise HTTPException(status_code=404, detail="Usage record not found")
(
usage_record,
has_request_body,
has_provider_request_body,
has_response_body,
has_client_response_body,
) = usage_row
user = db.query(User).filter(User.id == usage_record.user_id).first() user = db.query(User).filter(User.id == usage_record.user_id).first()
api_key = db.query(ApiKey).filter(ApiKey.id == usage_record.api_key_id).first() api_key = db.query(ApiKey).filter(ApiKey.id == usage_record.api_key_id).first()
@@ -1270,23 +1331,6 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
# 提取视频/图像/音频计费信息 # 提取视频/图像/音频计费信息
video_billing_info = self._extract_video_billing_info(usage_record) video_billing_info = self._extract_video_billing_info(usage_record)
has_request_body = bool(
usage_record.request_body is not None
or usage_record.request_body_compressed is not None
)
has_provider_request_body = bool(
usage_record.provider_request_body is not None
or usage_record.provider_request_body_compressed is not None
)
has_response_body = bool(
usage_record.response_body is not None
or usage_record.response_body_compressed is not None
)
has_client_response_body = bool(
usage_record.client_response_body is not None
or usage_record.client_response_body_compressed is not None
)
request_body = usage_record.get_request_body() if self.include_bodies else None request_body = usage_record.get_request_body() if self.include_bodies else None
provider_request_body = ( provider_request_body = (
usage_record.get_provider_request_body() if self.include_bodies else None usage_record.get_provider_request_body() if self.include_bodies else None

View File

@@ -322,20 +322,15 @@ class HandlerAdapterBase(ApiAdapter):
return build_test_request_body(cls.FORMAT_ID, request_data) return build_test_request_body(cls.FORMAT_ID, request_data)
@staticmethod @staticmethod
def _normalize_test_base_url(base_url: Any) -> str: def _validate_test_base_url(base_url: Any) -> str:
"""归一化 test-model 场景传入的 base_url。""" """校验 test-model 场景传入的 base_url。"""
if isinstance(base_url, str): if not isinstance(base_url, str):
normalized = base_url.strip() raise TypeError(f"base_url must be a non-empty string, got {type(base_url).__name__}")
if normalized:
return normalized
elif isinstance(base_url, dict):
for key in ("base_url", "url"):
value = base_url.get(key)
if isinstance(value, str) and value.strip():
logger.debug("[check_endpoint] 兼容字典形式的 base_url 输入: key={}", key)
return value.strip()
raise TypeError("base_url must be a non-empty string or a dict containing 'base_url'/'url'") normalized = base_url.strip()
if not normalized:
raise ValueError("base_url must be a non-empty string")
return normalized
@classmethod @classmethod
async def check_endpoint( async def check_endpoint(
@@ -375,7 +370,7 @@ class HandlerAdapterBase(ApiAdapter):
from src.core.api_format.headers import HeaderBuilder from src.core.api_format.headers import HeaderBuilder
from src.core.provider_types import ProviderType from src.core.provider_types import ProviderType
normalized_base_url = cls._normalize_test_base_url(base_url) validated_base_url = cls._validate_test_base_url(base_url)
is_antigravity = provider_type == ProviderType.ANTIGRAVITY is_antigravity = provider_type == ProviderType.ANTIGRAVITY
is_gemini_cli = provider_type == ProviderType.GEMINI_CLI is_gemini_cli = provider_type == ProviderType.GEMINI_CLI
is_vertex = provider_type == ProviderType.VERTEX_AI is_vertex = provider_type == ProviderType.VERTEX_AI
@@ -393,9 +388,9 @@ class HandlerAdapterBase(ApiAdapter):
_kiro_cfg = KiroAuthConfig.from_dict(decrypted_auth_config or {}) _kiro_cfg = KiroAuthConfig.from_dict(decrypted_auth_config or {})
region = _kiro_cfg.effective_api_region() region = _kiro_cfg.effective_api_region()
effective_base_url = ( effective_base_url = (
normalized_base_url.replace("{region}", region) validated_base_url.replace("{region}", region)
if "{region}" in normalized_base_url if "{region}" in validated_base_url
else normalized_base_url else validated_base_url
) )
url = f"{str(effective_base_url).rstrip('/')}{KIRO_GENERATE_ASSISTANT_PATH}" url = f"{str(effective_base_url).rstrip('/')}{KIRO_GENERATE_ASSISTANT_PATH}"
elif is_antigravity: elif is_antigravity:
@@ -408,13 +403,13 @@ class HandlerAdapterBase(ApiAdapter):
) )
ordered_urls = url_availability.get_ordered_urls(prefer_daily=True) ordered_urls = url_availability.get_ordered_urls(prefer_daily=True)
effective_base_url = ordered_urls[0] if ordered_urls else base_url effective_base_url = ordered_urls[0] if ordered_urls else validated_base_url
path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent") path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent")
url = f"{str(effective_base_url).rstrip('/')}{path}" url = f"{str(effective_base_url).rstrip('/')}{path}"
elif is_gemini_cli: elif is_gemini_cli:
from src.services.provider.adapters.gemini_cli.constants import V1INTERNAL_PATH_TEMPLATE from src.services.provider.adapters.gemini_cli.constants import V1INTERNAL_PATH_TEMPLATE
effective_base_url = base_url effective_base_url = validated_base_url
path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent") path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent")
url = f"{str(effective_base_url).rstrip('/')}{path}" url = f"{str(effective_base_url).rstrip('/')}{path}"
elif is_vertex and provider_endpoint is not None and provider_api_key is not None: elif is_vertex and provider_endpoint is not None and provider_api_key is not None:
@@ -441,7 +436,7 @@ class HandlerAdapterBase(ApiAdapter):
) )
else: else:
url = cls.build_endpoint_url( url = cls.build_endpoint_url(
normalized_base_url, validated_base_url,
request_data, request_data,
model_name, model_name,
provider_type=provider_type, provider_type=provider_type,
@@ -449,7 +444,7 @@ class HandlerAdapterBase(ApiAdapter):
# ---- Headers ---- # ---- Headers ----
cli_extra = cls.get_cli_extra_headers( cli_extra = cls.get_cli_extra_headers(
base_url=normalized_base_url, base_url=validated_base_url,
provider_type=provider_type, provider_type=provider_type,
) )
merged_extra = dict(extra_headers) if extra_headers else {} merged_extra = dict(extra_headers) if extra_headers else {}
@@ -507,7 +502,7 @@ class HandlerAdapterBase(ApiAdapter):
# ---- Body ---- # ---- Body ----
body = cls.build_request_body( body = cls.build_request_body(
request_data, request_data,
base_url=normalized_base_url, base_url=validated_base_url,
provider_type=provider_type, provider_type=provider_type,
) )

View File

@@ -7,7 +7,7 @@ from typing import Any
import pytest import pytest
from src.api.admin.usage.routes import AdminUsageRecordsAdapter from src.api.admin.usage.routes import AdminUsageDetailAdapter, AdminUsageRecordsAdapter
class _FakeQuery: class _FakeQuery:
@@ -16,9 +16,11 @@ class _FakeQuery:
*, *,
scalar_result: int | None = None, scalar_result: int | None = None,
all_result: list[Any] | None = None, all_result: list[Any] | None = None,
first_result: Any = None,
) -> None: ) -> None:
self.scalar_result = scalar_result self.scalar_result = scalar_result
self.all_result = all_result or [] self.all_result = all_result or []
self.first_result = first_result
self.options_args: tuple[Any, ...] = () self.options_args: tuple[Any, ...] = ()
def outerjoin(self, *args: Any, **kwargs: Any) -> _FakeQuery: def outerjoin(self, *args: Any, **kwargs: Any) -> _FakeQuery:
@@ -49,6 +51,9 @@ class _FakeQuery:
def all(self) -> list[Any]: def all(self) -> list[Any]:
return self.all_result return self.all_result
def first(self) -> Any:
return self.first_result
class _FakeDb: class _FakeDb:
def __init__(self, queries: list[_FakeQuery]) -> None: def __init__(self, queries: list[_FakeQuery]) -> None:
@@ -141,3 +146,130 @@ async def test_admin_usage_records_returns_model_version_without_request_metadat
usage_load_only = data_query.options_args[0] usage_load_only = data_query.options_args[0]
usage_paths = {str(option.path) for option in usage_load_only.context} usage_paths = {str(option.path) for option in usage_load_only.context}
assert "ORM Path[Mapper[Usage(usage)] -> Usage.request_metadata]" not in usage_paths assert "ORM Path[Mapper[Usage(usage)] -> Usage.request_metadata]" not in usage_paths
@pytest.mark.asyncio
async def test_admin_usage_detail_defers_large_body_columns_when_bodies_excluded(
monkeypatch: pytest.MonkeyPatch,
) -> None:
async def _fake_get_tiered_pricing_info(
self: AdminUsageDetailAdapter,
db: Any,
usage_record: Any,
) -> None:
return None
monkeypatch.setattr(
AdminUsageDetailAdapter,
"_get_tiered_pricing_info",
_fake_get_tiered_pricing_info,
)
monkeypatch.setattr(
AdminUsageDetailAdapter,
"_extract_video_billing_info",
lambda self, usage_record: None,
)
class _UsageRecord:
id = "usage-1"
request_id = "req-1"
user_id = "user-1"
api_key_id = "key-1"
provider_name = "openai"
api_format = "openai:cli"
model = "gpt-5.4"
target_model = None
input_tokens = 10
output_tokens = 20
total_tokens = 30
cache_creation_input_tokens = 0
cache_read_input_tokens = 0
cache_creation_input_tokens_5m = 0
cache_creation_input_tokens_1h = 0
input_cost_usd = Decimal("0.001")
output_cost_usd = Decimal("0.002")
total_cost_usd = Decimal("0.003")
cache_creation_cost_usd = Decimal("0")
cache_read_cost_usd = Decimal("0")
request_cost_usd = Decimal("0")
input_price_per_1m = Decimal("0.1")
output_price_per_1m = Decimal("0.2")
cache_creation_price_per_1m = None
cache_read_price_per_1m = None
price_per_request = None
request_type = "chat"
is_stream = True
status_code = 200
error_message = None
status = "completed"
response_time_ms = 1200
first_byte_time_ms = 200
created_at = datetime(2026, 3, 12, 7, 0, tzinfo=timezone.utc)
request_headers = {"x-test": "1"}
provider_request_headers = {"authorization": "***"}
response_headers = {"content-type": "text/event-stream"}
client_response_headers = {"content-type": "text/event-stream"}
request_metadata = {"trace_id": "trace-1"}
def get_request_body(self) -> Any:
raise AssertionError("request body should not be loaded")
def get_provider_request_body(self) -> Any:
raise AssertionError("provider request body should not be loaded")
def get_response_body(self) -> Any:
raise AssertionError("response body should not be loaded")
def get_client_response_body(self) -> Any:
raise AssertionError("client response body should not be loaded")
class _ApiKeyRecord:
id = "key-1"
name = "Primary"
def get_display_key(self) -> str:
return "sk-test"
usage_query = _FakeQuery(
first_result=(_UsageRecord(), True, True, True, True),
)
user_query = _FakeQuery(
first_result=SimpleNamespace(id="user-1", username="tester", email="u@example.com"),
)
api_key_query = _FakeQuery(first_result=_ApiKeyRecord())
db = _FakeDb([usage_query, user_query, api_key_query])
context = SimpleNamespace(
db=db,
user=SimpleNamespace(id="admin-1"),
add_audit_metadata=lambda **_: None,
)
adapter = AdminUsageDetailAdapter(usage_id="usage-1", include_bodies=False)
result = await adapter.handle(context) # type: ignore[arg-type]
assert result["request_body"] is None
assert result["provider_request_body"] is None
assert result["response_body"] is None
assert result["client_response_body"] is None
assert result["has_request_body"] is True
assert result["has_provider_request_body"] is True
assert result["has_response_body"] is True
assert result["has_client_response_body"] is True
deferred_paths = {
str(context.path)
for option in usage_query.options_args
for context in getattr(option, "context", ())
}
assert "ORM Path[Mapper[Usage(usage)] -> Usage.request_body]" in deferred_paths
assert "ORM Path[Mapper[Usage(usage)] -> Usage.provider_request_body]" in deferred_paths
assert "ORM Path[Mapper[Usage(usage)] -> Usage.response_body]" in deferred_paths
assert "ORM Path[Mapper[Usage(usage)] -> Usage.client_response_body]" in deferred_paths
assert "ORM Path[Mapper[Usage(usage)] -> Usage.request_body_compressed]" in deferred_paths
assert (
"ORM Path[Mapper[Usage(usage)] -> Usage.provider_request_body_compressed]" in deferred_paths
)
assert "ORM Path[Mapper[Usage(usage)] -> Usage.response_body_compressed]" in deferred_paths
assert (
"ORM Path[Mapper[Usage(usage)] -> Usage.client_response_body_compressed]" in deferred_paths
)

View File

@@ -48,6 +48,14 @@ def test_new_provider_priority_defaults_to_current_top() -> None:
assert needs_shift is False assert needs_shift is False
def test_new_provider_priority_clamps_at_zero_when_already_topmost() -> None:
priority, needs_shift = _resolve_new_provider_priority(
current_min_priority=0, requested_priority=None
)
assert priority == 0
assert needs_shift is True
def test_new_provider_priority_defaults_to_100_when_empty() -> None: def test_new_provider_priority_defaults_to_100_when_empty() -> None:
priority, needs_shift = _resolve_new_provider_priority( priority, needs_shift = _resolve_new_provider_priority(
current_min_priority=None, requested_priority=None current_min_priority=None, requested_priority=None

View File

@@ -1,4 +1,7 @@
import pytest
from src.api.admin.system import AdminExportConfigAdapter, AdminImportConfigAdapter from src.api.admin.system import AdminExportConfigAdapter, AdminImportConfigAdapter
from src.core.exceptions import InvalidRequestException
from src.services.provider_ops.types import SENSITIVE_CREDENTIAL_FIELDS from src.services.provider_ops.types import SENSITIVE_CREDENTIAL_FIELDS
@@ -114,3 +117,27 @@ def test_import_provider_config_encrypts_refresh_token() -> None:
assert result["provider_ops"]["connector"]["credentials"]["refresh_token"] == "enc:rt-1" assert result["provider_ops"]["connector"]["credentials"]["refresh_token"] == "enc:rt-1"
assert result["provider_ops"]["connector"]["credentials"]["api_key"] == "enc:key-1" assert result["provider_ops"]["connector"]["credentials"]["api_key"] == "enc:key-1"
assert config["provider_ops"]["connector"]["credentials"]["refresh_token"] == "rt-1" assert config["provider_ops"]["connector"]["credentials"]["refresh_token"] == "rt-1"
def test_import_endpoint_payload_rejects_dict_base_url() -> None:
with pytest.raises(InvalidRequestException, match="导入 Endpoint 失败"):
AdminImportConfigAdapter._normalize_import_endpoint_payload(
"provider-1",
{
"api_format": "claude:chat",
"base_url": {"url": "https://api.anthropic.com"},
},
)
def test_import_endpoint_payload_normalizes_base_url() -> None:
result = AdminImportConfigAdapter._normalize_import_endpoint_payload(
"provider-1",
{
"api_format": "claude:chat",
"base_url": "https://api.anthropic.com/",
},
)
assert result["base_url"] == "https://api.anthropic.com"
assert result["api_format"] == "claude:chat"

View File

@@ -4,21 +4,11 @@ from src.api.handlers.claude.adapter import ClaudeChatAdapter
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_check_endpoint_accepts_base_url_dict(monkeypatch: pytest.MonkeyPatch) -> None: async def test_check_endpoint_rejects_base_url_dict() -> None:
captured: dict[str, object] = {} with pytest.raises(TypeError, match="base_url must be a non-empty string"):
async def fake_run_endpoint_check(**kwargs):
captured.update(kwargs)
return {"status_code": 200, "headers": {}, "response_time_ms": 1, "request_id": "test"}
monkeypatch.setattr(
"src.api.handlers.base.endpoint_checker.run_endpoint_check",
fake_run_endpoint_check,
)
await ClaudeChatAdapter.check_endpoint( await ClaudeChatAdapter.check_endpoint(
client=None, client=None, # type: ignore[arg-type]
base_url={"base_url": "https://api.anthropic.com"}, base_url={"base_url": "https://api.anthropic.com"}, # type: ignore[arg-type]
api_key="test-key", api_key="test-key",
request_data={ request_data={
"model": "claude-sonnet-4-5-20250929", "model": "claude-sonnet-4-5-20250929",
@@ -28,35 +18,8 @@ async def test_check_endpoint_accepts_base_url_dict(monkeypatch: pytest.MonkeyPa
}, },
) )
assert captured["url"] == "https://api.anthropic.com/v1/messages"
assert isinstance(captured["json_body"], dict)
def test_validate_test_base_url_trims_whitespace() -> None:
@pytest.mark.asyncio assert ClaudeChatAdapter._validate_test_base_url(" https://api.anthropic.com/v1 ") == (
async def test_check_endpoint_accepts_url_key_in_base_url_dict( "https://api.anthropic.com/v1"
monkeypatch: pytest.MonkeyPatch,
) -> None:
captured: dict[str, object] = {}
async def fake_run_endpoint_check(**kwargs):
captured.update(kwargs)
return {"status_code": 200, "headers": {}, "response_time_ms": 1, "request_id": "test"}
monkeypatch.setattr(
"src.api.handlers.base.endpoint_checker.run_endpoint_check",
fake_run_endpoint_check,
) )
await ClaudeChatAdapter.check_endpoint(
client=None,
base_url={"url": "https://api.anthropic.com/v1"},
api_key="test-key",
request_data={
"model": "claude-sonnet-4-5-20250929",
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 32,
"stream": False,
},
)
assert captured["url"] == "https://api.anthropic.com/v1/messages"

View File

@@ -9,6 +9,7 @@ from src.api.admin.provider_query import (
_build_test_attempts_from_candidate_keys, _build_test_attempts_from_candidate_keys,
_filter_test_candidates_by_endpoint, _filter_test_candidates_by_endpoint,
_flatten_test_candidates_for_concurrency, _flatten_test_candidates_for_concurrency,
_require_test_endpoint_base_url,
_resolve_test_effective_model, _resolve_test_effective_model,
) )
from src.services.scheduling.schemas import PoolCandidate from src.services.scheduling.schemas import PoolCandidate
@@ -148,3 +149,18 @@ def test_test_model_failover_request_validates_concurrency_range() -> None:
model_name="gpt-4o-mini", model_name="gpt-4o-mini",
concurrency=0, concurrency=0,
) )
def test_require_test_endpoint_base_url_rejects_non_string() -> None:
endpoint = SimpleNamespace(id="ep-bad", api_format="claude:chat", base_url={"url": "https://x"})
with pytest.raises(ValueError, match="invalid base_url type"):
_require_test_endpoint_base_url(endpoint)
def test_require_test_endpoint_base_url_trims_whitespace() -> None:
endpoint = SimpleNamespace(
id="ep-ok", api_format="claude:chat", base_url=" https://api.anthropic.com "
)
assert _require_test_endpoint_base_url(endpoint) == "https://api.anthropic.com"