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 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 优先级编辑状态
const editingKeyPriority = ref<Record<string, string | null>>({}) // format -> keyId
@@ -642,6 +648,77 @@ function normalizePriorityMap(
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(
value: Record<string, unknown> | null | undefined
): Record<string, number> | null {
@@ -657,6 +734,69 @@ function normalizeRateMultipliers(
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 map = new Map<string, ProviderWithEndpointsSummary>()
sortedProviders.value.forEach((provider) => {
@@ -857,8 +997,13 @@ watch(internalOpen, async (open) => {
async function loadAllProviders() {
try {
const response = await getProvidersSummary({ page: 1, page_size: 9999 })
sortedProviders.value = sortProvidersByActiveAndPriority(response.items)
snapshotProviderBaseline(response.items)
sortedProviders.value = sortProvidersByActiveAndPriority(
normalizeProvidersForEditing(response.items)
)
} catch {
originalProviderPriorityById = new Map()
originalPoolPriorityByProviderId = new Map()
sortedProviders.value = []
}
}
@@ -995,12 +1140,14 @@ async function loadKeysByFormat() {
data[format] = sortKeysByActiveAndPriority(data[format])
}
keysByFormat.value = data
snapshotKeyBaseline()
const formats = sortApiFormats(Object.keys(data))
if (formats.length > 0 && !formats.includes(activeFormatTab.value)) {
activeFormatTab.value = formats[0]
}
} catch (err: unknown) {
originalKeyPriorityById = new Map()
showError(parseApiError(err, '加载 Key 列表失败'), '错误')
} finally {
loadingKeys.value = false
@@ -1307,35 +1454,43 @@ async function save() {
// 第一步:先保存所有 Provider 和 Key 的优先级数据
// 确保优先级数据全部到位后,再切换调度模式,避免瞬态不一致
const providerUpdates = sortedProviders.value.map((provider) => {
const payload: Parameters<typeof updateProvider>[1] = {
provider_priority: provider.provider_priority,
const providerTasks: Array<() => Promise<unknown>> = []
sortedProviders.value.forEach((provider) => {
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
? {
...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 的按格式优先级(保留原有其他格式的配置)
const keyPriorityByFormatMap = new Map<string, Record<string, number>>()
for (const format of Object.keys(keysByFormat.value)) {
const keys = keysByFormat.value[format].filter((key) => !isPoolManagedKey(key))
keys.forEach((key) => {
// 合并原有配置,避免丢失未显示格式的优先级
const existing = keyPriorityByFormatMap.get(key.id)
|| normalizePriorityMap(key.global_priority_by_format)
existing[normalizeApiFormatKey(format)] = key.priority
keyPriorityByFormatMap.set(key.id, existing)
})
}
const keyPriorityByFormatMap = buildEditableKeyPriorityMap()
const keyTasks = Array.from(keyPriorityByFormatMap.entries())
.filter(([keyId, priorityByFormat]) => !arePriorityMapsEqual(
originalKeyPriorityById.get(keyId),
priorityByFormat,
))
.map(([keyId, priorityByFormat]) => () =>
updateProviderKey(keyId, { global_priority_by_format: priorityByFormat })
)
const keyUpdates = Array.from(keyPriorityByFormatMap.entries()).map(([keyId, priorityByFormat]) =>
updateProviderKey(keyId, { global_priority_by_format: priorityByFormat })
)
await Promise.all([...providerUpdates, ...keyUpdates])
await runTasksWithConcurrency([...providerTasks, ...keyTasks])
// 第二步:优先级数据全部就绪后,顺序保存调度配置
// 先保存优先级模式,再保存调度模式,确保 Scheduler 状态完整切换
@@ -1350,6 +1505,7 @@ async function save() {
'调度模式cache_affinity(缓存亲和模式) 或 load_balance(负载均衡模式) 或 fixed_order(固定顺序模式)'
)
await loadAllProviders()
await loadKeysByFormat()
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)
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]:
"""按 tier/可用性对 Antigravity Key 降序排列。
@@ -917,7 +936,7 @@ async def test_model(
endpoint_config = {
"api_key": api_key_value,
"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,
"extra_headers": extra_headers if extra_headers else None,
"timeout": TimeoutDefaults.HTTP_REQUEST,
@@ -1601,7 +1620,7 @@ async def _execute_test_check(
async def _run_check(stream: bool) -> dict[str, Any]:
return await adapter_class.check_endpoint(
None,
endpoint.base_url,
_require_test_endpoint_base_url(endpoint),
api_key_value,
{
**request_payload,

View File

@@ -76,13 +76,17 @@ def _resolve_new_provider_priority(
"""Resolve insertion priority for a newly created provider.
Returns ``(priority, needs_shift)``. When the caller explicitly specifies
a priority we need to shift existing rows; when auto-topping we simply pick
``min - 1`` so no shift is required.
a priority we need to shift existing rows. For auto-top insertion we prefer
``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:
return int(requested_priority), True
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

View File

@@ -1255,6 +1255,41 @@ class AdminImportConfigAdapter(AdminApiAdapter):
return sorted(endpoint_formats)
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:
"""加密 Provider config 中的 provider_ops credentials"""
if not config:
@@ -1579,18 +1614,23 @@ class AdminImportConfigAdapter(AdminApiAdapter):
f"Endpoint '{ep_format}' 已存在于 Provider '{prov_data['name']}'"
)
elif merge_mode == "overwrite":
existing_ep.base_url = ep_data.get("base_url", existing_ep.base_url)
existing_ep.header_rules = ep_data.get("header_rules")
existing_ep.body_rules = ep_data.get("body_rules")
existing_ep.max_retries = ep_data.get("max_retries", 2)
normalized_ep = self._normalize_import_endpoint_payload(
provider_id,
{**ep_data, "api_format": ep_format},
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.custom_path = ep_data.get("custom_path")
existing_ep.config = ep_data.get("config")
existing_ep.format_acceptance_config = ep_data.get(
existing_ep.custom_path = normalized_ep.get("custom_path")
existing_ep.config = normalized_ep.get("config")
existing_ep.format_acceptance_config = normalized_ep.get(
"format_acceptance_config"
)
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)
existing_ep.api_format = sig.key # 使用归一化后的格式
@@ -1599,6 +1639,10 @@ class AdminImportConfigAdapter(AdminApiAdapter):
existing_ep.updated_at = datetime.now(timezone.utc)
stats["endpoints"]["updated"] += 1
else:
normalized_ep = self._normalize_import_endpoint_payload(
provider_id,
{**ep_data, "api_format": ep_format},
)
sig = parse_signature_key(ep_format)
api_family = sig.api_family.value
endpoint_kind = sig.endpoint_kind.value
@@ -1608,16 +1652,16 @@ class AdminImportConfigAdapter(AdminApiAdapter):
api_format=sig.key, # 使用归一化后的格式
api_family=api_family,
endpoint_kind=endpoint_kind,
base_url=ep_data["base_url"],
header_rules=ep_data.get("header_rules"),
body_rules=ep_data.get("body_rules"),
max_retries=ep_data.get("max_retries", 2),
base_url=normalized_ep["base_url"],
header_rules=normalized_ep.get("header_rules"),
body_rules=normalized_ep.get("body_rules"),
max_retries=normalized_ep.get("max_retries", 2),
is_active=ep_data.get("is_active", True),
custom_path=ep_data.get("custom_path"),
config=ep_data.get("config"),
format_acceptance_config=ep_data.get("format_acceptance_config"),
custom_path=normalized_ep.get("custom_path"),
config=normalized_ep.get("config"),
format_acceptance_config=normalized_ep.get("format_acceptance_config"),
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)

View File

@@ -9,7 +9,7 @@ from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Query, Request
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.context import ApiRequestContext
@@ -886,8 +886,8 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
count_query = count_query.outerjoin(ApiKey, Usage.api_key_id == ApiKey.id)
# -- 构建数据查询(完整 JOIN --
usage_model_version = Usage.request_metadata["model_version"].as_string().label(
"model_version"
usage_model_version = (
Usage.request_metadata["model_version"].as_string().label("model_version")
)
query = (
@@ -1246,16 +1246,77 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
usage_id: str
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]
db = context.db
# 先通过主键 id 查找,如果找不到再尝试通过 request_id 查找
usage_record = db.query(Usage).filter(Usage.id == self.usage_id).first()
if not usage_record:
# 兼容通过 request_id 查找(用于异步任务等场景)
usage_record = db.query(Usage).filter(Usage.request_id == self.usage_id).first()
if not usage_record:
usage_row = self._load_usage_detail_row(db)
if not usage_row:
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()
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)
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
provider_request_body = (
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)
@staticmethod
def _normalize_test_base_url(base_url: Any) -> str:
"""归一化 test-model 场景传入的 base_url。"""
if isinstance(base_url, str):
normalized = base_url.strip()
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()
def _validate_test_base_url(base_url: Any) -> str:
"""校验 test-model 场景传入的 base_url。"""
if not isinstance(base_url, str):
raise TypeError(f"base_url must be a non-empty string, got {type(base_url).__name__}")
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
async def check_endpoint(
@@ -375,7 +370,7 @@ class HandlerAdapterBase(ApiAdapter):
from src.core.api_format.headers import HeaderBuilder
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_gemini_cli = provider_type == ProviderType.GEMINI_CLI
is_vertex = provider_type == ProviderType.VERTEX_AI
@@ -393,9 +388,9 @@ class HandlerAdapterBase(ApiAdapter):
_kiro_cfg = KiroAuthConfig.from_dict(decrypted_auth_config or {})
region = _kiro_cfg.effective_api_region()
effective_base_url = (
normalized_base_url.replace("{region}", region)
if "{region}" in normalized_base_url
else normalized_base_url
validated_base_url.replace("{region}", region)
if "{region}" in validated_base_url
else validated_base_url
)
url = f"{str(effective_base_url).rstrip('/')}{KIRO_GENERATE_ASSISTANT_PATH}"
elif is_antigravity:
@@ -408,13 +403,13 @@ class HandlerAdapterBase(ApiAdapter):
)
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")
url = f"{str(effective_base_url).rstrip('/')}{path}"
elif is_gemini_cli:
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")
url = f"{str(effective_base_url).rstrip('/')}{path}"
elif is_vertex and provider_endpoint is not None and provider_api_key is not None:
@@ -441,7 +436,7 @@ class HandlerAdapterBase(ApiAdapter):
)
else:
url = cls.build_endpoint_url(
normalized_base_url,
validated_base_url,
request_data,
model_name,
provider_type=provider_type,
@@ -449,7 +444,7 @@ class HandlerAdapterBase(ApiAdapter):
# ---- Headers ----
cli_extra = cls.get_cli_extra_headers(
base_url=normalized_base_url,
base_url=validated_base_url,
provider_type=provider_type,
)
merged_extra = dict(extra_headers) if extra_headers else {}
@@ -507,7 +502,7 @@ class HandlerAdapterBase(ApiAdapter):
# ---- Body ----
body = cls.build_request_body(
request_data,
base_url=normalized_base_url,
base_url=validated_base_url,
provider_type=provider_type,
)

View File

@@ -7,7 +7,7 @@ from typing import Any
import pytest
from src.api.admin.usage.routes import AdminUsageRecordsAdapter
from src.api.admin.usage.routes import AdminUsageDetailAdapter, AdminUsageRecordsAdapter
class _FakeQuery:
@@ -16,9 +16,11 @@ class _FakeQuery:
*,
scalar_result: int | None = None,
all_result: list[Any] | None = None,
first_result: Any = None,
) -> None:
self.scalar_result = scalar_result
self.all_result = all_result or []
self.first_result = first_result
self.options_args: tuple[Any, ...] = ()
def outerjoin(self, *args: Any, **kwargs: Any) -> _FakeQuery:
@@ -49,6 +51,9 @@ class _FakeQuery:
def all(self) -> list[Any]:
return self.all_result
def first(self) -> Any:
return self.first_result
class _FakeDb:
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_paths = {str(option.path) for option in usage_load_only.context}
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
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:
priority, needs_shift = _resolve_new_provider_priority(
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.core.exceptions import InvalidRequestException
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"]["api_key"] == "enc:key-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,59 +4,22 @@ from src.api.handlers.claude.adapter import ClaudeChatAdapter
@pytest.mark.asyncio
async def test_check_endpoint_accepts_base_url_dict(monkeypatch: pytest.MonkeyPatch) -> None:
captured: dict[str, object] = {}
async def test_check_endpoint_rejects_base_url_dict() -> None:
with pytest.raises(TypeError, match="base_url must be a non-empty string"):
await ClaudeChatAdapter.check_endpoint(
client=None, # type: ignore[arg-type]
base_url={"base_url": "https://api.anthropic.com"}, # type: ignore[arg-type]
api_key="test-key",
request_data={
"model": "claude-sonnet-4-5-20250929",
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 32,
"stream": False,
},
)
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,
def test_validate_test_base_url_trims_whitespace() -> None:
assert ClaudeChatAdapter._validate_test_base_url(" https://api.anthropic.com/v1 ") == (
"https://api.anthropic.com/v1"
)
await ClaudeChatAdapter.check_endpoint(
client=None,
base_url={"base_url": "https://api.anthropic.com"},
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"
assert isinstance(captured["json_body"], dict)
@pytest.mark.asyncio
async def test_check_endpoint_accepts_url_key_in_base_url_dict(
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,
_filter_test_candidates_by_endpoint,
_flatten_test_candidates_for_concurrency,
_require_test_endpoint_base_url,
_resolve_test_effective_model,
)
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",
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"