refactor: 继续拆分大型模块并增强模块注册健壮性

后端:
- chat_handler_base 错误处理函数提取到 chat_error_utils 子模块
- CLI mixin 引入 CliHandlerProtocol 协议类改善类型标注
- aware_scheduler 拆分为 _candidate_builder 和 _candidate_sorter 子模块
- usage recording 拆分为 _billing_integration 和 _recording_helpers 子模块
- ModuleRegistry 添加循环依赖检测,将写操作从查询方法分离到 reconcile_module_state
- 修正 plugin manager 入度注释

前端:
- 路由守卫逻辑拆分为独立 guards 模块
- ProviderManagement 拆分为 TableHeader/TableRow/BalanceCell/MobileCard 子组件
- SystemSettings 拆分为多个 Section 子组件和 composables
- 提取 useEndpointStatus/useProviderBalance/useProviderFilters composables

测试适配重构后的子模块结构
This commit is contained in:
fawney19
2026-02-14 20:06:38 +08:00
parent 676e918edc
commit 8a670f5524
51 changed files with 7095 additions and 5359 deletions

View File

@@ -0,0 +1,150 @@
"""
Chat Error Utils - Chat Handler 错误处理工具函数
从 chat_handler_base.py 提取的模块级工具函数,用于错误响应的构建和转换。
"""
from __future__ import annotations
import json
from typing import Any
from src.api.handlers.base.utils import get_format_converter_registry
from src.core.exceptions import ThinkingSignatureException, UpstreamClientException
from src.core.logger import logger
from src.models.database import ProviderAPIKey
from src.services.cache.aware_scheduler import ProviderCandidate
from src.services.provider.transport import get_vertex_ai_effective_format
def _get_error_status_code(e: Exception, default: int = 400) -> int:
"""从异常中提取 HTTP 状态码"""
code = getattr(e, "status_code", None)
return code if isinstance(code, int) and code > 0 else default
def _resolve_vertex_ai_format(
key: ProviderAPIKey,
auth_info: Any,
model: str,
provider_api_format: str,
client_api_format: str,
candidate: ProviderCandidate | None,
) -> tuple[str, bool]:
"""
解析 Vertex AI 动态格式并计算 needs_conversion
当 auth_type=vertex_ai 时,同一个 GCP 项目可以访问 Gemini 和 Claude
但它们的请求/响应格式不同,需要根据模型名动态选择。
用户可通过 auth_config.model_format_mapping 配置自定义映射。
Args:
key: Provider API Key
auth_info: 认证信息(包含 decrypted_auth_config
model: 模型名
provider_api_format: 当前 provider API 格式
client_api_format: 客户端 API 格式
candidate: Provider 候选(用于获取原始 needs_conversion
Returns:
(effective_provider_format, needs_conversion) 元组
"""
key_auth_type = getattr(key, "auth_type", "api_key")
if key_auth_type == "vertex_ai":
vertex_auth_config = auth_info.decrypted_auth_config if auth_info else None
effective_format = get_vertex_ai_effective_format(model, vertex_auth_config)
if effective_format.upper() != provider_api_format.upper():
logger.debug(
f"Vertex AI 动态格式切换: {provider_api_format} -> {effective_format} "
f"(model={model})"
)
provider_api_format = effective_format
# Vertex AI 模式下,根据动态格式与客户端格式比较确定是否需要转换
needs_conversion = provider_api_format.upper() != client_api_format.upper()
else:
# 非 Vertex AI使用 candidate 的 needs_conversion
needs_conversion = (
bool(getattr(candidate, "needs_conversion", False)) if candidate else False
)
return provider_api_format, needs_conversion
def _convert_error_response_best_effort(
error_response: dict[str, Any],
source_format: str,
target_format: str,
) -> dict[str, Any]:
"""
将上游错误响应 best-effort 转换为客户端格式。
说明:错误转换走 Canonical registry。转换失败时构造安全的通用错误响应
避免泄露上游原始错误详情。
"""
try:
registry = get_format_converter_registry()
return registry.convert_error_response(error_response, source_format, target_format)
except Exception as e:
logger.debug(f"错误响应转换失败 ({source_format} -> {target_format}): {e}")
# 转换失败时构造安全的通用错误,避免泄露上游详情
return _build_client_error_response_best_effort("upstream error", target_format)
def _build_client_error_response_best_effort(
message: str,
target_format: str,
) -> dict[str, Any]:
"""
当无法解析上游错误 body 时构造一个目标格式的错误响应best-effort
"""
try:
from src.core.api_format.conversion.internal import ErrorType, InternalError
registry = get_format_converter_registry()
normalizer = registry.get_normalizer(target_format)
if normalizer and normalizer.capabilities.supports_error_conversion:
return normalizer.error_from_internal(
InternalError(type=ErrorType.INVALID_REQUEST, message=message, retryable=False)
)
except Exception as e:
logger.debug(f"构建客户端错误响应失败 (target={target_format}): {e}")
return {"error": {"type": "upstream_client_error", "message": message}}
def _build_error_json_payload(
e: ThinkingSignatureException | UpstreamClientException,
client_format: str,
provider_format: str,
needs_conversion: bool = True,
) -> dict[str, Any]:
"""
构建错误 JSON 响应 payload公共逻辑
从异常中提取上游错误信息,尝试转换为客户端格式。
Args:
e: ThinkingSignatureException 或 UpstreamClientException
client_format: 客户端 API 格式
provider_format: Provider API 格式
needs_conversion: 是否需要格式转换
Returns:
格式化的错误响应字典
"""
raw = getattr(e, "upstream_error", None)
message = getattr(e, "message", str(e))
if isinstance(raw, str) and raw:
try:
parsed = json.loads(raw)
except json.JSONDecodeError:
parsed = None
if isinstance(parsed, dict):
if needs_conversion:
return _convert_error_response_best_effort(parsed, provider_format, client_format)
return parsed
return _build_client_error_response_best_effort(message, client_format)

View File

