feat(test): 模型测试支持故障转移,展示每次尝试详情

- 后端新增 test-model-failover 接口,支持 global/direct 两种测试模式
- 利用 FailoverEngine 遍历候选并记录每次尝试的状态、延迟、错误等详情
- 前端 ModelsTab/ModelMappingTab 切换到新接口,移除格式选择下拉菜单
- 新增 TestResultDialog 组件,失败时展示候选尝试详情表格
- 简化组件 props 传递,移除不再需要的 endpoints/mappingPreview 依赖
This commit is contained in:
fawney19
2026-03-01 03:37:41 +08:00
parent fea3d183bf
commit b61fc4eb6b
7 changed files with 786 additions and 242 deletions

View File

@@ -7,7 +7,8 @@ from __future__ import annotations
import asyncio
import json
from typing import Any
import time
from typing import TYPE_CHECKING, Any
from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel
@@ -39,6 +40,9 @@ from src.services.provider.oauth_token import resolve_oauth_access_token
from src.services.proxy_node.resolver import resolve_effective_proxy
from src.utils.auth_utils import get_current_user
if TYPE_CHECKING:
from src.services.scheduling.schemas import ProviderCandidate
router = APIRouter(prefix="/api/admin/provider-query", tags=["Provider Query"])
@@ -194,6 +198,46 @@ class TestModelRequest(BaseModel):
api_format: str | None = None # 指定使用的API格式如果不指定则使用端点的默认格式
class TestModelFailoverRequest(BaseModel):
"""带故障转移的模型测试请求"""
provider_id: str
mode: str # "global" = 模拟外部请求(用全局模型名), "direct" = 直接测试(用provider_model_name)
model_name: str # global 模式传 global_model_name, direct 模式传 provider_model_name
api_format: str | None = None # 指定 API 格式endpoint signature
message: str | None = "Hello"
class TestAttemptDetail(BaseModel):
"""单次测试尝试的详情"""
candidate_index: int
endpoint_api_format: str
endpoint_base_url: str
key_name: str | None = None
key_id: str
auth_type: str
effective_model: str | None = None # 实际发送的模型名(映射后)
status: str # "success" | "failed" | "skipped"
skip_reason: str | None = None
error_message: str | None = None
status_code: int | None = None
latency_ms: int | None = None
class TestModelFailoverResponse(BaseModel):
"""带故障转移的模型测试响应"""
success: bool
model: str
provider: dict[str, str]
attempts: list[TestAttemptDetail]
total_candidates: int
total_attempts: int
data: dict | None = None
error: str | None = None
# ============ API Endpoints ============
@@ -1011,3 +1055,419 @@ async def test_model(
else None
),
}
# ---------------------------------------------------------------------------
# 带故障转移的模型测试
# ---------------------------------------------------------------------------
def _build_direct_test_candidates(
provider: Provider,
api_format: str | None = None,
) -> list[ProviderCandidate]:
"""
为直接测试模式构建候选列表。
遍历 Provider 的活跃 Endpoint 和 Key不经过 GlobalModel 解析。
"""
from src.services.scheduling.schemas import ProviderCandidate
candidates: list[ProviderCandidate] = []
for endpoint in provider.endpoints or []:
if not getattr(endpoint, "is_active", False):
continue
ep_format = str(getattr(endpoint, "api_format", "") or "")
if not ep_format:
continue
if api_format and ep_format != api_format:
continue
for key in provider.api_keys or []:
if not getattr(key, "is_active", False):
continue
key_formats = getattr(key, "api_formats", None)
if key_formats is not None and ep_format not in key_formats:
continue
candidates.append(
ProviderCandidate(
provider=provider,
endpoint=endpoint,
key=key,
is_skipped=False,
provider_api_format=ep_format,
)
)
return candidates
@router.post("/test-model-failover")
async def test_model_failover(
request: TestModelFailoverRequest,
db: Session = Depends(get_db),
current_user: User = Depends(get_current_user),
) -> Any:
"""
带故障转移的模型测试
支持两种模式:
- global: 模拟外部请求,用全局模型名走候选解析(限定当前 Provider
- direct: 直接测试 provider_model_name在当前 Provider 内多 Key 故障转移
"""
from src.services.candidate.failover import FailoverEngine
from src.services.candidate.policy import RetryMode, RetryPolicy, SkipPolicy
from src.services.task.protocol import AttemptKind, AttemptResult
# 1. 加载 Provider
provider = (
db.query(Provider)
.options(
joinedload(Provider.endpoints),
joinedload(Provider.api_keys),
joinedload(Provider.models),
)
.filter(Provider.id == request.provider_id)
.first()
)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
if request.mode not in ("global", "direct"):
raise HTTPException(status_code=400, detail="mode must be 'global' or 'direct'")
# 2. 构建候选列表
candidates = []
gm_obj = None # GlobalModel 对象global 模式下用于 fallback 映射
if request.mode == "global":
# 模拟外部请求:走 CandidateBuilder 候选解析
from src.services.scheduling.candidate_builder import CandidateBuilder
from src.services.scheduling.candidate_sorter import CandidateSorter
from src.services.scheduling.scheduling_config import SchedulingConfig
sorter = CandidateSorter(SchedulingConfig())
builder = CandidateBuilder(sorter)
# 确定 client_format
client_format = request.api_format
if not client_format:
# 取第一个活跃端点的格式
for ep in provider.endpoints or []:
if getattr(ep, "is_active", False):
client_format = str(getattr(ep, "api_format", "") or "")
if client_format:
break
if not client_format:
raise HTTPException(
status_code=400, detail="No active endpoint found to determine API format"
)
# 从 GlobalModel 提取 model_mappings正则映射规则用于 Key.allowed_models 匹配)
from src.services.cache.model_cache import ModelCacheService
model_mappings: list[str] = []
gm_obj = None
try:
gm_obj = await ModelCacheService.get_global_model_by_name(db, request.model_name)
if gm_obj and isinstance(gm_obj.config, dict):
raw_mappings = gm_obj.config.get("model_mappings", [])
if isinstance(raw_mappings, list):
model_mappings = raw_mappings
except Exception as e:
logger.warning("[test-model-failover] Failed to get GlobalModel mappings: {}", e)
try:
candidates = await builder._build_candidates(
db=db,
providers=[provider],
client_format=client_format,
model_name=request.model_name,
model_mappings=model_mappings if model_mappings else None,
affinity_key=None,
is_stream=False,
)
except Exception as e:
logger.warning("[test-model-failover] CandidateBuilder failed: {}", e)
candidates = []
else:
# 直接测试:简单匹配 Endpoint + Key
candidates = _build_direct_test_candidates(
provider=provider,
api_format=request.api_format,
)
if not candidates:
return TestModelFailoverResponse(
success=False,
model=request.model_name,
provider={"id": str(provider.id), "name": provider.name},
attempts=[],
total_candidates=0,
total_attempts=0,
error="No available candidates found for this model",
).model_dump()
# 3. 定义 attempt_func
attempts: list[TestAttemptDetail] = []
p_type = str(getattr(provider, "provider_type", "") or "").lower()
async def _attempt_func(candidate: Any) -> AttemptResult:
start_time = time.monotonic()
endpoint = candidate.endpoint
key = candidate.key
candidate_idx = getattr(candidate, "_utf_candidate_index", 0)
auth_type = str(getattr(key, "auth_type", "api_key") or "api_key").lower()
extra_headers: dict[str, str] = {}
oauth_meta: dict = {}
effective_model = request.model_name
attempt_recorded = False
try:
# 解析 Key复用统一的认证解析逻辑
effective_proxy = resolve_effective_proxy(
getattr(provider, "proxy", None), getattr(key, "proxy", None)
)
try:
api_key_value, auth_config = await _resolve_key_auth(
key, provider, provider_proxy_config=effective_proxy
)
except _KeyAuthError as e:
raise Exception(e.message) from e
oauth_meta = auth_config or {}
# OAuth 额外头
if auth_type == "oauth":
account_id = oauth_meta.get("account_id")
if account_id:
extra_headers["chatgpt-account-id"] = str(account_id)
ep_extra = get_extra_headers_from_endpoint(endpoint) or {}
extra_headers.update(ep_extra)
# 确定实际模型名
effective_model = request.model_name
if request.mode == "global":
if candidate.mapping_matched_model:
effective_model = candidate.mapping_matched_model
elif gm_obj:
# Fallback: 从 Provider.Model.provider_model_mappings 获取映射
# 与正常请求流程中 _get_mapped_model() 的逻辑一致
gm_id_str = str(gm_obj.id)
for m in provider.models or []:
if not getattr(m, "is_active", False):
continue
if str(getattr(m, "global_model_id", "")) != gm_id_str:
continue
ep_format = str(getattr(endpoint, "api_format", "") or "")
effective_model = m.select_provider_model_name(
affinity_key=None, api_format=ep_format
)
logger.info(
"[test-failover] Fallback mapping: {} -> {} "
"(provider_model_name={}, has_provider_model_mappings={})",
request.model_name,
effective_model,
m.provider_model_name,
bool(m.provider_model_mappings),
)
break
else:
logger.info(
"[test-failover] No matching Model found for gm_id={} in provider={}",
gm_id_str,
provider.name,
)
# 获取 adapter
adapter_class = get_adapter_for_format(endpoint.api_format)
if not adapter_class:
raise Exception(f"Unknown API format: {endpoint.api_format}")
# 构建测试请求
check_request = {
"model": effective_model,
"messages": [{"role": "user", "content": request.message or "Hello"}],
"max_tokens": 30,
"temperature": 0.7,
"stream": True,
}
body_rules = getattr(endpoint, "body_rules", None)
header_rules = getattr(endpoint, "header_rules", None)
# 执行检查
response = await adapter_class.check_endpoint(
None,
endpoint.base_url,
api_key_value,
check_request,
extra_headers if extra_headers else None,
body_rules=body_rules,
header_rules=header_rules,
db=db,
user=current_user,
provider_name=provider.name,
provider_id=str(provider.id),
api_key_id=str(key.id),
model_name=effective_model,
auth_type=auth_type,
provider_type=p_type if p_type else None,
decrypted_auth_config=oauth_meta if oauth_meta else None,
proxy_config=effective_proxy,
)
latency_ms = int((time.monotonic() - start_time) * 1000)
status_code = response.get("status_code", 0)
# 检查响应是否有错误
has_error = bool(response.get("error")) or status_code != 200
if not has_error:
resp_data = response.get("response", {})
resp_body = resp_data.get("response_body", {})
if isinstance(resp_body, str):
try:
parsed = json.loads(resp_body)
except (json.JSONDecodeError, ValueError):
parsed = resp_body
else:
parsed = resp_body
if isinstance(parsed, dict) and "error" in parsed:
has_error = True
if has_error:
error_msg = str(response.get("error", ""))[:300]
if not error_msg and status_code != 200:
error_msg = f"HTTP {status_code}"
if not error_msg and isinstance(parsed, dict) and "error" in parsed:
err_val = parsed["error"]
error_msg = str(
err_val.get("message", err_val)
if isinstance(err_val, dict)
else err_val
)[:300]
attempts.append(
TestAttemptDetail(
candidate_index=candidate_idx,
endpoint_api_format=str(endpoint.api_format),
endpoint_base_url=str(endpoint.base_url)[:80],
key_name=getattr(key, "name", None),
key_id=str(key.id),
auth_type=auth_type,
effective_model=effective_model,
status="failed",
error_message=error_msg,
status_code=status_code,
latency_ms=latency_ms,
)
)
attempt_recorded = True
raise Exception(f"Upstream error: status={status_code}, error={error_msg}")
# 成功
attempts.append(
TestAttemptDetail(
candidate_index=candidate_idx,
endpoint_api_format=str(endpoint.api_format),
endpoint_base_url=str(endpoint.base_url)[:80],
key_name=getattr(key, "name", None),
key_id=str(key.id),
auth_type=auth_type,
effective_model=effective_model,
status="success",
status_code=status_code,
latency_ms=latency_ms,
)
)
return AttemptResult(
kind=AttemptKind.SYNC_RESPONSE,
http_status=status_code,
http_headers={},
response_body=response.get("response", response),
)
except Exception as exc:
latency_ms = int((time.monotonic() - start_time) * 1000)
# has_error 路径已记录带 status_code 的详细 attempt此处仅补录早期异常
if not attempt_recorded:
attempts.append(
TestAttemptDetail(
candidate_index=candidate_idx,
endpoint_api_format=str(endpoint.api_format),
endpoint_base_url=str(endpoint.base_url)[:80],
key_name=getattr(key, "name", None),
key_id=str(key.id),
auth_type=auth_type,
effective_model=effective_model,
status="failed",
error_message=str(exc)[:300],
latency_ms=latency_ms,
)
)
raise
# 4. 预设 candidate indexFailoverEngine 也会 setattr此处兜底防止 setattr 失败)
for i, cand in enumerate(candidates):
cand._utf_candidate_index = i # type: ignore[attr-defined]
# 5. 执行故障转移
try:
engine = FailoverEngine(db)
result = await engine.execute(
candidates=candidates,
attempt_func=_attempt_func,
retry_policy=RetryPolicy(mode=RetryMode.DISABLED),
skip_policy=SkipPolicy(),
request_id=None,
)
# 补充 skipped 候选到 attempts
for i, cand in enumerate(candidates):
if cand.is_skipped and not any(a.candidate_index == i for a in attempts):
attempts.append(
TestAttemptDetail(
candidate_index=i,
endpoint_api_format=str(cand.endpoint.api_format),
endpoint_base_url=str(cand.endpoint.base_url)[:80],
key_name=getattr(cand.key, "name", None),
key_id=str(cand.key.id),
auth_type=str(getattr(cand.key, "auth_type", "") or ""),
status="skipped",
skip_reason=cand.skip_reason,
)
)
attempts.sort(key=lambda a: a.candidate_index)
# 提取成功时的数据
data = None
if result.success and result.attempt_result:
data = {
"stream": True,
"response": result.attempt_result.response_body,
}
return TestModelFailoverResponse(
success=result.success,
model=request.model_name,
provider={"id": str(provider.id), "name": provider.name},
attempts=attempts,
total_candidates=len(candidates),
total_attempts=result.attempt_count,
data=data,
error=result.error_message if not result.success else None,
).model_dump()
except Exception as e:
logger.error("[test-model-failover] Error: {}", e)
return TestModelFailoverResponse(
success=False,
model=request.model_name,
provider={"id": str(provider.id), "name": provider.name},
attempts=attempts,
total_candidates=len(candidates),
total_attempts=0,
error=str(e)[:500],
).model_dump()