mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
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:
150
src/api/handlers/base/chat_error_utils.py
Normal file
150
src/api/handlers/base/chat_error_utils.py
Normal 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)
|
||||
@@ -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}"
|
||||
|
||||
729
src/api/handlers/base/chat_sync_executor.py
Normal file
729
src/api/handlers/base/chat_sync_executor.py
Normal 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}"
|
||||
@@ -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,
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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",
|
||||
|
||||
252
src/api/handlers/base/cli_protocol.py
Normal file
252
src/api/handlers/base/cli_protocol.py
Normal 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: ...
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user