@@ -38,19 +38,19 @@ from src.api.handlers.base.base_handler import (
ClientDisconnectedException,
wait_for_with_disconnect_detection,
)
from src.api.handlers.base.chat_error_utils import (
_build_error_json_payload,
_get_error_status_code,
_resolve_vertex_ai_format,
)
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.request_builder import PassthroughRequestBuilder, get_provider_auth
from src.api.handlers.base.response_parser import ResponseParser
from src.api.handlers.base.stream_context import (
StreamContext,
extract_proxy_timing,
is_format_converted,
)
from src.api.handlers.base.stream_processor import StreamProcessor
from src.api.handlers.base.stream_telemetry import StreamTelemetryRecorder
from src.api.handlers.base.upstream_stream_bridge import (
aggregate_upstream_stream_to_internal_response,
)
from src.api.handlers.base.utils import (
build_sse_headers,
filter_proxy_response_headers,
@@ -60,12 +60,9 @@ from src.config.settings import config
from src.core.api_format.conversion.stream_bridge import (
iter_internal_response_as_stream_events,
)
from src.core.error_utils import extract_client_error_message
from src.core.exceptions import (
EmbeddedErrorException,
ProviderAuthException,
ProviderNotAvailableException,
ProviderRateLimitException,
ProviderTimeoutException,
ThinkingSignatureException,
UpstreamClientException,
@@ -87,145 +84,10 @@ from src.services.provider.stream_policy import (
)
from src.services.provider.transport import (
build_provider_url,
get_vertex_ai_effective_format,
redact_url_for_log,
)
from src.services.system.config import SystemConfigService
def _get_error_status_code(e: Exception, default: int = 400) -> int:
"""从异常中提取 HTTP 状态码"""
code = getattr(e, "status_code", None)
return code if isinstance(code, int) and code > 0 else default
def _resolve_vertex_ai_format(
key: ProviderAPIKey,
auth_info: Any,
model: str,
provider_api_format: str,
client_api_format: str,
candidate: ProviderCandidate | None,
) -> tuple[str, bool]:
"""
解析 Vertex AI 动态格式并计算 needs_conversion
当 auth_type=vertex_ai 时,同一个 GCP 项目可以访问 Gemini 和 Claude
但它们的请求/响应格式不同,需要根据模型名动态选择。
用户可通过 auth_config.model_format_mapping 配置自定义映射。
Args:
key: Provider API Key
auth_info: 认证信息(包含 decrypted_auth_config
model: 模型名
provider_api_format: 当前 provider API 格式
client_api_format: 客户端 API 格式
candidate: Provider 候选(用于获取原始 needs_conversion
Returns:
(effective_provider_format, needs_conversion) 元组
"""
key_auth_type = getattr(key, "auth_type", "api_key")
if key_auth_type == "vertex_ai":
vertex_auth_config = auth_info.decrypted_auth_config if auth_info else None
effective_format = get_vertex_ai_effective_format(model, vertex_auth_config)
if effective_format.upper() != provider_api_format.upper():
logger.debug(
f"Vertex AI 动态格式切换: {provider_api_format} -> {effective_format} "
f"(model={model})"
)
provider_api_format = effective_format
# Vertex AI 模式下,根据动态格式与客户端格式比较确定是否需要转换
needs_conversion = provider_api_format.upper() != client_api_format.upper()
else:
# 非 Vertex AI使用 candidate 的 needs_conversion
needs_conversion = (
bool(getattr(candidate, "needs_conversion", False)) if candidate else False
)
return provider_api_format, needs_conversion
def _convert_error_response_best_effort(
error_response: dict[str, Any],
source_format: str,
target_format: str,
) -> dict[str, Any]:
"""
将上游错误响应 best-effort 转换为客户端格式。
说明:错误转换走 Canonical registry。转换失败时构造安全的通用错误响应
避免泄露上游原始错误详情。
"""
try:
registry = get_format_converter_registry()
return registry.convert_error_response(error_response, source_format, target_format)
except Exception as e:
logger.debug(f"错误响应转换失败 ({source_format} -> {target_format}): {e}")
# 转换失败时构造安全的通用错误,避免泄露上游详情
return _build_client_error_response_best_effort("upstream error", target_format)
def _build_client_error_response_best_effort(
message: str,
target_format: str,
) -> dict[str, Any]:
"""
当无法解析上游错误 body 时构造一个目标格式的错误响应best-effort
"""
try:
from src.core.api_format.conversion.internal import ErrorType, InternalError
registry = get_format_converter_registry()
normalizer = registry.get_normalizer(target_format)
if normalizer and normalizer.capabilities.supports_error_conversion:
return normalizer.error_from_internal(
InternalError(type=ErrorType.INVALID_REQUEST, message=message, retryable=False)
)
except Exception as e:
logger.debug(f"构建客户端错误响应失败 (target={target_format}): {e}")
return {"error": {"type": "upstream_client_error", "message": message}}
def _build_error_json_payload(
e: ThinkingSignatureException | UpstreamClientException,
client_format: str,
provider_format: str,
needs_conversion: bool = True,
) -> dict[str, Any]:
"""
构建错误 JSON 响应 payload公共逻辑
从异常中提取上游错误信息,尝试转换为客户端格式。
Args:
e: ThinkingSignatureException 或 UpstreamClientException
client_format: 客户端 API 格式
provider_format: Provider API 格式
needs_conversion: 是否需要格式转换
Returns:
格式化的错误响应字典
"""
raw = getattr(e, "upstream_error", None)
message = getattr(e, "message", str(e))
if isinstance(raw, str) and raw:
try:
parsed = json.loads(raw)
except json.JSONDecodeError:
parsed = None
if isinstance(parsed, dict):
if needs_conversion:
return _convert_error_response_best_effort(parsed, provider_format, client_format)
return parsed
return _build_client_error_response_best_effort(message, client_format)
@dataclass
class ProviderRequestResult:
"""_prepare_provider_request() 的返回结果,封装请求构建阶段的所有产出。"""
@@ -756,7 +618,11 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
"签名错误" if isinstance(e, ThinkingSignatureException) else "上游客户端错误"
)
self._log_request_error(f"流式请求失败({error_type}", e)
await self._record_stream_failure(ctx, e, original_headers, original_request_body)
from src.api.handlers.base.chat_sync_executor import ChatSyncExecutor
await ChatSyncExecutor(self)._record_stream_failure(
ctx, e, original_headers, original_request_body
)
client_format = (ctx.client_api_format or "").upper()
provider_format = (ctx.provider_api_format or client_format).upper()
payload = _build_error_json_payload(
@@ -769,7 +635,11 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
except Exception as e:
self._log_request_error("流式请求失败", e)
await self._record_stream_failure(ctx, e, original_headers, original_request_body)
from src.api.handlers.base.chat_sync_executor import ChatSyncExecutor
await ChatSyncExecutor(self)._record_stream_failure(
ctx, e, original_headers, original_request_body
)
raise
async def _prepare_provider_request(
@@ -1351,7 +1221,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
response_ctx = None
continue
error_text = await self._extract_error_text(e)
from src.api.handlers.base.chat_sync_executor import ChatSyncExecutor
error_text = await ChatSyncExecutor(self)._extract_error_text(e)
logger.error(
f"Provider 返回错误: {e.response.status_code}\n Response: {error_text}"
)
@@ -1384,57 +1256,6 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
start_time=self.start_time,
)
async def _record_stream_failure(
self,
ctx: StreamContext,
error: Exception,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None:
"""记录流式请求失败"""
response_time_ms = self.elapsed_ms()
status_code = 503
if isinstance(error, ThinkingSignatureException):
status_code = 400
elif isinstance(error, UpstreamClientException):
status_code = _get_error_status_code(error)
elif isinstance(error, ProviderAuthException):
status_code = 503
elif isinstance(error, ProviderRateLimitException):
status_code = 429
elif isinstance(error, ProviderTimeoutException):
status_code = 504
actual_request_body = ctx.provider_request_body or original_request_body
# 失败时返回给客户端的是 JSON 错误响应
client_response_headers = {"content-type": "application/json"}
stream_fail_metadata: dict[str, Any] | None = None
if ctx.proxy_info:
stream_fail_metadata = {"proxy": ctx.proxy_info}
await self.telemetry.record_failure(
provider=ctx.provider_name or "unknown",
model=ctx.model,
response_time_ms=response_time_ms,
status_code=status_code,
error_message=extract_client_error_message(error),
request_headers=original_headers,
request_body=actual_request_body,
is_stream=True,
api_format=ctx.api_format,
provider_request_headers=ctx.provider_request_headers,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
# 格式转换追踪
endpoint_api_format=ctx.provider_api_format or None,
has_format_conversion=ctx.has_format_conversion,
target_model=ctx.mapped_model,
request_metadata=stream_fail_metadata,
)
# ==================== 非流式处理 ====================
async def process_sync(
@@ -1446,561 +1267,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
query_params: dict[str, str] | None = None,
) -> JSONResponse:
"""处理非流式响应"""
logger.debug(f"开始非流式响应处理 ({self.FORMAT_ID})")
from src.api.handlers.base.chat_sync_executor import ChatSyncExecutor
# 转换请求格式
converted_request = await self._convert_request(request)
model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
api_format = self.allowed_api_formats[0]
# 提前创建 pending 记录,让前端可以立即看到"处理中"
self._create_pending_usage(
model=model,
is_stream=False,
request_type="chat",
api_format=self.FORMAT_ID,
request_headers=original_headers,
request_body=original_request_body,
executor = ChatSyncExecutor(self)
return await executor.execute(
request, http_request, original_headers, original_request_body, query_params
)
# 可变请求体容器:允许 TaskService 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: dict[str, Any] = {"body": original_request_body}
# 用于跟踪的变量
provider_name: str | None = None
response_json: dict[str, Any] | None = None
status_code = 200
response_headers: dict[str, str] = {}
provider_request_headers: dict[str, str] = {}
provider_request_body: dict[str, Any] | None = None
provider_api_format_for_error: str | None = None
client_api_format_for_error: str | None = None
needs_conversion_for_error: bool = False # 用于构建错误 payload含 envelope rewrite
provider_id: str | None = None # Provider ID用于失败记录
endpoint_id: str | None = None # Endpoint ID用于失败记录
key_id: str | None = None # Key ID用于失败记录
mapped_model_result: str | None = None # 映射后的目标模型名(用于 Usage 记录)
sync_proxy_info: dict[str, Any] | None = None # 代理信息(用于 Usage 记录)
async def sync_request_func(
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> dict[str, Any]:
nonlocal provider_name, response_json, status_code, response_headers
nonlocal provider_request_headers, provider_request_body, mapped_model_result
nonlocal provider_api_format_for_error, client_api_format_for_error, needs_conversion_for_error
nonlocal sync_proxy_info
provider_name = str(provider.name)
provider_api_format = str(endpoint.api_format or api_format)
client_api_format = (
api_format.value if hasattr(api_format, "value") else str(api_format)
)
# 构建 Provider 请求模型映射、格式转换、envelope 包装)
prep = await self._prepare_provider_request(
model=model,
provider=provider,
endpoint=endpoint,
key=key,
original_request_body=request_body_ref["body"],
client_api_format=client_api_format,
provider_api_format=provider_api_format,
candidate=candidate,
client_is_stream=False,
)
provider_api_format = prep.provider_api_format
needs_conversion = prep.needs_conversion
provider_api_format_for_error = provider_api_format
client_api_format_for_error = client_api_format
needs_conversion_for_error = needs_conversion
mapped_model = prep.mapped_model
if mapped_model:
mapped_model_result = mapped_model
request_body = prep.request_body
url_model = prep.url_model
envelope = prep.envelope
upstream_is_stream = prep.upstream_is_stream
auth_info = prep.auth_info
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
provider_payload, provider_hdrs = self._request_builder.build(
request_body,
original_headers,
endpoint,
key,
is_stream=upstream_is_stream,
extra_headers=prep.extra_headers if prep.extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
)
if upstream_is_stream:
from src.core.api_format.headers import set_accept_if_absent
set_accept_if_absent(provider_hdrs)
provider_request_headers = provider_hdrs
provider_request_body = provider_payload
url = build_provider_url(
endpoint,
query_params=query_params,
path_params={"model": url_model},
is_stream=upstream_is_stream, # sync handler may still force upstream streaming
key=key,
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
)
# 非流式:必须在 build_provider_url 调用后立即缓存(避免 contextvar 被后续调用覆盖)
selected_base_url_cached = envelope.capture_selected_base_url() if envelope else None
# 解析有效代理Key 级别优先于 Provider 级别)
from src.services.proxy_node.resolver import (
get_proxy_label,
resolve_effective_proxy,
resolve_proxy_info,
)
_effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
sync_proxy_info = resolve_proxy_info(_effective_proxy)
_proxy_label = get_proxy_label(sync_proxy_info)
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
logger.info(
f" [{self.request_id}] 发送{'上游流式(聚合)' if upstream_is_stream else '非流式'}请求: "
f"Provider={provider.name}, 模型={model} -> {mapped_model or '无映射'}, "
f"代理={_proxy_label}"
)
logger.debug(f" [{self.request_id}] 请求URL: {redact_url_for_log(url)}")
# 获取复用的 HTTP 客户端支持代理配置Key 级别优先于 Provider 级别)
# 注意:使用 get_proxy_client 复用连接池,不再每次创建新客户端
from src.clients.http_client import HTTPClientPool
from src.services.proxy_node.resolver import (
build_post_kwargs,
build_stream_kwargs,
resolve_delegate_config,
)
# 非流式请求使用 http_request_timeout 作为整体超时
# 优先使用 Provider 配置,否则使用全局配置
request_timeout = provider.request_timeout or config.http_request_timeout
delegate_cfg = resolve_delegate_config(_effective_proxy)
http_client = await HTTPClientPool.get_upstream_client(
delegate_cfg, proxy_config=_effective_proxy
)
# 注意:不使用 async with因为复用的客户端不应该被关闭
# 超时通过 timeout 参数控制
resp: httpx.Response | None = None
if not upstream_is_stream:
try:
_pkw = build_post_kwargs(
delegate_cfg,
url=url,
headers=provider_hdrs,
payload=provider_payload,
timeout=request_timeout,
)
resp = await http_client.post(**_pkw)
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
if envelope:
envelope.on_connection_error(base_url=selected_base_url_cached, exc=e)
if selected_base_url_cached:
logger.warning(
f"[{envelope.name}] Connection error: {selected_base_url_cached} ({e})"
)
raise
else:
# Forced upstream streaming: aggregate SSE to a sync JSON response.
provider_parser = (
get_parser_for_format(provider_api_format) if provider_api_format else None
)
try:
_stream_args = build_stream_kwargs(
delegate_cfg,
url=url,
headers=provider_hdrs,
payload=provider_payload,
timeout=request_timeout,
)
async with http_client.stream(**_stream_args) as stream_resp:
resp = stream_resp
status_code = stream_resp.status_code
response_headers = dict(stream_resp.headers)
extract_proxy_timing(sync_proxy_info, response_headers)
if envelope:
envelope.on_http_status(
base_url=selected_base_url_cached,
status_code=status_code,
)
stream_resp.raise_for_status()
byte_iter = stream_resp.aiter_bytes()
if provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
from src.services.provider.adapters.kiro.eventstream_rewriter import (
apply_kiro_stream_rewrite,
)
byte_iter = apply_kiro_stream_rewrite(byte_iter, model=str(model or ""))
internal_resp = await aggregate_upstream_stream_to_internal_response(
byte_iter,
provider_api_format=provider_api_format,
provider_name=str(provider.name),
model=str(model or ""),
request_id=str(self.request_id or ""),
envelope=envelope,
provider_parser=provider_parser,
)
registry = get_format_converter_registry()
tgt_norm = (
registry.get_normalizer(client_api_format)
if client_api_format
else None
)
if tgt_norm is None:
raise RuntimeError(f"未注册 Normalizer: {client_api_format}")
response_json = tgt_norm.response_from_internal(
internal_resp,
requested_model=model,
)
response_json = response_json if isinstance(response_json, dict) else {}
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
if envelope:
envelope.on_connection_error(base_url=selected_base_url_cached, exc=e)
if selected_base_url_cached:
logger.warning(
f"[{envelope.name}] Connection error: {selected_base_url_cached} ({e})"
)
raise
status_code = resp.status_code
response_headers = dict(resp.headers)
extract_proxy_timing(sync_proxy_info, response_headers)
if envelope:
envelope.on_http_status(base_url=selected_base_url_cached, status_code=status_code)
# Forced upstream streaming already built response_json via aggregator.
if upstream_is_stream:
return response_json if isinstance(response_json, dict) else {}
# 统一使用 HTTPStatusError让 TaskService/error_classifier 负责分类(客户端错误/兼容性错误/限流等)
try:
resp.raise_for_status()
except httpx.HTTPStatusError as e:
error_body = ""
try:
error_body = resp.text[:4000] if resp.text else ""
except Exception:
error_body = ""
# 供 ErrorClassifier 优先读取
e.upstream_response = error_body # type: ignore[attr-defined]
raise
# 安全解析 JSON 响应,处理可能的编码错误
try:
response_json = resp.json()
except (UnicodeDecodeError, json.JSONDecodeError) as e:
# 获取原始响应内容用于调试(存入 upstream_response
raw_content = ""
try:
raw_content = resp.text[:500] if resp.text else "(empty)"
except Exception:
try:
raw_content = repr(resp.content[:500]) if resp.content else "(empty)"
except Exception:
raw_content = "(unable to read)"
logger.error(f"[{self.request_id}] 无法解析响应 JSON: {e}, 原始内容: {raw_content}")
# 判断错误类型,生成友好的客户端错误消息(不暴露提供商信息)
if raw_content == "(empty)" or not raw_content.strip():
client_message = "上游服务返回了空响应"
elif raw_content.strip().startswith(("<", "<!doctype", "<!DOCTYPE")):
client_message = "上游服务返回了非预期的响应格式"
else:
client_message = "上游服务返回了无效的响应"
raise ProviderNotAvailableException(
client_message,
provider_name=str(provider.name),
upstream_status=resp.status_code,
upstream_response=raw_content,
)
if envelope:
response_json = envelope.unwrap_response(response_json)
envelope.postprocess_unwrapped_response(model=model, data=response_json)
# 检查响应体中的嵌套错误HTTP 200 但响应体包含错误)
if isinstance(response_json, dict):
parser = get_parser_for_format(provider_api_format)
if parser.is_error_response(response_json):
parsed = parser.parse_response(response_json, 200)
logger.warning(
f" [{self.request_id}] 非流式检测到嵌套错误: "
f"Provider={provider.name}, "
f"error_type={parsed.error_type}, "
f"embedded_status={parsed.embedded_status_code}, "
f"message={parsed.error_message}"
)
raise EmbeddedErrorException(
provider_name=str(provider.name),
error_code=parsed.embedded_status_code,
error_message=parsed.error_message,
error_status=parsed.error_type,
)
# 跨格式:响应转换回 client_format失败触发 failover
if needs_conversion and isinstance(response_json, dict):
registry = get_format_converter_registry()
response_json = registry.convert_response(
response_json,
provider_api_format,
client_api_format,
requested_model=model, # 使用用户请求的原始模型名
)
return response_json if isinstance(response_json, dict) else {}
try:
# 解析能力需求
capability_requirements = self._resolve_capability_requirements(
model_name=model,
request_headers=original_headers,
request_body=original_request_body,
)
preferred_key_ids = await self._resolve_preferred_key_ids(
model_name=model,
request_body=original_request_body,
)
# 统一入口:总是通过 TaskService
from src.services.task import TaskService
from src.services.task.context import TaskMode
exec_result = await TaskService(self.db, self.redis).execute(
task_type="chat",
task_mode=TaskMode.SYNC,
api_format=api_format,
model_name=model,
user_api_key=self.api_key,
request_func=sync_request_func,
request_id=self.request_id,
is_stream=False,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body_ref=request_body_ref,
)
result = exec_result.response
actual_provider_name = exec_result.provider_name or "unknown"
attempt_id = exec_result.request_candidate_id
provider_id = exec_result.provider_id
endpoint_id = exec_result.endpoint_id
key_id = exec_result.key_id
provider_name = actual_provider_name
response_time_ms = self.elapsed_ms()
# 确保 response_json 不为 None
if response_json is None:
response_json = {}
# 规范化响应
response_json = self._normalize_response(response_json)
# 提取 usage
usage_info = self._extract_usage(response_json)
input_tokens = usage_info.get("input_tokens", 0)
output_tokens = usage_info.get("output_tokens", 0)
cache_creation_tokens = usage_info.get("cache_creation_input_tokens", 0)
cached_tokens = usage_info.get("cache_read_input_tokens", 0)
actual_request_body = provider_request_body or original_request_body
# 非流式成功时,返回给客户端的是提供商响应头(透传)
# JSONResponse 会自动设置 content-type但我们记录实际返回的完整头
client_response_headers = filter_proxy_response_headers(response_headers)
client_response_headers["content-type"] = "application/json"
request_metadata = self._build_request_metadata() or {}
if sync_proxy_info:
request_metadata["proxy"] = sync_proxy_info
total_cost = await self.telemetry.record_success(
provider=provider_name,
model=model,
input_tokens=input_tokens,
output_tokens=output_tokens,
response_time_ms=response_time_ms,
status_code=status_code,
request_headers=original_headers,
request_body=actual_request_body,
response_headers=response_headers,
client_response_headers=client_response_headers,
response_body=response_json,
cache_creation_tokens=cache_creation_tokens,
cache_read_tokens=cached_tokens,
is_stream=False,
provider_request_headers=provider_request_headers,
api_format=api_format,
# 格式转换追踪
endpoint_api_format=provider_api_format_for_error or None,
has_format_conversion=is_format_converted(
provider_api_format_for_error, client_api_format_for_error
),
provider_id=provider_id,
provider_endpoint_id=endpoint_id,
provider_api_key_id=key_id,
# 模型映射信息
target_model=mapped_model_result,
request_metadata=request_metadata or None,
)
logger.debug(f"{self.FORMAT_ID} 非流式响应完成")
# 简洁的请求完成摘要
logger.info(
f"[OK] {self.request_id[:8]} | {model} | {provider_name or 'unknown'} | {response_time_ms}ms | "
f"in:{input_tokens or 0} out:{output_tokens or 0}"
)
# 透传提供商的响应头
return JSONResponse(
status_code=status_code,
content=response_json,
headers=client_response_headers,
)
except ThinkingSignatureException as e:
# Thinking 签名错误TaskService 层已处理整流重试但仍失败
# 记录实际发送给 Provider 的请求体,便于排查问题根因
response_time_ms = self.elapsed_ms()
actual_request_body = provider_request_body or original_request_body
request_metadata = self._build_request_metadata() or {}
if sync_proxy_info:
request_metadata["proxy"] = sync_proxy_info
await self.telemetry.record_failure(
provider=provider_name or "unknown",
model=model,
response_time_ms=response_time_ms,
status_code=e.status_code or 400,
request_headers=original_headers,
request_body=actual_request_body,
error_message=str(e),
is_stream=False,
request_metadata=request_metadata or None,
)
client_format = (client_api_format_for_error or "").upper()
provider_format = (provider_api_format_for_error or client_format).upper()
payload = _build_error_json_payload(
e, client_format, provider_format, needs_conversion=needs_conversion_for_error
)
return JSONResponse(
status_code=_get_error_status_code(e),
content=payload,
)
except UpstreamClientException as e:
response_time_ms = self.elapsed_ms()
actual_request_body = provider_request_body or original_request_body
request_metadata = self._build_request_metadata() or {}
if sync_proxy_info:
request_metadata["proxy"] = sync_proxy_info
await self.telemetry.record_failure(
provider=provider_name or "unknown",
model=model,
response_time_ms=response_time_ms,
status_code=_get_error_status_code(e),
request_headers=original_headers,
request_body=actual_request_body,
error_message=str(e),
is_stream=False,
api_format=api_format,
provider_request_headers=provider_request_headers,
response_headers=response_headers,
client_response_headers={"content-type": "application/json"},
# 格式转换追踪
endpoint_api_format=provider_api_format_for_error or None,
has_format_conversion=is_format_converted(
provider_api_format_for_error, client_api_format_for_error
),
target_model=mapped_model_result,
request_metadata=request_metadata,
)
client_format = (client_api_format_for_error or "").upper()
provider_format = (provider_api_format_for_error or client_format).upper()
payload = _build_error_json_payload(
e, client_format, provider_format, needs_conversion=needs_conversion_for_error
)
return JSONResponse(
status_code=_get_error_status_code(e),
content=payload,
)
except Exception as e:
response_time_ms = self.elapsed_ms()
status_code = 503
if isinstance(e, ProviderAuthException):
status_code = 503
elif isinstance(e, ProviderRateLimitException):
status_code = 429
elif isinstance(e, ProviderTimeoutException):
status_code = 504
actual_request_body = provider_request_body or original_request_body
# 尝试从异常中提取响应头
error_response_headers: dict[str, str] = {}
if isinstance(e, ProviderRateLimitException) and e.response_headers:
error_response_headers = e.response_headers
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
error_response_headers = dict(e.response.headers)
request_metadata = self._build_request_metadata() or {}
if sync_proxy_info:
request_metadata["proxy"] = sync_proxy_info
await self.telemetry.record_failure(
provider=provider_name or "unknown",
model=model,
response_time_ms=response_time_ms,
status_code=status_code,
error_message=extract_client_error_message(e),
request_headers=original_headers,
request_body=actual_request_body,
is_stream=False,
api_format=api_format,
provider_request_headers=provider_request_headers,
response_headers=error_response_headers,
# 非流式失败返回给客户端的是 JSON 错误响应
client_response_headers={"content-type": "application/json"},
# 格式转换追踪
endpoint_api_format=provider_api_format_for_error or None,
has_format_conversion=is_format_converted(
provider_api_format_for_error, client_api_format_for_error
),
# 模型映射信息
target_model=mapped_model_result,
request_metadata=request_metadata,
)
raise
async def _extract_error_text(self, e: httpx.HTTPStatusError) -> str:
"""从 HTTP 错误中提取错误文本"""
try:
if hasattr(e.response, "is_stream_consumed") and not e.response.is_stream_consumed:
error_bytes = await e.response.aread()
return error_bytes.decode("utf-8", errors="replace")
else:
return e.response.text if hasattr(e.response, "_content") else "Unable to read"
except Exception as decode_error:
return f"Unable to read error: {decode_error}"

View File

@@ -0,0 +1,729 @@
"""
ChatSyncExecutor - 非流式请求执行器
从 ChatHandlerBase.process_sync() 提取的独立类,负责:
- 非流式请求的完整执行流程(请求构建、发送、响应解析)
- 通过 SyncRequestContext 管理可变状态(替代原来的 nonlocal 变量)
- 异常处理与 telemetry 记录
- 流式失败记录_record_stream_failure
- HTTP 错误文本提取_extract_error_text
"""
from __future__ import annotations
import json
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
import httpx
from fastapi.responses import JSONResponse
from src.api.handlers.base.chat_error_utils import (
_build_error_json_payload,
_get_error_status_code,
)
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.stream_context import (
StreamContext,
extract_proxy_timing,
is_format_converted,
)
from src.api.handlers.base.utils import (
filter_proxy_response_headers,
get_format_converter_registry,
)
from src.core.error_utils import extract_client_error_message
from src.core.exceptions import (
EmbeddedErrorException,
ProviderAuthException,
ProviderNotAvailableException,
ProviderRateLimitException,
ProviderTimeoutException,
ThinkingSignatureException,
UpstreamClientException,
)
from src.core.logger import logger
if TYPE_CHECKING:
from fastapi import Request
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.cache.aware_scheduler import ProviderCandidate
@dataclass
class SyncRequestContext:
"""同步请求的可变状态容器,替代原来的 nonlocal 变量"""
provider_name: str | None = None
response_json: dict[str, Any] | None = None
status_code: int = 200
response_headers: dict[str, str] = field(default_factory=dict)
provider_request_headers: dict[str, str] = field(default_factory=dict)
provider_request_body: dict[str, Any] | None = None
provider_api_format_for_error: str | None = None
client_api_format_for_error: str | None = None
needs_conversion_for_error: bool = False
provider_id: str | None = None
endpoint_id: str | None = None
key_id: str | None = None
mapped_model_result: str | None = None
sync_proxy_info: dict[str, Any] | None = None
class ChatSyncExecutor:
"""非流式请求执行器,从 ChatHandlerBase 提取"""
def __init__(self, handler: ChatHandlerBase) -> None:
self._handler = handler
self._ctx = SyncRequestContext()
async def execute(
self,
request: Any,
http_request: Request,
original_headers: dict[str, Any],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
) -> JSONResponse:
"""处理非流式响应(原 process_sync 的完整逻辑)"""
handler = self._handler
logger.debug(f"开始非流式响应处理 ({handler.FORMAT_ID})")
# 转换请求格式
converted_request = await handler._convert_request(request)
model = getattr(converted_request, "model", original_request_body.get("model", "unknown"))
api_format = handler.allowed_api_formats[0]
# 提前创建 pending 记录,让前端可以立即看到"处理中"
handler._create_pending_usage(
model=model,
is_stream=False,
request_type="chat",
api_format=handler.FORMAT_ID,
request_headers=original_headers,
request_body=original_request_body,
)
# 可变请求体容器:允许 TaskService 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: dict[str, Any] = {"body": original_request_body}
# 捕获的上下文变量
ctx = self._ctx
async def sync_request_func(
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> dict[str, Any]:
return await self._sync_request_func(
provider,
endpoint,
key,
candidate,
model=model,
api_format=api_format,
original_headers=original_headers,
request_body_ref=request_body_ref,
query_params=query_params,
)
try:
# 解析能力需求
capability_requirements = handler._resolve_capability_requirements(
model_name=model,
request_headers=original_headers,
request_body=original_request_body,
)
preferred_key_ids = await handler._resolve_preferred_key_ids(
model_name=model,
request_body=original_request_body,
)
# 统一入口:总是通过 TaskService
from src.services.task import TaskService
from src.services.task.context import TaskMode
exec_result = await TaskService(handler.db, handler.redis).execute(
task_type="chat",
task_mode=TaskMode.SYNC,
api_format=api_format,
model_name=model,
user_api_key=handler.api_key,
request_func=sync_request_func,
request_id=handler.request_id,
is_stream=False,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body_ref=request_body_ref,
)
actual_provider_name = exec_result.provider_name or "unknown"
ctx.provider_id = exec_result.provider_id
ctx.endpoint_id = exec_result.endpoint_id
ctx.key_id = exec_result.key_id
ctx.provider_name = actual_provider_name
response_time_ms = handler.elapsed_ms()
# 确保 response_json 不为 None
if ctx.response_json is None:
ctx.response_json = {}
# 规范化响应
ctx.response_json = handler._normalize_response(ctx.response_json)
# 提取 usage
usage_info = handler._extract_usage(ctx.response_json)
input_tokens = usage_info.get("input_tokens", 0)
output_tokens = usage_info.get("output_tokens", 0)
cache_creation_tokens = usage_info.get("cache_creation_input_tokens", 0)
cached_tokens = usage_info.get("cache_read_input_tokens", 0)
actual_request_body = ctx.provider_request_body or original_request_body
# 非流式成功时,返回给客户端的是提供商响应头(透传)
# JSONResponse 会自动设置 content-type但我们记录实际返回的完整头
client_response_headers = filter_proxy_response_headers(ctx.response_headers)
client_response_headers["content-type"] = "application/json"
request_metadata = handler._build_request_metadata() or {}
if ctx.sync_proxy_info:
request_metadata["proxy"] = ctx.sync_proxy_info
total_cost = await handler.telemetry.record_success( # noqa: F841
provider=ctx.provider_name,
model=model,
input_tokens=input_tokens,
output_tokens=output_tokens,
response_time_ms=response_time_ms,
status_code=ctx.status_code,
request_headers=original_headers,
request_body=actual_request_body,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
response_body=ctx.response_json,
cache_creation_tokens=cache_creation_tokens,
cache_read_tokens=cached_tokens,
is_stream=False,
provider_request_headers=ctx.provider_request_headers,
api_format=api_format,
# 格式转换追踪
endpoint_api_format=ctx.provider_api_format_for_error or None,
has_format_conversion=is_format_converted(
ctx.provider_api_format_for_error, ctx.client_api_format_for_error
),
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
# 模型映射信息
target_model=ctx.mapped_model_result,
request_metadata=request_metadata or None,
)
logger.debug(f"{handler.FORMAT_ID} 非流式响应完成")
# 简洁的请求完成摘要
logger.info(
f"[OK] {handler.request_id[:8]} | {model} | "
f"{ctx.provider_name or 'unknown'} | {response_time_ms}ms | "
f"in:{input_tokens or 0} out:{output_tokens or 0}"
)
# 透传提供商的响应头
return JSONResponse(
status_code=ctx.status_code,
content=ctx.response_json,
headers=client_response_headers,
)
except ThinkingSignatureException as e:
# Thinking 签名错误TaskService 层已处理整流重试但仍失败
# 记录实际发送给 Provider 的请求体,便于排查问题根因
response_time_ms = handler.elapsed_ms()
actual_request_body = ctx.provider_request_body or original_request_body
request_metadata = handler._build_request_metadata() or {}
if ctx.sync_proxy_info:
request_metadata["proxy"] = ctx.sync_proxy_info
await handler.telemetry.record_failure(
provider=ctx.provider_name or "unknown",
model=model,
response_time_ms=response_time_ms,
status_code=e.status_code or 400,
request_headers=original_headers,
request_body=actual_request_body,
error_message=str(e),
is_stream=False,
request_metadata=request_metadata or None,
)
client_format = (ctx.client_api_format_for_error or "").upper()
provider_format = (ctx.provider_api_format_for_error or client_format).upper()
payload = _build_error_json_payload(
e,
client_format,
provider_format,
needs_conversion=ctx.needs_conversion_for_error,
)
return JSONResponse(
status_code=_get_error_status_code(e),
content=payload,
)
except UpstreamClientException as e:
response_time_ms = handler.elapsed_ms()
actual_request_body = ctx.provider_request_body or original_request_body
request_metadata = handler._build_request_metadata() or {}
if ctx.sync_proxy_info:
request_metadata["proxy"] = ctx.sync_proxy_info
await handler.telemetry.record_failure(
provider=ctx.provider_name or "unknown",
model=model,
response_time_ms=response_time_ms,
status_code=_get_error_status_code(e),
request_headers=original_headers,
request_body=actual_request_body,
error_message=str(e),
is_stream=False,
api_format=api_format,
provider_request_headers=ctx.provider_request_headers,
response_headers=ctx.response_headers,
client_response_headers={"content-type": "application/json"},
# 格式转换追踪
endpoint_api_format=ctx.provider_api_format_for_error or None,
has_format_conversion=is_format_converted(
ctx.provider_api_format_for_error, ctx.client_api_format_for_error
),
target_model=ctx.mapped_model_result,
request_metadata=request_metadata,
)
client_format = (ctx.client_api_format_for_error or "").upper()
provider_format = (ctx.provider_api_format_for_error or client_format).upper()
payload = _build_error_json_payload(
e,
client_format,
provider_format,
needs_conversion=ctx.needs_conversion_for_error,
)
return JSONResponse(
status_code=_get_error_status_code(e),
content=payload,
)
except Exception as e:
response_time_ms = handler.elapsed_ms()
status_code = 503
if isinstance(e, ProviderAuthException):
status_code = 503
elif isinstance(e, ProviderRateLimitException):
status_code = 429
elif isinstance(e, ProviderTimeoutException):
status_code = 504
actual_request_body = ctx.provider_request_body or original_request_body
# 尝试从异常中提取响应头
error_response_headers: dict[str, str] = {}
if isinstance(e, ProviderRateLimitException) and e.response_headers:
error_response_headers = e.response_headers
elif isinstance(e, httpx.HTTPStatusError) and hasattr(e, "response"):
error_response_headers = dict(e.response.headers)
request_metadata = handler._build_request_metadata() or {}
if ctx.sync_proxy_info:
request_metadata["proxy"] = ctx.sync_proxy_info
await handler.telemetry.record_failure(
provider=ctx.provider_name or "unknown",
model=model,
response_time_ms=response_time_ms,
status_code=status_code,
error_message=extract_client_error_message(e),
request_headers=original_headers,
request_body=actual_request_body,
is_stream=False,
api_format=api_format,
provider_request_headers=ctx.provider_request_headers,
response_headers=error_response_headers,
# 非流式失败返回给客户端的是 JSON 错误响应
client_response_headers={"content-type": "application/json"},
# 格式转换追踪
endpoint_api_format=ctx.provider_api_format_for_error or None,
has_format_conversion=is_format_converted(
ctx.provider_api_format_for_error, ctx.client_api_format_for_error
),
# 模型映射信息
target_model=ctx.mapped_model_result,
request_metadata=request_metadata,
)
raise
async def _sync_request_func(
self,
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
*,
model: str,
api_format: Any,
original_headers: dict[str, Any],
request_body_ref: dict[str, Any],
query_params: dict[str, str] | None = None,
) -> dict[str, Any]:
"""单次同步请求(原 sync_request_func 内嵌函数)"""
handler = self._handler
ctx = self._ctx
ctx.provider_name = str(provider.name)
provider_api_format = str(endpoint.api_format or api_format)
client_api_format = api_format.value if hasattr(api_format, "value") else str(api_format)
# 构建 Provider 请求模型映射、格式转换、envelope 包装)
prep = await handler._prepare_provider_request(
model=model,
provider=provider,
endpoint=endpoint,
key=key,
original_request_body=request_body_ref["body"],
client_api_format=client_api_format,
provider_api_format=provider_api_format,
candidate=candidate,
client_is_stream=False,
)
provider_api_format = prep.provider_api_format
needs_conversion = prep.needs_conversion
ctx.provider_api_format_for_error = provider_api_format
ctx.client_api_format_for_error = client_api_format
ctx.needs_conversion_for_error = needs_conversion
mapped_model = prep.mapped_model
if mapped_model:
ctx.mapped_model_result = mapped_model
request_body = prep.request_body
url_model = prep.url_model
envelope = prep.envelope
upstream_is_stream = prep.upstream_is_stream
auth_info = prep.auth_info
# 构建请求(上游始终使用 header 认证,不跟随客户端的 query 方式)
provider_payload, provider_hdrs = handler._request_builder.build(
request_body,
original_headers,
endpoint,
key,
is_stream=upstream_is_stream,
extra_headers=prep.extra_headers if prep.extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
)
if upstream_is_stream:
from src.core.api_format.headers import set_accept_if_absent
set_accept_if_absent(provider_hdrs)
ctx.provider_request_headers = provider_hdrs
ctx.provider_request_body = provider_payload
from src.services.provider.transport import (
build_provider_url,
redact_url_for_log,
)
url = build_provider_url(
endpoint,
query_params=query_params,
path_params={"model": url_model},
is_stream=upstream_is_stream, # sync handler may still force upstream streaming
key=key,
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
)
# 非流式:必须在 build_provider_url 调用后立即缓存(避免 contextvar 被后续调用覆盖)
selected_base_url_cached = envelope.capture_selected_base_url() if envelope else None
# 解析有效代理Key 级别优先于 Provider 级别)
from src.services.proxy_node.resolver import (
get_proxy_label,
resolve_effective_proxy,
resolve_proxy_info,
)
_effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
ctx.sync_proxy_info = resolve_proxy_info(_effective_proxy)
_proxy_label = get_proxy_label(ctx.sync_proxy_info)
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
logger.info(
f" [{handler.request_id}] "
f"发送{'上游流式(聚合)' if upstream_is_stream else '非流式'}请求: "
f"Provider={provider.name}, 模型={model} -> {mapped_model or '无映射'}, "
f"代理={_proxy_label}"
)
logger.debug(f" [{handler.request_id}] 请求URL: {redact_url_for_log(url)}")
# 获取复用的 HTTP 客户端支持代理配置Key 级别优先于 Provider 级别)
# 注意:使用 get_proxy_client 复用连接池,不再每次创建新客户端
from src.clients.http_client import HTTPClientPool
from src.config.settings import config
from src.services.proxy_node.resolver import (
build_post_kwargs,
build_stream_kwargs,
resolve_delegate_config,
)
# 非流式请求使用 http_request_timeout 作为整体超时
# 优先使用 Provider 配置,否则使用全局配置
request_timeout = provider.request_timeout or config.http_request_timeout
delegate_cfg = resolve_delegate_config(_effective_proxy)
http_client = await HTTPClientPool.get_upstream_client(
delegate_cfg, proxy_config=_effective_proxy
)
# 注意:不使用 async with因为复用的客户端不应该被关闭
# 超时通过 timeout 参数控制
resp: httpx.Response | None = None
if not upstream_is_stream:
try:
_pkw = build_post_kwargs(
delegate_cfg,
url=url,
headers=provider_hdrs,
payload=provider_payload,
timeout=request_timeout,
)
resp = await http_client.post(**_pkw)
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
if envelope:
envelope.on_connection_error(base_url=selected_base_url_cached, exc=e)
if selected_base_url_cached:
logger.warning(
f"[{envelope.name}] Connection error: "
f"{selected_base_url_cached} ({e})"
)
raise
else:
# Forced upstream streaming: aggregate SSE to a sync JSON response.
provider_parser = (
get_parser_for_format(provider_api_format) if provider_api_format else None
)
try:
_stream_args = build_stream_kwargs(
delegate_cfg,
url=url,
headers=provider_hdrs,
payload=provider_payload,
timeout=request_timeout,
)
async with http_client.stream(**_stream_args) as stream_resp:
resp = stream_resp
ctx.status_code = stream_resp.status_code
ctx.response_headers = dict(stream_resp.headers)
extract_proxy_timing(ctx.sync_proxy_info, ctx.response_headers)
if envelope:
envelope.on_http_status(
base_url=selected_base_url_cached,
status_code=ctx.status_code,
)
stream_resp.raise_for_status()
byte_iter = stream_resp.aiter_bytes()
if provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
from src.services.provider.adapters.kiro.eventstream_rewriter import (
apply_kiro_stream_rewrite,
)
byte_iter = apply_kiro_stream_rewrite(byte_iter, model=str(model or ""))
from src.api.handlers.base.upstream_stream_bridge import (
aggregate_upstream_stream_to_internal_response,
)
internal_resp = await aggregate_upstream_stream_to_internal_response(
byte_iter,
provider_api_format=provider_api_format,
provider_name=str(provider.name),
model=str(model or ""),
request_id=str(handler.request_id or ""),
envelope=envelope,
provider_parser=provider_parser,
)
registry = get_format_converter_registry()
tgt_norm = (
registry.get_normalizer(client_api_format) if client_api_format else None
)
if tgt_norm is None:
raise RuntimeError(f"未注册 Normalizer: {client_api_format}")
ctx.response_json = tgt_norm.response_from_internal(
internal_resp,
requested_model=model,
)
ctx.response_json = (
ctx.response_json if isinstance(ctx.response_json, dict) else {}
)
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.TimeoutException) as e:
if envelope:
envelope.on_connection_error(base_url=selected_base_url_cached, exc=e)
if selected_base_url_cached:
logger.warning(
f"[{envelope.name}] Connection error: "
f"{selected_base_url_cached} ({e})"
)
raise
ctx.status_code = resp.status_code
ctx.response_headers = dict(resp.headers)
extract_proxy_timing(ctx.sync_proxy_info, ctx.response_headers)
if envelope:
envelope.on_http_status(base_url=selected_base_url_cached, status_code=ctx.status_code)
# Forced upstream streaming already built response_json via aggregator.
if upstream_is_stream:
return ctx.response_json if isinstance(ctx.response_json, dict) else {}
# 统一使用 HTTPStatusError让 TaskService/error_classifier 负责分类
# (客户端错误/兼容性错误/限流等)
try:
resp.raise_for_status()
except httpx.HTTPStatusError as e:
error_body = ""
try:
error_body = resp.text[:4000] if resp.text else ""
except Exception:
error_body = ""
# 供 ErrorClassifier 优先读取
e.upstream_response = error_body # type: ignore[attr-defined]
raise
# 安全解析 JSON 响应,处理可能的编码错误
try:
ctx.response_json = resp.json()
except (UnicodeDecodeError, json.JSONDecodeError) as e:
# 获取原始响应内容用于调试(存入 upstream_response
raw_content = ""
try:
raw_content = resp.text[:500] if resp.text else "(empty)"
except Exception:
try:
raw_content = repr(resp.content[:500]) if resp.content else "(empty)"
except Exception:
raw_content = "(unable to read)"
logger.error(f"[{handler.request_id}] 无法解析响应 JSON: {e}, 原始内容: {raw_content}")
# 判断错误类型,生成友好的客户端错误消息(不暴露提供商信息)
if raw_content == "(empty)" or not raw_content.strip():
client_message = "上游服务返回了空响应"
elif raw_content.strip().startswith(("<", "<!doctype", "<!DOCTYPE")):
client_message = "上游服务返回了非预期的响应格式"
else:
client_message = "上游服务返回了无效的响应"
raise ProviderNotAvailableException(
client_message,
provider_name=str(provider.name),
upstream_status=resp.status_code,
upstream_response=raw_content,
)
if envelope:
ctx.response_json = envelope.unwrap_response(ctx.response_json)
envelope.postprocess_unwrapped_response(model=model, data=ctx.response_json)
# 检查响应体中的嵌套错误HTTP 200 但响应体包含错误)
if isinstance(ctx.response_json, dict):
parser = get_parser_for_format(provider_api_format)
if parser.is_error_response(ctx.response_json):
parsed = parser.parse_response(ctx.response_json, 200)
logger.warning(
f" [{handler.request_id}] 非流式检测到嵌套错误: "
f"Provider={provider.name}, "
f"error_type={parsed.error_type}, "
f"embedded_status={parsed.embedded_status_code}, "
f"message={parsed.error_message}"
)
raise EmbeddedErrorException(
provider_name=str(provider.name),
error_code=parsed.embedded_status_code,
error_message=parsed.error_message,
error_status=parsed.error_type,
)
# 跨格式:响应转换回 client_format失败触发 failover
if needs_conversion and isinstance(ctx.response_json, dict):
registry = get_format_converter_registry()
ctx.response_json = registry.convert_response(
ctx.response_json,
provider_api_format,
client_api_format,
requested_model=model, # 使用用户请求的原始模型名
)
return ctx.response_json if isinstance(ctx.response_json, dict) else {}
async def _record_stream_failure(
self,
ctx: StreamContext,
error: Exception,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None:
"""记录流式请求失败"""
handler = self._handler
response_time_ms = handler.elapsed_ms()
status_code = 503
if isinstance(error, ThinkingSignatureException):
status_code = 400
elif isinstance(error, UpstreamClientException):
status_code = _get_error_status_code(error)
elif isinstance(error, ProviderAuthException):
status_code = 503
elif isinstance(error, ProviderRateLimitException):
status_code = 429
elif isinstance(error, ProviderTimeoutException):
status_code = 504
actual_request_body = ctx.provider_request_body or original_request_body
# 失败时返回给客户端的是 JSON 错误响应
client_response_headers = {"content-type": "application/json"}
stream_fail_metadata: dict[str, Any] | None = None
if ctx.proxy_info:
stream_fail_metadata = {"proxy": ctx.proxy_info}
await handler.telemetry.record_failure(
provider=ctx.provider_name or "unknown",
model=ctx.model,
response_time_ms=response_time_ms,
status_code=status_code,
error_message=extract_client_error_message(error),
request_headers=original_headers,
request_body=actual_request_body,
is_stream=True,
api_format=ctx.api_format,
provider_request_headers=ctx.provider_request_headers,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
# 格式转换追踪
endpoint_api_format=ctx.provider_api_format or None,
has_format_conversion=ctx.has_format_conversion,
target_model=ctx.mapped_model,
request_metadata=stream_fail_metadata,
)
async def _extract_error_text(self, e: httpx.HTTPStatusError) -> str:
"""从 HTTP 错误中提取错误文本"""
try:
if hasattr(e.response, "is_stream_consumed") and not e.response.is_stream_consumed:
error_bytes = await e.response.aread()
return error_bytes.decode("utf-8", errors="replace")
else:
return e.response.text if hasattr(e.response, "_content") else "Unable to read"
except Exception as decode_error:
return f"Unable to read error: {decode_error}"

View File

@@ -3,7 +3,7 @@
from __future__ import annotations
import json
from typing import Any
from typing import TYPE_CHECKING, Any
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.stream_context import StreamContext
@@ -18,12 +18,15 @@ from .cli_sse_helpers import (
_parse_sse_event_data_line,
)
if TYPE_CHECKING:
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
class CliEventMixin:
"""SSE 事件处理和格式转换相关方法的 Mixin"""
def _handle_sse_event(
self,
self: CliHandlerProtocol,
ctx: StreamContext,
event_name: str | None,
data_str: str,

View File

@@ -5,7 +5,7 @@ from __future__ import annotations
import asyncio
import time
from collections.abc import AsyncGenerator
from typing import Any
from typing import TYPE_CHECKING, Any
import httpx
from fastapi import Request
@@ -25,6 +25,9 @@ from src.database import get_db
from src.models.database import User
from src.services.provider.behavior import get_provider_behavior
if TYPE_CHECKING:
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
class CliMonitorMixin:
"""监控和统计相关方法的 Mixin"""
@@ -157,7 +160,7 @@ class CliMonitorMixin:
raise
async def _record_stream_stats(
self,
self: CliHandlerProtocol,
ctx: StreamContext,
original_headers: dict[str, str],
original_request_body: dict[str, Any],

View File

@@ -29,6 +29,7 @@ from src.utils.sse_parser import SSEEventParser
from src.utils.timeout import read_first_chunk_with_ttfb_timeout
if TYPE_CHECKING:
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
from src.models.database import Provider, ProviderEndpoint
@@ -136,7 +137,7 @@ class CliPrefetchMixin:
)
async def _prefetch_and_check_embedded_error(
self,
self: CliHandlerProtocol,
byte_iterator: Any,
provider: "Provider",
endpoint: "ProviderEndpoint",

View File

@@ -0,0 +1,252 @@
"""
CLI Handler Mixin Protocol -- Mixin 隐式依赖的编译时契约
各 Mixin (CliStreamMixin, CliSyncMixin, CliRequestMixin, CliMonitorMixin,
CliPrefetchMixin, CliEventMixin) 通过 duck typing 访问宿主类的属性和方法。
本模块将这些隐式依赖显式声明为 Protocol使 mypy/pyright 能在编辑期捕获
缺失属性或类型不匹配的错误。
渐进式采用:仅在各 Mixin 的公开方法签名中标注 `self: CliHandlerProtocol`
不修改方法体或私有 helper。
"""
from __future__ import annotations
from typing import (
TYPE_CHECKING,
Any,
Protocol,
runtime_checkable,
)
if TYPE_CHECKING:
from redis import Redis
from sqlalchemy.orm import Session
from src.api.handlers.base.base_handler import MessageTelemetry
from src.api.handlers.base.request_builder import RequestBuilder
from src.api.handlers.base.response_parser import ResponseParser
from src.api.handlers.base.stream_context import StreamContext
from src.models.database import ApiKey, User
@runtime_checkable
class CliHandlerProtocol(Protocol):
"""CLI Handler Mixin 宿主需要满足的属性/方法契约。
声明范围仅覆盖 Mixin 实际引用的 self.xxx不要求宿主实现全部
BaseMessageHandler 接口。
"""
# ------------------------------------------------------------------
# 实例属性 -- 来自 BaseMessageHandler.__init__
# ------------------------------------------------------------------
db: Session
user: User
api_key: ApiKey
request_id: str
client_ip: str
user_agent: str
start_time: float
allowed_api_formats: list[str]
redis: Redis # type: ignore[type-arg]
telemetry: MessageTelemetry
perf_metrics: dict[str, Any] | None
# ------------------------------------------------------------------
# 类属性 -- 来自 CliMessageHandlerBase
# ------------------------------------------------------------------
FORMAT_ID: str
DATA_TIMEOUT: int
EMPTY_CHUNK_THRESHOLD: int
# ------------------------------------------------------------------
# 属性/方法 -- 来自 CliMessageHandlerBase / BaseMessageHandler
# ------------------------------------------------------------------
@property
def parser(self) -> ResponseParser: ...
_request_builder: RequestBuilder
# ------------------------------------------------------------------
# 方法 -- 来自 BaseMessageHandler (被多个 Mixin 引用)
# ------------------------------------------------------------------
def _create_pending_usage(
self,
model: str,
is_stream: bool,
request_type: str = ...,
api_format: str | None = ...,
request_headers: dict[str, Any] | None = ...,
request_body: dict[str, Any] | None = ...,
) -> None: ...
def _build_request_metadata(
self,
http_request: Any | None = ...,
) -> dict[str, Any] | None: ...
def _resolve_capability_requirements(
self,
model_name: str,
request_headers: dict[str, str] | None = ...,
request_body: dict[str, Any] | None = ...,
) -> dict[str, bool]: ...
async def _resolve_preferred_key_ids(
self,
model_name: str,
request_body: dict[str, Any] | None = ...,
) -> list[str] | None: ...
def _update_usage_to_streaming(
self,
request_id: str | None = ...,
) -> None: ...
def _update_usage_to_streaming_with_ctx(
self,
ctx: StreamContext,
) -> None: ...
def _log_request_error(
self,
message: str,
error: Exception,
) -> None: ...
# ------------------------------------------------------------------
# 方法 -- 来自 CliRequestMixin (被 CliStreamMixin / CliSyncMixin 引用)
# ------------------------------------------------------------------
def extract_model_from_request(
self,
request_body: dict[str, Any],
path_params: dict[str, Any] | None = ...,
) -> str: ...
async def _get_mapped_model(
self,
source_model: str,
provider_id: str,
) -> str | None: ...
def apply_mapped_model(
self,
request_body: dict[str, Any],
mapped_model: str,
) -> dict[str, Any]: ...
def prepare_provider_request_body(
self,
request_body: dict[str, Any],
) -> dict[str, Any]: ...
def finalize_provider_request(
self,
request_body: dict[str, Any],
*,
mapped_model: str | None,
provider_api_format: str | None,
) -> dict[str, Any]: ...
def get_model_for_url(
self,
request_body: dict[str, Any],
mapped_model: str | None,
) -> str | None: ...
def _convert_request_for_cross_format(
self,
request_body: dict[str, Any],
client_api_format: str,
provider_api_format: str,
mapped_model: str | None,
fallback_model: str,
is_stream: bool,
*,
target_variant: str | None = ...,
) -> tuple[dict[str, Any], str]: ...
def _extract_response_metadata(
self,
response: dict[str, Any],
) -> dict[str, Any]: ...
# ------------------------------------------------------------------
# 方法 -- 来自 CliEventMixin (被 CliStreamMixin / CliPrefetchMixin 引用)
# ------------------------------------------------------------------
def _handle_sse_event(
self,
ctx: StreamContext,
event_name: str | None,
data_str: str,
record_chunk: bool = ...,
) -> None: ...
def _mark_first_output(
self,
ctx: StreamContext,
state: dict[str, bool],
) -> None: ...
def _convert_sse_line(
self,
ctx: StreamContext,
line: str,
events: list[Any],
) -> tuple[list[str], list[dict[str, Any]]]: ...
def _record_converted_chunks(
self,
ctx: StreamContext,
converted_events: list[dict[str, Any]],
) -> None: ...
def _finalize_stream_metadata(
self,
ctx: StreamContext,
) -> None: ...
# ------------------------------------------------------------------
# 方法 -- 来自 CliPrefetchMixin (被 CliStreamMixin 引用)
# ------------------------------------------------------------------
def _flush_remaining_sse_data(
self,
ctx: StreamContext,
buffer: bytes,
decoder: Any,
sse_parser: Any,
*,
record_chunk: bool = ...,
) -> None: ...
def _estimate_tokens_for_incomplete_stream(
self,
ctx: StreamContext,
request_body: dict[str, Any],
) -> None: ...
# ------------------------------------------------------------------
# 方法 -- 来自 CliMonitorMixin (被 CliStreamMixin 引用)
# ------------------------------------------------------------------
async def _create_monitored_stream(
self,
ctx: StreamContext,
stream_generator: Any,
http_request: Any | None = ...,
) -> Any: ...
async def _record_stream_stats(
self,
ctx: StreamContext,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None: ...
async def _record_stream_failure(
self,
ctx: StreamContext,
error: Exception,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None: ...

View File

@@ -11,6 +11,7 @@ from src.api.handlers.base.utils import get_format_converter_registry
from src.core.logger import logger
if TYPE_CHECKING:
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
from src.core.api_format import EndpointDefinition
@@ -63,7 +64,7 @@ class CliRequestMixin:
return None
def extract_model_from_request(
self,
self: CliHandlerProtocol,
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002 - 子类使用
) -> str:

View File

@@ -53,6 +53,7 @@ from src.utils.timeout import read_first_chunk_with_ttfb_timeout
from .cli_sse_helpers import _format_converted_events_to_sse
if TYPE_CHECKING:
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
@@ -60,7 +61,7 @@ class CliStreamMixin:
"""流式处理核心方法的 Mixin"""
async def process_stream(
self,
self: CliHandlerProtocol,
original_request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,

View File

@@ -39,6 +39,7 @@ from src.services.provider.stream_policy import (
from src.services.provider.transport import build_provider_url
if TYPE_CHECKING:
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
@@ -46,7 +47,7 @@ class CliSyncMixin:
"""同步处理相关方法的 Mixin"""
async def process_sync(
self,
self: CliHandlerProtocol,
original_request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,