mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
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:
@@ -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-***abc),OAuth 类型不返回
|
||||
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,
|
||||
|
||||
@@ -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 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,
|
||||
},
|
||||
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()
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
# =========================================================================
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user