mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
@@ -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('优先级已保存')
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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"
|
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
Reference in New Issue
Block a user