feat(test,export,headers): 模型测试复用统一运行时、实时进度展示、导出增强与请求头大小写保留

- 模型测试 failover 从手动 FailoverEngine 改为 TaskService.execute_sync_candidates 统一运行时
- 前端新增实时 trace 轮询进度展示(候选状态、测试账号、进度条)
- 用户导出/导入支持明文 Key 优先(版本升至 1.2),新增 email_verified 字段
- SENSITIVE_CREDENTIAL_FIELDS 统一到 provider_ops/types.py,补充 refresh_token
- 请求头大小写保留机制(resolve_header_name_case + HeaderBuilder.add 语义修改)
- Codex envelope 移除合成头部,保留客户端原始请求头
- endpoint_checker 支持自定义超时透传
- 新增 x-forwarded-scheme 到上游丢弃头部列表
This commit is contained in:
fawney19
2026-03-06 21:06:45 +08:00
parent 7950ba7dc5
commit 90760da499
28 changed files with 1485 additions and 429 deletions

View File

@@ -39,6 +39,7 @@ class CandidateResponse(BaseModel):
endpoint_name: str | None = None # 端点显示名称api_format
key_id: str | None = None
key_name: str | None = None # 密钥名称
key_account_label: str | None = None # 更适合展示的测试账号标签(优先 OAuth 邮箱)
key_preview: str | None = None # 密钥脱敏预览(如 sk-***abcOAuth 类型不返回
key_auth_type: str | None = None # 密钥认证类型api_key, service_account, oauth
key_oauth_plan_type: str | None = None # OAuth 账号套餐类型free/plus/team/enterprise
@@ -257,6 +258,7 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
key_ids = {c.key_id for c in candidates if c.key_id}
key_map: dict[str, str] = {}
key_preview_map: dict[str, str] = {}
key_account_label_map: dict[str, str | None] = {}
key_capabilities_map: dict[str, dict | None] = {}
key_auth_type_map: dict[str, str] = {}
key_oauth_plan_map: dict[str, str | None] = {}
@@ -268,6 +270,7 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
keys = db.query(ProviderAPIKey).filter(ProviderAPIKey.id.in_(key_ids)).all()
for k in keys:
key_map[k.id] = k.name
key_account_label_map[k.id] = k.name
key_capabilities_map[k.id] = k.capabilities
is_oauth = k.auth_type == "oauth"
@@ -286,6 +289,9 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
try:
decrypted_config = crypto_service.decrypt(k.auth_config)
auth_config = json.loads(decrypted_config)
email = auth_config.get("email")
if isinstance(email, str) and email.strip():
key_account_label_map[k.id] = email.strip()
oauth_plan_type = auth_config.get("plan_type")
if not oauth_plan_type:
ag_tier = auth_config.get("tier")
@@ -348,6 +354,9 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
endpoint_map.get(candidate.endpoint_id) if candidate.endpoint_id else None
)
key_name = key_map.get(candidate.key_id) if candidate.key_id else None
key_account_label = (
key_account_label_map.get(candidate.key_id) if candidate.key_id else None
)
key_preview = key_preview_map.get(candidate.key_id) if candidate.key_id else None
key_auth_type = key_auth_type_map.get(candidate.key_id) if candidate.key_id else None
key_oauth_plan_type = (
@@ -370,6 +379,7 @@ class AdminGetRequestTraceAdapter(AdminApiAdapter):
endpoint_name=endpoint_name,
key_id=candidate.key_id,
key_name=key_name,
key_account_label=key_account_label,
key_preview=key_preview,
key_auth_type=key_auth_type,
key_oauth_plan_type=key_oauth_plan_type,

View File

@@ -7,9 +7,10 @@ from __future__ import annotations
import asyncio
import json
import time
from typing import TYPE_CHECKING, Any
from uuid import uuid4
import httpx
from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel
from sqlalchemy.orm import Session, joinedload
@@ -207,12 +208,14 @@ class TestModelFailoverRequest(BaseModel):
api_format: str | None = None # 指定 API 格式endpoint signature
endpoint_id: str | None = None # 指定仅使用该端点测试
message: str | None = "Hello"
request_id: str | None = None
class TestAttemptDetail(BaseModel):
"""单次测试尝试的详情"""
candidate_index: int
retry_index: int = 0
endpoint_api_format: str
endpoint_base_url: str
key_name: str | None = None
@@ -1180,6 +1183,264 @@ def _filter_test_candidates_by_endpoint(
]
def _parse_jsonish(value: Any) -> Any:
if isinstance(value, str):
try:
return json.loads(value)
except (json.JSONDecodeError, TypeError, ValueError):
return value
return value
def _resolve_test_effective_model(
*,
provider: Provider,
candidate: Any,
request: TestModelFailoverRequest,
gm_obj: Any,
key: Any | None = None,
) -> str:
effective_model = request.model_name
if request.mode != "global":
return effective_model
current_key = key or getattr(candidate, "key", None)
pool_mapping = (
getattr(current_key, "_pool_mapping_matched_model", None) if current_key else None
)
mapping_matched_model = pool_mapping or getattr(candidate, "mapping_matched_model", None)
if mapping_matched_model:
return str(mapping_matched_model)
if not gm_obj:
return effective_model
gm_id_str = str(gm_obj.id)
endpoint = getattr(candidate, "endpoint", None)
ep_format = str(getattr(endpoint, "api_format", "") or "")
for model in provider.models or []:
if not getattr(model, "is_active", False):
continue
if str(getattr(model, "global_model_id", "") or "") != gm_id_str:
continue
selected = model.select_provider_model_name(affinity_key=None, api_format=ep_format)
if selected:
return str(selected)
return effective_model
def _build_test_candidate_meta(
*,
candidates: list[ProviderCandidate],
provider: Provider,
request: TestModelFailoverRequest,
gm_obj: Any,
) -> tuple[dict[tuple[int, str], dict[str, Any]], dict[int, dict[str, Any]]]:
from src.services.scheduling.schemas import PoolCandidate
by_pair: dict[tuple[int, str], dict[str, Any]] = {}
by_candidate: dict[int, dict[str, Any]] = {}
for candidate_index, candidate in enumerate(candidates):
endpoint = candidate.endpoint
base_meta = {
"endpoint_api_format": str(getattr(endpoint, "api_format", "") or ""),
"endpoint_base_url": str(getattr(endpoint, "base_url", "") or "")[:80],
"effective_model": _resolve_test_effective_model(
provider=provider,
candidate=candidate,
request=request,
gm_obj=gm_obj,
),
}
by_candidate[candidate_index] = base_meta
key = getattr(candidate, "key", None)
if key is not None and getattr(key, "id", None):
by_pair[(candidate_index, str(key.id))] = dict(base_meta)
if isinstance(candidate, PoolCandidate):
for pool_key in candidate.pool_keys or []:
if not getattr(pool_key, "id", None):
continue
by_pair[(candidate_index, str(pool_key.id))] = {
"endpoint_api_format": base_meta["endpoint_api_format"],
"endpoint_base_url": base_meta["endpoint_base_url"],
"effective_model": _resolve_test_effective_model(
provider=provider,
candidate=candidate,
request=request,
gm_obj=gm_obj,
key=pool_key,
),
}
return by_pair, by_candidate
def _maybe_mark_test_oauth_key_invalid(
*,
db: Session,
key: Any,
auth_type: str,
error_payload: Any,
) -> None:
if auth_type != "oauth" or not isinstance(error_payload, dict):
return
error_obj = error_payload.get("error")
if not isinstance(error_obj, dict):
return
error_message = str(error_obj.get("message", "") or "")
if error_obj.get("code") != 403:
return
if (
"verify" not in error_message.lower()
and "permission" not in str(error_obj.get("status", "") or "").lower()
):
return
from datetime import datetime, timezone
from src.services.provider.oauth_token import OAUTH_ACCOUNT_BLOCK_PREFIX
key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = f"{OAUTH_ACCOUNT_BLOCK_PREFIX}Google 要求验证账号"
key.is_active = False
db.commit()
def _extract_test_response_or_raise(
*,
response: dict[str, Any],
endpoint: Any,
provider_name: str,
auth_type: str,
api_key: Any,
db: Session,
) -> dict[str, Any]:
status_code = int(response.get("status_code", 0) or 0)
response_payload = response.get("response", {})
parsed_payload = _parse_jsonish(response_payload)
if isinstance(parsed_payload, dict) and "response_body" in parsed_payload:
parsed_payload = _parse_jsonish(parsed_payload.get("response_body"))
if isinstance(parsed_payload, dict) and "error" in parsed_payload:
_maybe_mark_test_oauth_key_invalid(
db=db,
key=api_key,
auth_type=auth_type,
error_payload=parsed_payload,
)
error_obj = parsed_payload["error"]
error_code = error_obj.get("code") if isinstance(error_obj, dict) else status_code or 500
error_message = (
error_obj.get("message") if isinstance(error_obj, dict) else str(error_obj or "")
)
error_status = error_obj.get("status") if isinstance(error_obj, dict) else None
from src.core.exceptions import EmbeddedErrorException
raise EmbeddedErrorException(
provider_name=provider_name,
error_code=int(error_code) if error_code else None,
error_message=str(error_message or ""),
error_status=str(error_status) if error_status else None,
)
if status_code == 200 and not response.get("error"):
return parsed_payload if isinstance(parsed_payload, dict) else response_payload
error_meta = response_payload if isinstance(response_payload, dict) else {}
error_type = str(error_meta.get("error_type", "") or "")
error_message = str(response.get("error", "") or "")
if not error_message and isinstance(parsed_payload, dict):
embedded_error = parsed_payload.get("error")
if isinstance(embedded_error, dict):
error_message = str(embedded_error.get("message", "") or "")
elif embedded_error:
error_message = str(embedded_error)
if not error_message and isinstance(parsed_payload, str):
error_message = parsed_payload
if not error_message and status_code:
error_message = f"HTTP {status_code}"
request_obj = httpx.Request("POST", str(getattr(endpoint, "base_url", "") or ""))
if status_code > 0:
body_text = error_message[:4000] if error_message else ""
synthetic_response = httpx.Response(
status_code=status_code,
request=request_obj,
text=body_text,
headers=response.get("headers", {}),
)
http_error = httpx.HTTPStatusError(
message=body_text or f"HTTP {status_code}",
request=request_obj,
response=synthetic_response,
)
http_error.upstream_response = body_text # type: ignore[attr-defined]
raise http_error
if error_type == "timeout":
raise httpx.TimeoutException(error_message or "Request timeout")
if error_type in {"network_error", "connection_failed"}:
raise httpx.ConnectError(error_message or "Connection failed", request=request_obj)
from src.core.exceptions import ProviderNotAvailableException
raise ProviderNotAvailableException(
error_message or "服务暂时不可用,请稍后重试",
provider_name=provider_name,
upstream_response=error_message or None,
)
def _build_test_attempts_from_candidate_keys(
*,
candidate_keys: list[Any],
candidate_meta_by_pair: dict[tuple[int, str], dict[str, Any]],
candidate_meta_by_index: dict[int, dict[str, Any]],
) -> list[TestAttemptDetail]:
attempts: list[TestAttemptDetail] = []
for candidate_key in candidate_keys:
status = str(getattr(candidate_key, "status", "") or "").strip().lower()
if status in {"", "available", "unused"}:
continue
candidate_index = int(getattr(candidate_key, "candidate_index", 0) or 0)
retry_index = int(getattr(candidate_key, "retry_index", 0) or 0)
key_id = str(getattr(candidate_key, "key_id", "") or "")
meta = candidate_meta_by_pair.get((candidate_index, key_id)) or candidate_meta_by_index.get(
candidate_index, {}
)
attempts.append(
TestAttemptDetail(
candidate_index=candidate_index,
retry_index=retry_index,
endpoint_api_format=str(meta.get("endpoint_api_format", "") or ""),
endpoint_base_url=str(meta.get("endpoint_base_url", "") or ""),
key_name=getattr(candidate_key, "key_name", None),
key_id=key_id,
auth_type=str(getattr(candidate_key, "auth_type", "") or ""),
effective_model=(
str(meta.get("effective_model")) if meta.get("effective_model") else None
),
status=status,
skip_reason=getattr(candidate_key, "skip_reason", None),
error_message=getattr(candidate_key, "error_message", None),
status_code=getattr(candidate_key, "status_code", None),
latency_ms=getattr(candidate_key, "latency_ms", None),
)
)
attempts.sort(key=lambda attempt: (attempt.candidate_index, attempt.retry_index))
return attempts
@router.post("/test-model-failover")
async def test_model_failover(
request: TestModelFailoverRequest,
@@ -1193,11 +1454,13 @@ async def test_model_failover(
- 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
from src.core.exceptions import ProviderNotAvailableException
from src.services.candidate.recorder import CandidateRecorder
from src.services.scheduling.candidate_builder import CandidateBuilder
from src.services.scheduling.candidate_sorter import CandidateSorter
from src.services.scheduling.scheduling_config import SchedulingConfig
from src.services.task import TaskService
# 1. 加载 Provider
provider = (
db.query(Provider)
.options(
@@ -1214,9 +1477,8 @@ async def test_model_failover(
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 映射
candidates: list[ProviderCandidate] = []
gm_obj = None
endpoint_by_id = {
str(getattr(ep, "id", "") or ""): ep
for ep in (provider.endpoints or [])
@@ -1231,21 +1493,14 @@ async def test_model_failover(
if request.api_format and ep_format != request.api_format:
raise HTTPException(status_code=400, detail="endpoint_id does not match api_format")
client_format = request.api_format
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 and requested_endpoint is not None:
client_format = str(getattr(requested_endpoint, "api_format", "") or "")
if not client_format:
# 取第一个活跃端点的格式
for ep in provider.endpoints or []:
if getattr(ep, "is_active", False):
client_format = str(getattr(ep, "api_format", "") or "")
@@ -1256,11 +1511,9 @@ async def test_model_failover(
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):
@@ -1285,7 +1538,10 @@ async def test_model_failover(
candidates = []
candidates = _filter_test_candidates_by_endpoint(candidates, request.endpoint_id)
else:
# 直接测试:简单匹配 Endpoint + Key
if not client_format and requested_endpoint is not None:
client_format = str(getattr(requested_endpoint, "api_format", "") or "")
if not client_format and provider.endpoints:
client_format = str(getattr(provider.endpoints[0], "api_format", "") or "")
candidates = _build_direct_test_candidates(
provider=provider,
api_format=request.api_format,
@@ -1303,266 +1559,165 @@ async def test_model_failover(
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()
request_payload = {
"model": request.model_name,
"messages": [{"role": "user", "content": request.message or "Hello"}],
"max_tokens": 30,
"temperature": 0.7,
"stream": True,
}
request_id = str(request.request_id or f"provider-test-{uuid4().hex[:12]}")
request_timeout = float(getattr(provider, "request_timeout", 0) or TimeoutDefaults.HTTP_REQUEST)
provider_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)
async def _request_func(provider_obj: Any, endpoint: Any, key: Any, candidate: Any) -> Any:
effective_proxy = resolve_effective_proxy(
getattr(provider_obj, "proxy", None), getattr(key, "proxy", None)
)
try:
api_key_value, auth_config = await _resolve_key_auth(
key,
provider_obj,
provider_proxy_config=effective_proxy,
)
except _KeyAuthError as e:
raise RuntimeError(e.message) from e
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
extra_headers = get_extra_headers_from_endpoint(endpoint) or {}
if auth_type == "oauth":
account_id = (auth_config or {}).get("account_id")
if account_id:
extra_headers["chatgpt-account-id"] = str(account_id)
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 {}
effective_model = _resolve_test_effective_model(
provider=provider,
candidate=candidate,
request=request,
gm_obj=gm_obj,
key=key,
)
adapter_class = get_adapter_for_format(endpoint.api_format)
if not adapter_class:
raise ValueError(f"Unknown API format: {endpoint.api_format}")
# 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 = {
response = await adapter_class.check_endpoint(
None,
endpoint.base_url,
api_key_value,
{
**request_payload,
"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,
provider_endpoint=endpoint,
provider_api_key=key,
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,
},
extra_headers if extra_headers else None,
body_rules=getattr(endpoint, "body_rules", None),
header_rules=getattr(endpoint, "header_rules", None),
db=db,
user=current_user,
provider_name=provider_obj.name,
provider_id=str(provider_obj.id),
api_key_id=str(key.id),
model_name=effective_model,
auth_type=auth_type,
provider_type=provider_type if provider_type else None,
decrypted_auth_config=auth_config if auth_config else None,
provider_endpoint=endpoint,
provider_api_key=key,
proxy_config=effective_proxy,
timeout_seconds=request_timeout,
)
return _extract_test_response_or_raise(
response=response,
endpoint=endpoint,
provider_name=str(provider_obj.name),
auth_type=auth_type,
api_key=key,
db=db,
)
# 补充 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,
)
)
candidate_recorder = CandidateRecorder(db)
task_service = TaskService(db)
exec_result = None
run_error: Exception | None = None
attempts.sort(key=lambda a: a.candidate_index)
try:
exec_result = await task_service.execute_sync_candidates(
api_format=client_format or "openai:chat",
model_name=request.model_name,
candidates=candidates,
request_func=_request_func,
request_id=request_id,
current_user=current_user,
user_api_key=None,
is_stream=False,
capability_requirements=None,
request_body_ref={"body": dict(request_payload)},
request_headers=None,
request_body=dict(request_payload),
affinity_key=f"provider-test:{provider.id}",
create_pending_usage=False,
enable_cache_affinity=False,
)
except Exception as exc:
run_error = exc
logger.error("[test-model-failover] Error: {}", exc)
# 提取成功时的数据
data = None
if result.success and result.attempt_result:
data = {
try:
candidate_keys = candidate_recorder.get_candidate_keys(request_id)
except Exception:
candidate_keys = list(exec_result.candidate_keys) if exec_result else []
candidate_meta_by_pair, candidate_meta_by_index = _build_test_candidate_meta(
candidates=candidates,
provider=provider,
request=request,
gm_obj=gm_obj,
)
attempts = _build_test_attempts_from_candidate_keys(
candidate_keys=candidate_keys,
candidate_meta_by_pair=candidate_meta_by_pair,
candidate_meta_by_index=candidate_meta_by_index,
)
total_attempts = sum(1 for attempt in attempts if attempt.status != "skipped")
if exec_result and exec_result.success:
return TestModelFailoverResponse(
success=True,
model=request.model_name,
provider={"id": str(provider.id), "name": provider.name},
attempts=attempts,
total_candidates=len(candidates),
total_attempts=exec_result.attempt_count,
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,
"response": exec_result.response,
},
error=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()
error_message = None
if run_error is not None:
if isinstance(run_error, ProviderNotAvailableException) and getattr(
run_error, "upstream_response", None
):
error_message = str(run_error.upstream_response)[:500]
if not error_message:
error_message = str(run_error)
if not error_message:
failed_attempt = next(
(attempt for attempt in reversed(attempts) if attempt.error_message),
None,
)
error_message = (
failed_attempt.error_message if failed_attempt else "服务暂时不可用,请稍后重试"
)
return TestModelFailoverResponse(
success=False,
model=request.model_name,
provider={"id": str(provider.id), "name": provider.name},
attempts=attempts,
total_candidates=len(candidates),
total_attempts=total_attempts,
error=str(error_message)[:500],
).model_dump()

View File

@@ -22,6 +22,7 @@ from src.database import get_db
from src.models.api import SystemSettingsRequest, SystemSettingsResponse
from src.models.database import ApiKey, Provider, Usage, User
from src.services.email.email_template import EmailTemplate
from src.services.provider_ops.types import SENSITIVE_CREDENTIAL_FIELDS
from src.services.system.config import SystemConfigService
from src.utils.cache_decorator import cache_result
@@ -899,16 +900,7 @@ class AdminExportConfigAdapter(AdminApiAdapter):
"""导出提供商和模型配置"""
# Provider Ops 中需要解密的敏感字段
SENSITIVE_CREDENTIALS = {
"api_key",
"password",
"session_token",
"session_cookie",
"token_cookie",
"auth_cookie",
"cookie_string",
"cookie",
}
SENSITIVE_CREDENTIALS = SENSITIVE_CREDENTIAL_FIELDS
@staticmethod
def _normalize_api_formats(raw_formats: Any) -> list[str]:
@@ -1180,16 +1172,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
"""导入提供商和模型配置"""
# Provider Ops 中需要加密的敏感字段
SENSITIVE_CREDENTIALS = {
"api_key",
"password",
"session_token",
"session_cookie",
"token_cookie",
"auth_cookie",
"cookie_string",
"cookie",
}
SENSITIVE_CREDENTIALS = SENSITIVE_CREDENTIAL_FIELDS
@staticmethod
def _extract_import_key_api_formats(
@@ -1962,8 +1945,45 @@ class AdminImportConfigAdapter(AdminApiAdapter):
class AdminExportUsersAdapter(AdminApiAdapter):
@staticmethod
def _serialize_api_key(key: ApiKey, include_is_standalone: bool = False) -> dict[str, Any]:
"""序列化用户 API Key 为导出格式。"""
from src.core.crypto import crypto_service
data = {
"key_hash": key.key_hash,
"name": key.name,
"balance_used_usd": key.balance_used_usd,
"current_balance_usd": key.current_balance_usd,
"allowed_providers": key.allowed_providers,
"allowed_api_formats": key.allowed_api_formats,
"allowed_models": key.allowed_models,
"rate_limit": key.rate_limit,
"concurrent_limit": key.concurrent_limit,
"force_capabilities": key.force_capabilities,
"is_active": key.is_active,
"expires_at": key.expires_at.isoformat() if key.expires_at else None,
"auto_delete_on_expiry": key.auto_delete_on_expiry,
"total_requests": key.total_requests,
"total_cost_usd": key.total_cost_usd,
}
if key.key_encrypted:
try:
data["key"] = crypto_service.decrypt(key.key_encrypted, silent=True)
except Exception:
logger.warning(
"[USERS_EXPORT] API Key 解密失败,回退为 legacy 密文字段: key_id={}", key.id
)
data["key_encrypted"] = key.key_encrypted
if include_is_standalone:
data["is_standalone"] = key.is_standalone
return data
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""导出用户数据(保留加密数据,排除管理员)"""
"""导出用户数据(优先导出解密后的完整 Key,排除管理员)"""
from datetime import datetime, timezone
from src.core.enums import UserRole
@@ -1971,30 +1991,6 @@ class AdminExportUsersAdapter(AdminApiAdapter):
db = context.db
def _serialize_api_key(key: ApiKey, include_is_standalone: bool = False) -> dict:
"""序列化 API Key 为导出格式"""
data = {
"key_hash": key.key_hash,
"key_encrypted": key.key_encrypted,
"name": key.name,
"balance_used_usd": key.balance_used_usd,
"current_balance_usd": key.current_balance_usd,
"allowed_providers": key.allowed_providers,
"allowed_api_formats": key.allowed_api_formats,
"allowed_models": key.allowed_models,
"rate_limit": key.rate_limit,
"concurrent_limit": key.concurrent_limit,
"force_capabilities": key.force_capabilities,
"is_active": key.is_active,
"expires_at": key.expires_at.isoformat() if key.expires_at else None,
"auto_delete_on_expiry": key.auto_delete_on_expiry,
"total_requests": key.total_requests,
"total_cost_usd": key.total_cost_usd,
}
if include_is_standalone:
data["is_standalone"] = key.is_standalone
return data
# 导出 Users排除管理员
users = db.query(User).filter(User.is_deleted.is_(False), User.role != UserRole.ADMIN).all()
users_data = []
@@ -2006,12 +2002,13 @@ class AdminExportUsersAdapter(AdminApiAdapter):
.all()
)
api_keys_data = [
_serialize_api_key(key, include_is_standalone=True) for key in api_keys
self._serialize_api_key(key, include_is_standalone=True) for key in api_keys
]
users_data.append(
{
"email": user.email,
"email_verified": user.email_verified,
"username": user.username,
"password_hash": user.password_hash,
"role": user.role.value if user.role else "user",
@@ -2029,10 +2026,10 @@ class AdminExportUsersAdapter(AdminApiAdapter):
# 导出独立余额 Keys管理员创建的不属于普通用户
standalone_keys = db.query(ApiKey).filter(ApiKey.is_standalone.is_(True)).all()
standalone_keys_data = [_serialize_api_key(key) for key in standalone_keys]
standalone_keys_data = [self._serialize_api_key(key) for key in standalone_keys]
return {
"version": "1.1",
"version": "1.2",
"exported_at": datetime.now(timezone.utc).isoformat(),
"users": users_data,
"standalone_keys": standalone_keys_data,
@@ -2040,6 +2037,22 @@ class AdminExportUsersAdapter(AdminApiAdapter):
class AdminImportUsersAdapter(AdminApiAdapter):
@staticmethod
def _resolve_api_key_material(key_data: dict[str, Any]) -> tuple[str | None, str | None]:
"""解析用户 API Key 导入材料,优先使用明文 key。"""
from src.core.crypto import crypto_service
from src.models.database import ApiKey
plaintext_key = key_data.get("key")
if isinstance(plaintext_key, str):
normalized = plaintext_key.strip()
if normalized:
return ApiKey.hash_key(normalized), crypto_service.encrypt(normalized)
key_hash = str(key_data.get("key_hash") or "").strip() or None
key_encrypted = key_data.get("key_encrypted")
return key_hash, key_encrypted
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""导入用户数据"""
import uuid
@@ -2079,7 +2092,7 @@ class AdminImportUsersAdapter(AdminApiAdapter):
(None, "skipped"): key 已存在,跳过
(None, "invalid"): 数据无效,跳过
"""
key_hash = key_data.get("key_hash", "").strip()
key_hash, key_encrypted = self._resolve_api_key_material(key_data)
if not key_hash:
return None, "invalid"
@@ -2103,7 +2116,7 @@ class AdminImportUsersAdapter(AdminApiAdapter):
id=str(uuid.uuid4()),
user_id=owner_id,
key_hash=key_hash,
key_encrypted=key_data.get("key_encrypted"),
key_encrypted=key_encrypted,
name=key_data.get("name"),
is_standalone=is_standalone or key_data.get("is_standalone", False),
balance_used_usd=key_data.get("balance_used_usd", 0.0),

View File

@@ -74,6 +74,7 @@ class SyncRequestContext:
mapped_model_result: str | None = None
sync_proxy_info: dict[str, Any] | None = None
provider_response_json: dict[str, Any] | None = None # 格式转换前的提供商原始响应
pool_summary: dict[str, Any] | None = None
class ChatSyncExecutor:

View File

@@ -74,6 +74,7 @@ async def run_endpoint_check(
user: Any | None = None, # User对象
proxy_config: dict[str, Any] | None = None, # 原始代理配置(支持 tunnel 模式)
is_stream: bool | None = None, # 显式流式标记(优先于 body/url 推断)
timeout: float | None = None,
) -> dict[str, Any]:
"""
执行端点检查(重构版本,使用新的架构):
@@ -97,6 +98,7 @@ async def run_endpoint_check(
request_id=str(uuid.uuid4())[:8],
proxy_config=proxy_config,
is_stream=is_stream,
timeout=float(timeout) if timeout is not None else 30.0,
)
# 使用协调器执行检查
@@ -595,6 +597,7 @@ class HttpRequestExecutor:
"""执行HTTP请求支持流式和非流式响应"""
start_time = time.time()
request_id = request.request_id or str(uuid.uuid4())[:8]
effective_timeout = float(request.timeout if request.timeout is not None else self.timeout)
# 检查是否是流式请求(优先显式参数,其次 body最后 URL 推断)
if request.is_stream is not None:
@@ -618,7 +621,7 @@ class HttpRequestExecutor:
# 统一通过 build_proxy_client_kwargs 构建(支持 tunnel 模式 + 普通代理 + 系统默认回退)
client_kwargs = build_proxy_client_kwargs(
proxy_config=request.proxy_config, timeout=self.timeout
proxy_config=request.proxy_config, timeout=effective_timeout
)
async with httpx.AsyncClient(**client_kwargs) as client:

View File

@@ -30,6 +30,7 @@ from src.core.api_format import (
get_adapter_protected_keys_for_endpoint,
get_auth_handler,
get_default_auth_method_for_endpoint,
resolve_header_name_case,
)
from src.core.exceptions import (
ProviderAuthException,
@@ -343,6 +344,7 @@ class HandlerAdapterBase(ApiAdapter):
provider_api_key: Any | None = None,
# 代理配置
proxy_config: dict[str, Any] | None = None,
timeout_seconds: float | None = None,
) -> dict[str, Any]:
"""
测试模型连接性(非流式)
@@ -458,7 +460,8 @@ class HandlerAdapterBase(ApiAdapter):
default_auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
if default_auth_header.lower() != "authorization":
headers.pop(default_auth_header, None)
headers["Authorization"] = f"Bearer {api_key}"
auth_header_name = resolve_header_name_case(extra_headers, "Authorization")
headers[auth_header_name] = f"Bearer {api_key}"
# ---- Body ----
body = cls.build_request_body(request_data, base_url=base_url, provider_type=provider_type)
@@ -527,6 +530,7 @@ class HandlerAdapterBase(ApiAdapter):
api_key_id=api_key_id,
model_name=effective_model_name,
proxy_config=proxy_config,
timeout=timeout_seconds,
)
# =========================================================================

View File

@@ -24,6 +24,7 @@ from src.core.api_format import (
HeaderBuilder,
get_auth_config_for_endpoint,
make_signature_key,
resolve_header_name_case,
)
from src.core.crypto import crypto_service
from src.models.endpoint_models import _CONDITION_OPS, _TYPE_IS_VALUES, parse_re_flags
@@ -1250,7 +1251,7 @@ class PassthroughRequestBuilder(RequestBuilder):
builder.add_many(effective_extra_headers)
# 5. 设置认证头(最高优先级,上游始终使用 header 认证)
builder.add(auth_header, auth_value)
builder.add(resolve_header_name_case(original_headers, auth_header), auth_value)
# 6. 确保有 Content-Type
headers = builder.build()

View File

@@ -14,7 +14,7 @@ from fastapi.responses import JSONResponse
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase, register_adapter
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.core.api_format import ApiFamily, get_auth_handler
from src.core.api_format import ApiFamily, get_auth_handler, resolve_header_name_case
from src.core.api_format.enums import AuthMethod
from src.core.api_format.headers import BROWSER_FINGERPRINT_HEADERS
from src.core.logger import logger
@@ -302,6 +302,7 @@ class GeminiChatAdapter(ChatAdapterBase):
provider_api_key: Any | None = None,
# 代理配置
proxy_config: dict[str, Any] | None = None,
timeout_seconds: float | None = None,
) -> dict[str, Any]:
"""测试 Gemini API 模型连接性(非流式)"""
from src.api.handlers.base.endpoint_checker import run_endpoint_check
@@ -382,7 +383,8 @@ class GeminiChatAdapter(ChatAdapterBase):
default_auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
if default_auth_header.lower() != "authorization":
headers.pop(default_auth_header, None)
headers["Authorization"] = f"Bearer {api_key}"
auth_header_name = resolve_header_name_case(extra_headers, "Authorization")
headers[auth_header_name] = f"Bearer {api_key}"
body = cls.build_request_body(request_data)
@@ -432,6 +434,7 @@ class GeminiChatAdapter(ChatAdapterBase):
api_key_id=api_key_id,
model_name=effective_model_name,
proxy_config=proxy_config,
timeout=timeout_seconds,
)