mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
feat(test): 模型测试支持故障转移,展示每次尝试详情
- 后端新增 test-model-failover 接口,支持 global/direct 两种测试模式 - 利用 FailoverEngine 遍历候选并记录每次尝试的状态、延迟、错误等详情 - 前端 ModelsTab/ModelMappingTab 切换到新接口,移除格式选择下拉菜单 - 新增 TestResultDialog 组件,失败时展示候选尝试详情表格 - 简化组件 props 传递,移除不再需要的 endpoints/mappingPreview 依赖
This commit is contained in:
@@ -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 index(FailoverEngine 也会 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()
|
||||
|
||||
Reference in New Issue
Block a user