refactor: 移除 Python 后端源码,全面迁移至 Rust gateway 架构

- 删除全部 Python 源码 (src/) 及 Alembic 迁移脚本,归档至 _deprecated_py_src/
- 重构 Rust gateway ai_pipeline: 拆分 planner/finalize 模块,新增 contracts/adaptation 层
- 重组 handlers 模块为 admin/public/proxy/internal/shared 子模块结构
- 新增 executor 模块,引入 Rust 原生数据库迁移 (aether-data/migrations)
- 简化 CI/Docker 构建流程,移除 base image 二级构建,统一为单一 app image
- 移除 Python 相关基础设施文件 (entrypoint.sh, gunicorn_conf.py, Dockerfile.base)
This commit is contained in:
fawney19
2026-04-03 16:26:16 +08:00
parent 8f26e1a31f
commit 1d9c77522a
868 changed files with 1735 additions and 2433 deletions

View File

@@ -0,0 +1,68 @@
"""
Handler 基类模块
提供 Adapter、Handler 的抽象基类,以及请求构建器和响应解析器。
注意Handler 基类ChatHandlerBase, CliMessageHandlerBase 等)不在这里导出,
因为它们依赖 services.usage.stream而后者又需要导入 response_parser
会形成循环导入。请直接从具体模块导入 Handler 基类。
"""
# Chat Adapter 基类(不会引起循环导入)
from src.api.handlers.base.chat_adapter_base import (
ChatAdapterBase,
get_adapter_class,
get_adapter_instance,
list_registered_formats,
register_adapter,
)
# CLI Adapter 基类
from src.api.handlers.base.cli_adapter_base import (
CliAdapterBase,
get_cli_adapter_class,
get_cli_adapter_instance,
list_registered_cli_formats,
register_cli_adapter,
)
# 请求构建器
from src.api.handlers.base.request_builder import (
SENSITIVE_HEADERS,
PassthroughRequestBuilder,
RequestBuilder,
build_passthrough_request,
)
# 响应解析器
from src.api.handlers.base.response_parser import (
ParsedChunk,
ParsedResponse,
ResponseParser,
StreamStats,
)
__all__ = [
# Chat Adapter
"ChatAdapterBase",
"register_adapter",
"get_adapter_class",
"get_adapter_instance",
"list_registered_formats",
# CLI Adapter
"CliAdapterBase",
"register_cli_adapter",
"get_cli_adapter_class",
"get_cli_adapter_instance",
"list_registered_cli_formats",
# 请求构建器
"RequestBuilder",
"PassthroughRequestBuilder",
"build_passthrough_request",
"SENSITIVE_HEADERS",
# 响应解析器
"ResponseParser",
"ParsedChunk",
"ParsedResponse",
"StreamStats",
]

View File

@@ -0,0 +1,740 @@
"""
基础消息处理器,封装通用的编排、转换、遥测逻辑。
接口约定:
- process_stream: 处理流式请求,返回 StreamingResponse
- process_sync: 处理非流式请求,返回 JSONResponse
签名规范(推荐):
async def process_stream(
self,
request: Any, # 解析后的请求模型
http_request: Request, # FastAPI Request 对象
original_headers: dict[str, str], # 原始请求头
original_request_body: dict[str, Any], # 原始请求体
query_params: dict[str, str] | None = None, # 查询参数
) -> StreamingResponse: ...
async def process_sync(
self,
request: Any,
http_request: Request,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
) -> JSONResponse: ...
"""
from __future__ import annotations
import asyncio
import time
from collections.abc import Awaitable, Callable, Coroutine
from typing import (
TYPE_CHECKING,
Any,
Protocol,
TypeVar,
runtime_checkable,
)
from fastapi import Request
from fastapi.responses import JSONResponse, StreamingResponse
from sqlalchemy.orm import Session
from src.clients.redis_client import get_redis_client_sync
from src.core.logger import logger
from src.services.provider.format import normalize_endpoint_signature
from src.services.usage.service import UsageService
from src.services.usage.telemetry import MessageTelemetry # re-export
if TYPE_CHECKING:
from src.api.handlers.base.stream_context import StreamContext
# Adapter 检测器类型:接受 headers 和可选的 request_body返回能力需求字典
type AdapterDetectorType = Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
# MessageTelemetry -- re-export from src.services.usage.telemetry (see import above)
__all__ = ["MessageTelemetry", "MessageHandlerProtocol", "AdapterDetectorType"]
@runtime_checkable
class MessageHandlerProtocol(Protocol):
"""
消息处理器协议 - 定义标准接口
ChatHandlerBase 和 CliMessageHandlerBase 均支持 http_request 参数用于客户端断连检测。
"""
async def process_stream(
self,
request: Any,
http_request: Request,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
) -> StreamingResponse:
"""处理流式请求"""
...
async def process_sync(
self,
request: Any,
http_request: Request,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
) -> JSONResponse:
"""处理非流式请求"""
...
class BaseMessageHandler:
"""
消息处理器基类,所有具体格式的 handler 可以继承它。
子类需要实现:
- process_stream: 处理流式请求
- process_sync: 处理非流式请求
推荐使用 MessageHandlerProtocol 中定义的签名。
"""
def __init__(
self,
*,
db: Session,
user: Any,
api_key: Any,
request_id: str,
client_ip: str,
user_agent: str,
start_time: float,
allowed_api_formats: list[str] | None = None,
adapter_detector: AdapterDetectorType | None = None,
perf_metrics: dict[str, Any] | None = None,
api_family: str | None = None,
endpoint_kind: str | None = None,
) -> None:
self.db = db
self.user = user
self.api_key = api_key
self.request_id = request_id
self.client_ip = client_ip
self.user_agent = user_agent
self.start_time = start_time
# 新模式endpoint signature keyfamily:kind如 "claude:chat"
self.allowed_api_formats = allowed_api_formats or ["claude:chat"]
self.primary_api_format = normalize_endpoint_signature(self.allowed_api_formats[0])
self.adapter_detector = adapter_detector
self.perf_metrics = perf_metrics
# 结构化格式维度(从 Adapter 层透传)
self.api_family = api_family
self.endpoint_kind = endpoint_kind
redis_client = get_redis_client_sync()
self.redis = redis_client
self.telemetry = MessageTelemetry(db, user, api_key, request_id, client_ip)
def elapsed_ms(self) -> int:
return int((time.time() - self.start_time) * 1000)
def _build_request_metadata(self, http_request: Request | None = None) -> dict[str, Any] | None:
if not isinstance(self.perf_metrics, dict) or not self.perf_metrics:
return None
return {"perf": self.perf_metrics}
@staticmethod
def _normalize_candidate_status(candidate: dict[str, Any]) -> str:
status = candidate.get("status")
if isinstance(status, str) and status.strip():
return status.strip().lower()
attempt_status = candidate.get("attempt_status")
if isinstance(attempt_status, str) and attempt_status.strip():
return attempt_status.strip().lower()
if candidate.get("skipped"):
return "skipped"
return ""
@staticmethod
def _to_int(value: Any, default: int = 0) -> int:
try:
return int(value)
except Exception:
return default
def _load_request_candidate_keys(self) -> list[Any]:
if not self.request_id:
return []
try:
from src.services.candidate.recorder import CandidateRecorder
return CandidateRecorder(self.db).get_candidate_keys(self.request_id)
except Exception:
return []
def _compact_candidate_key_snapshot(self, item: Any) -> dict[str, Any] | None:
raw: dict[str, Any] | None = None
if isinstance(item, dict):
raw = dict(item)
elif hasattr(item, "to_dict"):
try:
converted = item.to_dict()
if isinstance(converted, dict):
raw = dict(converted)
except Exception:
raw = None
if raw is None:
return None
status = self._normalize_candidate_status(raw)
candidate_index = raw.get("candidate_index", raw.get("index", 0))
retry_index = raw.get("retry_index", 0)
snapshot: dict[str, Any] = {
"candidate_index": self._to_int(candidate_index, 0),
"retry_index": self._to_int(retry_index, 0),
}
passthrough_fields = (
"provider_id",
"provider_name",
"endpoint_id",
"key_id",
"key_name",
"auth_type",
"priority",
"is_cached",
"skip_reason",
"error_type",
"status_code",
"latency_ms",
)
for field in passthrough_fields:
value = raw.get(field)
if value is not None and value != "":
snapshot[field] = value
if status:
snapshot["status"] = status
if raw.get("skipped"):
snapshot["skipped"] = True
snapshot.setdefault("status", "skipped")
if "selected" in raw:
snapshot["selected"] = bool(raw.get("selected"))
error_message = raw.get("error_message")
if isinstance(error_message, str) and error_message:
snapshot["error_message"] = error_message[:240]
return snapshot
def _collect_candidate_snapshots(
self,
*,
candidate_keys: list[Any] | None = None,
fallback_from_request: bool = False,
) -> list[dict[str, Any]]:
source = candidate_keys
if (not source) and fallback_from_request:
source = self._load_request_candidate_keys()
snapshots: list[dict[str, Any]] = []
for item in source or []:
snapshot = self._compact_candidate_key_snapshot(item)
if snapshot:
snapshots.append(snapshot)
snapshots.sort(
key=lambda it: (
self._to_int(it.get("candidate_index"), 0),
self._to_int(it.get("retry_index"), 0),
)
)
return snapshots[:64]
def _build_scheduling_audit(
self,
snapshots: list[dict[str, Any]],
*,
selected_key_id: str | None = None,
) -> dict[str, Any] | None:
if not snapshots:
return None
# "unused" means the candidate was pre-created for audit but never actually attempted.
executed_status_exclude = {"", "available", "pending", "skipped", "unused"}
executed_count = 0
attempts: list[dict[str, Any]] = []
account_map: dict[str, dict[str, Any]] = {}
candidate_indices: set[int] = set()
key_ids: set[str] = set()
for snapshot in snapshots:
status = str(snapshot.get("status", "") or "").lower()
if status in executed_status_exclude:
continue
executed_count += 1
candidate_index = self._to_int(snapshot.get("candidate_index"), 0)
retry_index = self._to_int(snapshot.get("retry_index"), 0)
key_id = snapshot.get("key_id")
key_name = snapshot.get("key_name")
provider_id = snapshot.get("provider_id")
provider_name = snapshot.get("provider_name")
candidate_indices.add(candidate_index)
if isinstance(key_id, str) and key_id:
key_ids.add(key_id)
if len(attempts) < 24:
attempts.append(
{
"candidate_index": candidate_index,
"retry_index": retry_index,
"provider_id": provider_id,
"provider_name": provider_name,
"key_id": key_id,
"key_name": key_name,
"status": status,
"status_code": snapshot.get("status_code"),
"error_type": snapshot.get("error_type"),
}
)
if not isinstance(key_id, str) or not key_id:
continue
account = account_map.get(key_id)
if account is None:
account = {
"key_id": key_id,
"key_name": key_name,
"provider_id": provider_id,
"provider_name": provider_name,
"attempts": 0,
"successes": 0,
"last_status": status,
}
account_map[key_id] = account
account["attempts"] = self._to_int(account.get("attempts"), 0) + 1
if status in {"success", "streaming"}:
account["successes"] = self._to_int(account.get("successes"), 0) + 1
account["last_status"] = status
if executed_count == 0:
return {
"mode": "internal",
"attempted_count": 0,
"account_count": 0,
"retry_occurred": False,
"failover_occurred": False,
"accounts": [],
"attempts": [],
}
selected_key_id_norm = str(selected_key_id) if selected_key_id else None
accounts = list(account_map.values())[:12]
selected_account: dict[str, Any] | None = None
if selected_key_id_norm and selected_key_id_norm in account_map:
selected_account = dict(account_map[selected_key_id_norm])
else:
for account in account_map.values():
if self._to_int(account.get("successes"), 0) > 0:
selected_account = dict(account)
selected_key_id_norm = str(account.get("key_id", ""))
break
if selected_account is not None:
for account in accounts:
if account.get("key_id") == selected_account.get("key_id"):
account["selected"] = True
failover_occurred = executed_count > 1 and (len(candidate_indices) > 1 or len(key_ids) > 1)
return {
"mode": "internal",
"attempted_count": executed_count,
"account_count": len(account_map),
"retry_occurred": executed_count > 1,
"failover_occurred": bool(failover_occurred),
"selected_key_id": selected_key_id_norm,
"selected_account": selected_account,
"accounts": accounts,
"attempts": attempts,
}
def _build_scheduling_metadata(
self,
*,
candidate_keys: list[Any] | None = None,
selected_key_id: str | None = None,
pool_summary: dict[str, Any] | None = None,
fallback_from_request: bool = False,
) -> dict[str, Any]:
snapshots = self._collect_candidate_snapshots(
candidate_keys=candidate_keys,
fallback_from_request=fallback_from_request,
)
metadata: dict[str, Any] = {}
if pool_summary:
metadata["pool_summary"] = pool_summary
if snapshots:
metadata["candidate_keys"] = snapshots
scheduling_audit = self._build_scheduling_audit(
snapshots,
selected_key_id=selected_key_id,
)
if scheduling_audit:
metadata["scheduling_audit"] = scheduling_audit
return metadata
def _merge_scheduling_metadata(
self,
request_metadata: dict[str, Any] | None,
*,
exec_result: Any | None = None,
selected_key_id: str | None = None,
candidate_keys: list[Any] | None = None,
pool_summary: dict[str, Any] | None = None,
fallback_from_request: bool = True,
) -> dict[str, Any] | None:
merged = dict(request_metadata or {})
resolved_candidate_keys = (
candidate_keys
if candidate_keys is not None
else getattr(exec_result, "candidate_keys", None)
)
resolved_key_id = selected_key_id or getattr(exec_result, "key_id", None)
resolved_pool_summary = (
pool_summary if pool_summary is not None else getattr(exec_result, "pool_summary", None)
)
merged.update(
self._build_scheduling_metadata(
candidate_keys=resolved_candidate_keys,
selected_key_id=resolved_key_id,
pool_summary=resolved_pool_summary,
fallback_from_request=fallback_from_request,
)
)
return merged or None
def _resolve_capability_requirements(
self,
model_name: str,
request_headers: dict[str, str] | None = None,
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""
解析请求的能力需求
来源:
1. 用户模型级配置 (User.model_capability_settings)
2. 用户 API Key 强制配置 (ApiKey.force_capabilities)
3. 请求头 X-Require-Capability
4. Adapter 的 detect_capability_requirements如 Claude 的 anthropic-beta
Args:
model_name: 模型名称
request_headers: 请求头
request_body: 请求体(可选)
Returns:
能力需求字典
"""
from src.services.capability.resolver import CapabilityResolver
return CapabilityResolver.resolve_requirements(
user=self.user,
user_api_key=self.api_key,
model_name=model_name,
request_headers=request_headers,
request_body=request_body,
adapter_detector=self.adapter_detector,
)
async def _resolve_preferred_key_ids(
self,
model_name: str,
request_body: dict[str, Any] | None = None,
) -> list[str] | None:
"""可选的 Key 优先级解析钩子(默认不启用)。"""
return None
def build_provider_payload(
self,
original_body: dict[str, Any],
*,
mapped_model: str | None = None,
) -> dict[str, Any]:
"""构建发送给 Provider 的请求体,替换 model 名称"""
payload = dict(original_body)
if mapped_model:
payload["model"] = mapped_model
return payload
def _create_pending_usage(
self,
model: str,
is_stream: bool,
request_type: str = "chat",
api_format: str | None = None,
request_headers: dict[str, Any] | None = None,
request_body: dict[str, Any] | None = None,
) -> bool:
"""在请求开始时创建 pending 状态的 Usage 记录
让前端可以立即看到"处理中"的请求,提升用户体验。
如果创建失败不影响主流程,仅记录警告日志。
Args:
model: 模型名称
is_stream: 是否为流式请求
request_type: 请求类型chat, video 等)
api_format: API 格式
request_headers: 原始请求头
request_body: 原始请求体
Returns:
bool: True 表示已成功创建False 表示创建失败(调用方可按需回退处理)。
"""
try:
UsageService.create_pending_usage(
db=self.db,
request_id=self.request_id,
user=self.user,
api_key=self.api_key,
model=model,
is_stream=is_stream,
request_type=request_type,
api_format=api_format,
request_headers=request_headers,
request_body=request_body,
)
return True
except Exception as exc:
# 创建失败不影响主流程
logger.warning(f"[{self.request_id}] Failed to create pending usage: {exc}")
return False
def _update_usage_to_streaming(self, request_id: str | None = None) -> None:
"""更新 Usage 状态为 streaming流式传输开始时调用
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输
注意TTFB首字节时间由 StreamContext.record_first_byte_time() 记录,
并在最终 record_success 时传递到数据库,避免重复记录导致数据不一致。
Args:
request_id: 请求 ID如果不传则使用 self.request_id
"""
import asyncio
from src.database.database import get_db
target_request_id = request_id or self.request_id
def _sync_update() -> None:
db_gen = get_db()
db = next(db_gen)
try:
UsageService.update_usage_status(
db=db,
request_id=target_request_id,
status="streaming",
)
finally:
db.close()
async def _do_update() -> None:
try:
await asyncio.to_thread(_sync_update)
except Exception as e:
logger.warning(f"[{target_request_id}] 更新 Usage 状态为 streaming 失败: {e}")
# 创建后台任务,不阻塞当前流
from src.utils.async_utils import safe_create_task
safe_create_task(_do_update())
def _update_usage_to_streaming_with_ctx(self, ctx: StreamContext) -> None:
"""更新 Usage 状态为 streaming同时更新 provider 相关信息
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输
注意TTFB首字节时间由 StreamContext.record_first_byte_time() 记录,
并在最终 record_success 时传递到数据库,避免重复记录导致数据不一致。
Args:
ctx: 流式上下文,包含 provider 相关信息
"""
import asyncio
from src.database.database import get_db
target_request_id = self.request_id
provider = ctx.provider_name
target_model = ctx.mapped_model
provider_id = ctx.provider_id
endpoint_id = ctx.endpoint_id
key_id = ctx.key_id
first_byte_time_ms = ctx.first_byte_time_ms
api_format = ctx.api_format
# 格式转换追踪
endpoint_api_format = ctx.provider_api_format or None
has_format_conversion = ctx.has_format_conversion
# 如果 provider 为空,记录警告(不应该发生,但用于调试)
if not provider:
logger.warning(
f"[{target_request_id}] 更新 streaming 状态时 provider 为空: "
f"ctx.provider_name={ctx.provider_name}, ctx.provider_id={ctx.provider_id}"
)
# Capture mutable ctx attrs before handing off to thread
provider_request_headers = ctx.provider_request_headers or None
provider_request_body = ctx.provider_request_body
def _sync_update() -> None:
db_gen = get_db()
db = next(db_gen)
try:
UsageService.update_usage_status(
db=db,
request_id=target_request_id,
status="streaming",
provider=provider,
target_model=target_model,
provider_id=provider_id,
provider_endpoint_id=endpoint_id,
provider_api_key_id=key_id,
first_byte_time_ms=first_byte_time_ms,
api_format=api_format,
endpoint_api_format=endpoint_api_format,
has_format_conversion=has_format_conversion,
provider_request_headers=provider_request_headers,
provider_request_body=provider_request_body,
)
finally:
db.close()
async def _do_update() -> None:
try:
await asyncio.to_thread(_sync_update)
except Exception as e:
logger.warning(f"[{target_request_id}] 更新 Usage 状态为 streaming 失败: {e}")
# 创建后台任务,不阻塞当前流
from src.utils.async_utils import safe_create_task
safe_create_task(_do_update())
def _log_request_error(self, message: str, error: Exception) -> None:
"""记录请求错误日志,对业务异常不打印堆栈
Args:
message: 错误消息前缀
error: 异常对象
"""
from src.core.exceptions import (
BalanceInsufficientException,
ModelNotSupportedException,
ProviderException,
RateLimitException,
UpstreamClientException,
)
if isinstance(
error,
(
ProviderException,
BalanceInsufficientException,
RateLimitException,
ModelNotSupportedException,
UpstreamClientException,
),
):
# 业务异常:简洁日志,不打印堆栈
logger.error(f"{message}: [{type(error).__name__}] {error}")
else:
# 未知异常:完整堆栈
logger.exception(f"{message}: {error}")
# ============================================================================
# 客户端断连检测
# ============================================================================
class ClientDisconnectedException(Exception):
"""客户端在等待首字节时断开连接"""
pass
_T = TypeVar("_T")
async def wait_for_with_disconnect_detection(
coro: Coroutine[Any, Any, _T],
timeout: float,
is_disconnected: Callable[[], Awaitable[bool]],
request_id: str,
check_interval: float = 0.5,
) -> _T:
"""
等待协程完成,同时检测客户端断连
在等待上游响应(如首字节)时,定期检测客户端是否已断连。
若检测到断连,取消任务并抛出 ClientDisconnectedException。
Args:
coro: 要等待的协程
timeout: 超时时间(秒)
is_disconnected: 异步断连检测函数(如 http_request.is_disconnected
request_id: 请求 ID用于日志
check_interval: 断连检测间隔(秒),默认 0.5s
Returns:
协程的返回值
Raises:
ClientDisconnectedException: 客户端断连
asyncio.TimeoutError: 超时
asyncio.CancelledError: 任务被外部取消
"""
task = asyncio.create_task(coro)
client_disconnected = False
async def check_client_disconnect() -> None:
nonlocal client_disconnected
while not task.done():
await asyncio.sleep(check_interval)
try:
if await is_disconnected():
client_disconnected = True
logger.debug(f" [{request_id}] 检测到客户端断连,取消预取任务")
task.cancel()
break
except Exception as e:
logger.debug(f" [{request_id}] 断连检测异常: {e}")
disconnect_task = asyncio.create_task(check_client_disconnect())
try:
return await asyncio.wait_for(task, timeout=timeout)
except asyncio.CancelledError:
if client_disconnected:
raise ClientDisconnectedException("Client disconnected during prefetch")
raise
finally:
disconnect_task.cancel()
try:
await disconnect_task
except asyncio.CancelledError:
pass

View File

@@ -0,0 +1,344 @@
"""
Chat Adapter 通用基类
提供 Chat 格式(进行请求验证和标准化)的通用适配器逻辑:
- 请求解析和验证Pydantic
- 审计日志记录
- Handler 创建和调用
公共逻辑(异常处理、计费、头部构建等)继承自 HandlerAdapterBase。
计费策略、模型抓取与 provider 格式能力由 `core.api_format` 注册表统一提供。
子类只需提供:
- FORMAT_ID: API 格式标识
- HANDLER_CLASS: 对应的 ChatHandlerBase 子类
- _validate_request_body(): 请求验证逻辑
"""
from __future__ import annotations
from abc import abstractmethod
from typing import Any
from fastapi import HTTPException
from fastapi.responses import JSONResponse
from sqlalchemy.orm import Session
from src.api.base.adapter import ApiMode
from src.api.base.context import ApiRequestContext
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.api.handlers.base.handler_adapter_base import HandlerAdapterBase
from src.core.exceptions import (
BalanceInsufficientException,
InvalidRequestException,
ModelNotSupportedException,
ProviderAuthException,
ProviderNotAvailableException,
ProviderRateLimitException,
ProviderTimeoutException,
UpstreamClientException,
)
from src.core.logger import logger
class ChatAdapterBase(HandlerAdapterBase):
"""
Chat Adapter 通用基类
提供 Chat 格式的通用适配器逻辑,子类只需配置:
- FORMAT_ID: API 格式标识
- HANDLER_CLASS: ChatHandlerBase 子类
- name: 适配器名称
"""
HANDLER_CLASS: type[ChatHandlerBase]
# 适配器配置
name: str = "chat.base"
mode = ApiMode.STANDARD
eager_request_body = False
async def handle(self, context: ApiRequestContext) -> Any:
"""处理 Chat API 请求"""
http_request = context.request
user = context.user
api_key = context.api_key
db = context.db
request_id = context.request_id
balance_remaining_value = context.balance_remaining
start_time = context.start_time
client_ip = context.client_ip
user_agent = context.user_agent
original_headers = context.original_headers
query_params = context.query_params
original_request_body = await context.ensure_json_body_async()
# 合并 path_params 到请求体(如 Gemini API 的 model 在 URL 路径中)
if context.path_params:
original_request_body = self._merge_path_params(
original_request_body, context.path_params
)
# 验证和解析请求
request_obj = self._validate_request_body(original_request_body, context.path_params)
if isinstance(request_obj, JSONResponse):
return request_obj
stream = getattr(request_obj, "stream", False)
model = getattr(request_obj, "model", "unknown")
# 添加审计元数据
audit_metadata = self._build_audit_metadata(original_request_body, request_obj)
context.add_audit_metadata(**audit_metadata)
# 格式化额度显示
balance_display = (
"unlimited" if balance_remaining_value is None else f"${balance_remaining_value:.2f}"
)
# 请求开始日志
logger.info(
f"[REQ] {request_id[:8]} | {self.FORMAT_ID} | {getattr(api_key, 'name', 'unknown')} | "
f"{model} | {'stream' if stream else 'sync'} | balance:{balance_display}"
)
try:
# 检查客户端连接
if await http_request.is_disconnected():
logger.warning("客户端连接断开")
raise HTTPException(status_code=499, detail="Client disconnected")
# 创建 Handler
handler = self._create_handler(
db=db,
user=user,
api_key=api_key,
request_id=request_id,
client_ip=client_ip,
user_agent=user_agent,
start_time=start_time,
perf_metrics=context.extra.get("perf"),
api_family=self.API_FAMILY.value if self.API_FAMILY else None,
endpoint_kind=self.ENDPOINT_KIND.value if self.ENDPOINT_KIND else None,
)
# 处理请求
if stream:
return await handler.process_stream(
request=request_obj,
http_request=http_request,
original_headers=original_headers,
original_request_body=original_request_body,
query_params=query_params,
client_content_encoding=context.client_content_encoding,
)
return await handler.process_sync(
request=request_obj,
http_request=http_request,
original_headers=original_headers,
original_request_body=original_request_body,
query_params=query_params,
client_content_encoding=context.client_content_encoding,
client_accept_encoding=context.client_accept_encoding,
)
except HTTPException:
raise
except (
ModelNotSupportedException,
BalanceInsufficientException,
InvalidRequestException,
) as e:
logger.info(f"客户端请求错误: {e.error_type}")
return self._error_response(
status_code=e.status_code,
error_type=("invalid_request_error" if e.status_code == 400 else "quota_exceeded"),
message=e.message,
)
except (
ProviderAuthException,
ProviderRateLimitException,
ProviderNotAvailableException,
ProviderTimeoutException,
UpstreamClientException,
) as e:
return await self._handle_provider_exception(
e,
db=db,
user=user,
api_key=api_key,
model=model,
stream=stream,
start_time=start_time,
original_headers=original_headers,
original_request_body=original_request_body,
client_ip=client_ip,
request_id=request_id,
)
except Exception as e:
return await self._handle_unexpected_exception(
e,
db=db,
user=user,
api_key=api_key,
model=model,
stream=stream,
start_time=start_time,
original_headers=original_headers,
original_request_body=original_request_body,
client_ip=client_ip,
request_id=request_id,
)
def _create_handler(
self,
*,
db: Session,
user: Any,
api_key: Any,
request_id: str,
client_ip: str,
user_agent: str,
start_time: float,
perf_metrics: dict[str, Any] | None = None,
api_family: str | None = None,
endpoint_kind: str | None = None,
) -> Any:
"""创建 Handler 实例 - 子类可覆盖"""
return self.HANDLER_CLASS(
db=db,
user=user,
api_key=api_key,
request_id=request_id,
client_ip=client_ip,
user_agent=user_agent,
start_time=start_time,
allowed_api_formats=self.allowed_api_formats,
adapter_detector=self.detect_capability_requirements,
perf_metrics=perf_metrics,
api_family=api_family,
endpoint_kind=endpoint_kind,
)
@abstractmethod
def _validate_request_body(
self, original_request_body: dict, path_params: dict | None = None
) -> None:
"""验证请求体 - 子类必须实现"""
pass
def _extract_message_count(self, payload: dict[str, Any], request_obj: Any) -> int:
"""提取消息数量 - 子类可覆盖"""
messages = payload.get("messages", [])
if hasattr(request_obj, "messages"):
messages = request_obj.messages
return len(messages) if isinstance(messages, list) else 0
def _build_audit_metadata(self, payload: dict[str, Any], request_obj: Any) -> dict[str, Any]:
"""构建审计日志元数据 - 子类可覆盖"""
model = getattr(request_obj, "model", payload.get("model", "unknown"))
stream = getattr(request_obj, "stream", payload.get("stream", False))
messages_count = self._extract_message_count(payload, request_obj)
return {
"action": f"{self.FORMAT_ID.lower()}_request",
"model": model,
"stream": bool(stream),
"max_tokens": getattr(request_obj, "max_tokens", payload.get("max_tokens")),
"messages_count": messages_count,
"temperature": getattr(request_obj, "temperature", payload.get("temperature")),
"top_p": getattr(request_obj, "top_p", payload.get("top_p")),
}
# =========================================================================
# Adapter 注册表
# =========================================================================
_ADAPTER_REGISTRY: dict[str, type[ChatAdapterBase]] = {}
_ADAPTERS_LOADED = False
def register_adapter(adapter_class: type[ChatAdapterBase]) -> type[ChatAdapterBase]:
"""
注册 Adapter 类到注册表
用法:
@register_adapter
class ClaudeChatAdapter(ChatAdapterBase):
FORMAT_ID = "CLAUDE"
...
Args:
adapter_class: Adapter 类
Returns:
注册的 Adapter 类(支持作为装饰器使用)
"""
format_id = adapter_class.FORMAT_ID
if format_id and format_id != "UNKNOWN":
_ADAPTER_REGISTRY[format_id.upper()] = adapter_class
return adapter_class
def _ensure_adapters_loaded() -> None:
"""确保所有 Adapter 已被加载(触发注册)"""
global _ADAPTERS_LOADED
if _ADAPTERS_LOADED:
return
# 导入各个 Adapter 模块以触发 @register_adapter 装饰器
try:
from src.api.handlers.claude import adapter as _ # noqa: F401
except ImportError:
pass
try:
from src.api.handlers.openai import adapter as _ # noqa: F401
except ImportError:
pass
try:
from src.api.handlers.gemini import adapter as _ # noqa: F401
except ImportError:
pass
_ADAPTERS_LOADED = True
def get_adapter_class(api_format: str) -> type[ChatAdapterBase] | None:
"""
根据 API format 获取 Adapter 类
Args:
api_format: API 格式标识(如 "openai:chat", "claude:chat", "gemini:chat"
Returns:
对应的 Adapter 类,如果未找到返回 None
"""
_ensure_adapters_loaded()
return _ADAPTER_REGISTRY.get(api_format.upper()) if api_format else None
def get_adapter_instance(api_format: str) -> ChatAdapterBase | None:
"""
根据 API format 获取 Adapter 实例
Args:
api_format: API 格式标识
Returns:
Adapter 实例,如果未找到返回 None
"""
adapter_class = get_adapter_class(api_format)
if adapter_class:
return adapter_class()
return None
def list_registered_formats() -> list[str]:
"""返回所有已注册的 API 格式"""
_ensure_adapters_loaded()
return list(_ADAPTER_REGISTRY.keys())

View File

@@ -0,0 +1,156 @@
"""
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.provider.adapters.vertex_ai.transport import (
get_effective_format as get_vertex_ai_effective_format,
)
from src.services.scheduling.aware_scheduler import ProviderCandidate
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_dynamic_format(
key: ProviderAPIKey,
auth_info: Any,
model: str,
provider_api_format: str,
client_api_format: str,
candidate: ProviderCandidate | None,
) -> tuple[str, bool]:
"""
解析动态格式并计算 needs_conversion
对于 Vertex AI 等跨格式 Provider同一个项目可以访问 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) 元组
"""
from src.core.provider_types import ProviderType
# 判断是否为 Vertex AI provider基于 provider_type 而非 auth_type
provider = getattr(key, "provider", None)
provider_type = getattr(provider, "provider_type", None) if provider else None
if provider_type == ProviderType.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)

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,971 @@
"""
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 (
build_json_response_for_client,
filter_proxy_response_headers,
get_format_converter_registry,
resolve_client_accept_encoding,
resolve_client_content_encoding,
)
from src.config.settings import config
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
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
PreparedExecutionPlan,
build_execution_plan_body,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
from src.services.task.request_state import MutableRequestBodyState
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.scheduling.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
provider_response_json: dict[str, Any] | None = None # 格式转换前的提供商原始响应
pool_summary: 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,
client_content_encoding: str | None = None,
client_accept_encoding: str | None = None,
) -> JSONResponse:
"""处理非流式响应(原 process_sync 的完整逻辑)"""
handler = self._handler
logger.debug(f"开始非流式响应处理 ({handler.FORMAT_ID})")
effective_client_content_encoding = resolve_client_content_encoding(
original_headers,
client_content_encoding,
)
effective_client_accept_encoding = resolve_client_accept_encoding(
original_headers,
client_accept_encoding,
)
# 转换请求格式
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 记录,让前端可以立即看到"处理中"
pending_usage_created = 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,
)
request_state = MutableRequestBodyState(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_state=request_state,
query_params=query_params,
client_content_encoding=effective_client_content_encoding,
)
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.core.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_state=request_state,
request_headers=original_headers,
request_body=original_request_body,
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
create_pending_usage=not pending_usage_created,
)
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)
cache_creation_tokens_5m = usage_info.get("cache_creation_input_tokens_5m", 0)
cache_creation_tokens_1h = usage_info.get("cache_creation_input_tokens_1h", 0)
# 非流式成功时,返回给客户端的是提供商响应头(透传)
# JSONResponse 会自动设置 content-type但我们记录实际返回的完整头
client_response_headers = filter_proxy_response_headers(ctx.response_headers)
client_response_headers["content-type"] = "application/json"
client_response = build_json_response_for_client(
status_code=ctx.status_code,
content=ctx.response_json,
headers=client_response_headers,
client_accept_encoding=effective_client_accept_encoding,
)
actual_client_response_headers = dict(client_response.headers)
request_metadata = handler._build_request_metadata() or {}
if ctx.sync_proxy_info:
request_metadata["proxy"] = ctx.sync_proxy_info
request_metadata = handler._merge_scheduling_metadata(
request_metadata,
exec_result=exec_result,
selected_key_id=ctx.key_id,
)
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=original_request_body,
response_headers=ctx.response_headers,
client_response_headers=actual_client_response_headers,
response_body=ctx.provider_response_json or ctx.response_json,
client_response_body=ctx.response_json if ctx.provider_response_json else None,
provider_request_body=ctx.provider_request_body,
cache_creation_tokens=cache_creation_tokens,
cache_read_tokens=cached_tokens,
cache_creation_tokens_5m=cache_creation_tokens_5m,
cache_creation_tokens_1h=cache_creation_tokens_1h,
is_stream=False,
provider_request_headers=ctx.provider_request_headers,
api_format=api_format,
api_family=handler.api_family,
endpoint_kind=handler.endpoint_kind,
# 格式转换追踪
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,
)
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 client_response
except ThinkingSignatureException as e:
# Thinking 签名错误TaskService 层已处理整流重试但仍失败
# 记录实际发送给 Provider 的请求体,便于排查问题根因
response_time_ms = handler.elapsed_ms()
request_metadata = handler._build_request_metadata() or {}
if ctx.sync_proxy_info:
request_metadata["proxy"] = ctx.sync_proxy_info
request_metadata = handler._merge_scheduling_metadata(
request_metadata,
selected_key_id=ctx.key_id,
pool_summary=ctx.pool_summary,
fallback_from_request=True,
)
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=original_request_body,
provider_request_body=ctx.provider_request_body,
error_message=str(e),
is_stream=False,
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
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 build_json_response_for_client(
status_code=_get_error_status_code(e),
content=payload,
headers={"content-type": "application/json"},
client_accept_encoding=effective_client_accept_encoding,
)
except UpstreamClientException as e:
response_time_ms = handler.elapsed_ms()
request_metadata = handler._build_request_metadata() or {}
if ctx.sync_proxy_info:
request_metadata["proxy"] = ctx.sync_proxy_info
request_metadata = handler._merge_scheduling_metadata(
request_metadata,
selected_key_id=ctx.key_id,
pool_summary=ctx.pool_summary,
fallback_from_request=True,
)
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,
)
error_response = build_json_response_for_client(
status_code=_get_error_status_code(e),
content=payload,
headers={"content-type": "application/json"},
client_accept_encoding=effective_client_accept_encoding,
)
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=original_request_body,
provider_request_body=ctx.provider_request_body,
error_message=str(e),
is_stream=False,
api_format=api_format,
api_family=handler.api_family,
endpoint_kind=handler.endpoint_kind,
provider_request_headers=ctx.provider_request_headers,
response_headers=ctx.response_headers,
client_response_headers=dict(error_response.headers),
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
# 格式转换追踪
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,
)
return error_response
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
# 尝试从异常中提取响应头
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
request_metadata = handler._merge_scheduling_metadata(
request_metadata,
selected_key_id=ctx.key_id,
pool_summary=ctx.pool_summary,
fallback_from_request=True,
)
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=original_request_body,
provider_request_body=ctx.provider_request_body,
is_stream=False,
api_format=api_format,
api_family=handler.api_family,
endpoint_kind=handler.endpoint_kind,
provider_request_headers=ctx.provider_request_headers,
response_headers=error_response_headers,
# 非流式失败返回给客户端的是 JSON 错误响应
client_response_headers={"content-type": "application/json"},
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
# 格式转换追踪
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_state: MutableRequestBodyState,
query_params: dict[str, str] | None = None,
client_content_encoding: str | None = None,
) -> dict[str, Any]:
"""单次同步请求(原 sync_request_func 内嵌函数)"""
prepared_plan = await self._build_sync_execution_plan(
provider,
endpoint,
key,
candidate,
model=model,
api_format=api_format,
original_headers=original_headers,
request_state=request_state,
query_params=query_params,
client_content_encoding=client_content_encoding,
)
return await self._execute_sync_plan(
prepared_plan=prepared_plan,
provider=provider,
model=model,
)
async def _build_sync_execution_plan(
self,
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
*,
model: str,
api_format: Any,
original_headers: dict[str, Any],
request_state: MutableRequestBodyState,
query_params: dict[str, str] | None = None,
client_content_encoding: str | None = None,
) -> PreparedExecutionPlan:
"""构建可序列化的执行计划,并保留本地执行所需的运行时上下文。"""
handler = self._handler
ctx = self._ctx
ctx.provider_name = str(provider.name)
ctx.provider_id = str(provider.id)
ctx.endpoint_id = str(endpoint.id)
ctx.key_id = str(key.id)
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,
working_request_body=request_state.build_attempt_body(),
original_headers=original_headers,
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
tls_profile = prep.tls_profile
# 构建请求(上游始终使用 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,
envelope=envelope,
provider_api_format=prep.provider_api_format,
)
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 (
build_proxy_url_async,
get_proxy_label,
get_system_proxy_config_async,
resolve_effective_proxy,
resolve_proxy_info_async,
)
_effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
ctx.sync_proxy_info = await resolve_proxy_info_async(_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)}")
# 解析 delegate 配置,用于后续本地执行或交给 Rust executor。
from src.services.proxy_node.resolver import (
resolve_delegate_config_async,
)
# 非流式请求使用 http_request_timeout 作为整体超时
# 优先使用 Provider 配置,否则使用全局配置
request_timeout = provider.request_timeout or config.http_request_timeout
delegate_cfg = await resolve_delegate_config_async(_effective_proxy)
is_tunnel_delegate = bool(delegate_cfg and delegate_cfg.get("tunnel"))
effective_proxy_for_contract = _effective_proxy
if not effective_proxy_for_contract or not effective_proxy_for_contract.get(
"enabled", True
):
effective_proxy_for_contract = await get_system_proxy_config_async()
proxy_url: str | None = None
if effective_proxy_for_contract and not is_tunnel_delegate:
proxy_url = await build_proxy_url_async(effective_proxy_for_contract)
return PreparedExecutionPlan(
contract=ExecutionPlan(
request_id=str(handler.request_id or ""),
candidate_id=str(
getattr(candidate, "request_candidate_id", "")
or getattr(candidate, "id", "")
or ""
)
or None,
provider_name=str(provider.name),
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
method="POST",
url=url,
headers=dict(provider_hdrs),
body=build_execution_plan_body(
provider_payload,
content_type=str(provider_hdrs.get("content-type") or "").strip() or None,
),
stream=upstream_is_stream,
provider_api_format=provider_api_format,
client_api_format=client_api_format,
model_name=str(model or ""),
content_type=str(provider_hdrs.get("content-type") or "").strip() or None,
content_encoding=client_content_encoding,
proxy=ExecutionProxySnapshot.from_proxy_info(
ctx.sync_proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if is_tunnel_delegate else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None
if is_tunnel_delegate
else None
),
),
tls_profile=tls_profile,
timeouts=ExecutionPlanTimeouts(
connect_ms=int(config.http_connect_timeout * 1000),
read_ms=int(config.http_read_timeout * 1000),
write_ms=int(config.http_write_timeout * 1000),
pool_ms=int(config.http_pool_timeout * 1000),
total_ms=int(request_timeout * 1000),
),
),
payload=provider_payload,
headers=dict(provider_hdrs),
upstream_is_stream=upstream_is_stream,
needs_conversion=needs_conversion,
provider_type=provider_type,
request_timeout=request_timeout,
delegate_config=delegate_cfg,
proxy_config=_effective_proxy,
envelope=envelope,
selected_base_url=selected_base_url_cached,
client_content_encoding=client_content_encoding,
proxy_info=ctx.sync_proxy_info,
)
async def _execute_sync_plan(
self,
*,
prepared_plan: PreparedExecutionPlan,
provider: Provider,
model: str,
) -> dict[str, Any]:
handler = self._handler
if not prepared_plan.remote_eligible:
raise ProviderNotAvailableException(
"执行器暂时不可用,请稍后重试",
provider_name=str(provider.name),
upstream_response="execution contract is not eligible for rust executor",
)
try:
rust_result = await ExecutionRuntimeClient().execute_sync_json(prepared_plan.contract)
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError) as exc:
logger.warning(
"[{}] Rust executor unavailable: {}",
handler.request_id,
exc,
)
raise ProviderNotAvailableException(
"执行器暂时不可用,请稍后重试",
provider_name=str(provider.name),
upstream_response=str(exc),
) from exc
else:
result = await self._finalize_rust_sync_result(
prepared_plan=prepared_plan,
provider=provider,
model=model,
status_code=rust_result.status_code,
response_headers=rust_result.headers,
response_json=rust_result.response_json,
provider_response_json=rust_result.provider_response_json,
response_body_bytes=rust_result.response_body_bytes,
)
logger.debug(
"[{}] sync chat 请求由 Rust executor 执行完成",
handler.request_id,
)
return result
async def _finalize_rust_sync_result(
self,
*,
prepared_plan: PreparedExecutionPlan,
provider: Provider,
model: str,
status_code: int,
response_headers: dict[str, str],
response_json: dict[str, Any] | None,
provider_response_json: dict[str, Any] | None = None,
response_body_bytes: bytes | None = None,
) -> dict[str, Any]:
ctx = self._ctx
synthetic_content = response_body_bytes
if synthetic_content is None:
synthetic_content = json.dumps(
response_json or {},
ensure_ascii=False,
).encode("utf-8")
request = httpx.Request(
prepared_plan.contract.method,
prepared_plan.contract.url,
headers=prepared_plan.headers,
)
synthetic_response = httpx.Response(
status_code,
request=request,
headers=response_headers,
content=synthetic_content,
)
ctx.status_code = status_code
ctx.response_headers = dict(synthetic_response.headers)
extract_proxy_timing(ctx.sync_proxy_info, ctx.response_headers)
if prepared_plan.envelope:
prepared_plan.envelope.on_http_status(
base_url=prepared_plan.selected_base_url,
status_code=ctx.status_code,
)
if status_code >= 400:
error = httpx.HTTPStatusError(
f"Upstream status error: {status_code}",
request=request,
response=synthetic_response,
)
error_body = ""
try:
if prepared_plan.envelope and hasattr(prepared_plan.envelope, "extract_error_text"):
error_body = await prepared_plan.envelope.extract_error_text(synthetic_response)
else:
error_body = synthetic_response.text[:4000] if synthetic_response.text else ""
except Exception:
error_body = synthetic_response.text[:4000] if synthetic_response.text else ""
error.upstream_response = error_body[:4000] # type: ignore[attr-defined]
raise error
if prepared_plan.upstream_is_stream:
if response_body_bytes is None:
raise ExecutionRuntimeClientError("Rust executor stream result must contain body bytes")
return await self._finalize_rust_stream_sync_result(
prepared_plan=prepared_plan,
provider=provider,
model=model,
response_body_bytes=response_body_bytes,
)
if response_json is None:
raise ExecutionRuntimeClientError("Rust executor sync result must contain response_json")
ctx.response_json = dict(response_json)
if provider_response_json is not None:
ctx.provider_response_json = dict(provider_response_json)
if prepared_plan.envelope:
ctx.response_json = prepared_plan.envelope.unwrap_response(ctx.response_json)
prepared_plan.envelope.postprocess_unwrapped_response(
model=model,
data=ctx.response_json,
)
if isinstance(ctx.response_json, dict):
parser = get_parser_for_format(ctx.provider_api_format_for_error or "")
if parser.is_error_response(ctx.response_json):
parsed = parser.parse_response(ctx.response_json, status_code)
raise EmbeddedErrorException(
provider_name=str(provider.name),
error_code=parsed.embedded_status_code,
error_message=parsed.error_message,
error_status=parsed.error_type,
)
if ctx.needs_conversion_for_error and isinstance(ctx.response_json, dict):
ctx.provider_response_json = ctx.response_json.copy()
registry = get_format_converter_registry()
ctx.response_json = registry.convert_response(
ctx.response_json,
ctx.provider_api_format_for_error or "",
ctx.client_api_format_for_error or "",
requested_model=model,
)
return ctx.response_json if isinstance(ctx.response_json, dict) else {}
async def _finalize_rust_stream_sync_result(
self,
*,
prepared_plan: PreparedExecutionPlan,
provider: Provider,
model: str,
response_body_bytes: bytes,
) -> dict[str, Any]:
async def _iter_body() -> Any:
yield response_body_bytes
return await self._aggregate_upstream_stream_response(
byte_iter=_iter_body(),
prepared_plan=prepared_plan,
provider=provider,
model=model,
)
async def _aggregate_upstream_stream_response(
self,
*,
byte_iter: Any,
prepared_plan: PreparedExecutionPlan,
provider: Provider,
model: str,
) -> dict[str, Any]:
ctx = self._ctx
provider_parser = (
get_parser_for_format(ctx.provider_api_format_for_error)
if ctx.provider_api_format_for_error
else None
)
if (
prepared_plan.provider_type == "kiro"
and prepared_plan.envelope
and prepared_plan.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=ctx.provider_api_format_for_error or "",
provider_name=str(provider.name),
model=str(model or ""),
request_id=str(self._handler.request_id or ""),
envelope=prepared_plan.envelope,
provider_parser=provider_parser,
)
registry = get_format_converter_registry()
tgt_norm = (
registry.get_normalizer(ctx.client_api_format_for_error)
if ctx.client_api_format_for_error
else None
)
if tgt_norm is None:
raise RuntimeError(f"未注册 Normalizer: {ctx.client_api_format_for_error}")
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 {}
return ctx.response_json
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
# 失败时返回给客户端的是 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}
stream_fail_metadata = handler._merge_scheduling_metadata(
stream_fail_metadata,
selected_key_id=ctx.key_id,
pool_summary=ctx.pool_summary,
fallback_from_request=True,
)
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=original_request_body,
provider_request_body=ctx.provider_request_body,
is_stream=True,
api_format=ctx.api_format,
api_family=handler.api_family,
endpoint_kind=handler.endpoint_kind,
provider_request_headers=ctx.provider_request_headers,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
# 格式转换追踪
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,
*,
envelope: Any = None,
) -> str:
"""从 HTTP 错误中提取错误文本"""
if envelope and hasattr(envelope, "extract_error_text"):
return await envelope.extract_error_text(e)
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,313 @@
"""
CLI Adapter 通用基类
提供 CLI 格式(直接透传请求)的通用适配器逻辑:
- 请求解析和验证
- 审计日志记录
- Handler 创建和调用
公共逻辑(异常处理、计费、头部构建等)继承自 HandlerAdapterBase。
计费策略、模型抓取与 provider 格式能力由 `core.api_format` 注册表统一提供。
子类只需提供:
- FORMAT_ID: API 格式标识
- HANDLER_CLASS: 对应的 MessageHandler 类
- 可选覆盖 _extract_message_count() 自定义消息计数逻辑
"""
from __future__ import annotations
from typing import Any
from fastapi import HTTPException
from src.api.base.adapter import ApiMode
from src.api.base.context import ApiRequestContext
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
from src.api.handlers.base.handler_adapter_base import HandlerAdapterBase
from src.core.api_format import EndpointKind
from src.core.exceptions import (
BalanceInsufficientException,
InvalidRequestException,
ModelNotSupportedException,
ProviderAuthException,
ProviderNotAvailableException,
ProviderRateLimitException,
ProviderTimeoutException,
UpstreamClientException,
)
from src.core.logger import logger
class CliAdapterBase(HandlerAdapterBase):
"""
CLI Adapter 通用基类
提供 CLI 格式的通用适配器逻辑,子类只需配置:
- FORMAT_ID: API 格式标识
- HANDLER_CLASS: MessageHandler 类
- name: 适配器名称
"""
HANDLER_CLASS: type[CliMessageHandlerBase]
# CLI 端点类型覆盖(基类默认 CHAT
ENDPOINT_KIND = EndpointKind.CLI
# 适配器配置
name: str = "cli.base"
mode = ApiMode.PROXY
eager_request_body = False
async def handle(self, context: ApiRequestContext) -> Any:
"""处理 CLI API 请求"""
http_request = context.request
user = context.user
api_key = context.api_key
db = context.db
request_id = context.request_id
balance_remaining_value = context.balance_remaining
start_time = context.start_time
client_ip = context.client_ip
user_agent = context.user_agent
original_headers = context.original_headers
query_params = context.query_params
# Store original headers for downstream envelope checks (e.g. CLI-only restriction).
# Only relevant for Claude Code CLI format; skip for others to avoid unnecessary coupling.
if self.FORMAT_ID == "claude:cli":
from src.services.provider.adapters.claude_code.client_restriction import (
set_original_request_headers,
)
set_original_request_headers(original_headers)
original_request_body = await context.ensure_json_body_async()
# 合并 path_params 到请求体(如 Gemini API 的 model 在 URL 路径中)
if context.path_params:
original_request_body = self._merge_path_params(
original_request_body, context.path_params
)
# 获取 stream优先从请求体其次从 path_params如 Gemini 通过 URL 端点区分)
stream = original_request_body.get("stream")
if stream is None and context.path_params:
stream = context.path_params.get("stream", False)
stream = bool(stream)
# 获取 model优先从请求体其次从 path_params如 Gemini 的 model 在 URL 路径中)
model = original_request_body.get("model")
if model is None and context.path_params:
model = context.path_params.get("model", "unknown")
model = model or "unknown"
# 提取请求元数据
audit_metadata = self._build_audit_metadata(original_request_body, context.path_params)
context.add_audit_metadata(**audit_metadata)
# 格式化额度显示
balance_display = (
"unlimited" if balance_remaining_value is None else f"${balance_remaining_value:.2f}"
)
# 请求开始日志
logger.info(
f"[REQ] {request_id[:8]} | {self.FORMAT_ID} | {getattr(api_key, 'name', 'unknown')} | "
f"{model} | {'stream' if stream else 'sync'} | balance:{balance_display}"
)
try:
# 检查客户端连接
if await http_request.is_disconnected():
logger.warning("客户端连接断开")
raise HTTPException(status_code=499, detail="Client disconnected")
# 创建 Handler
handler = self.HANDLER_CLASS(
db=db,
user=user,
api_key=api_key,
request_id=request_id,
client_ip=client_ip,
user_agent=user_agent,
start_time=start_time,
allowed_api_formats=self.allowed_api_formats,
adapter_detector=self.detect_capability_requirements,
perf_metrics=context.extra.get("perf"),
api_family=self.API_FAMILY.value if self.API_FAMILY else None,
endpoint_kind=self.ENDPOINT_KIND.value if self.ENDPOINT_KIND else None,
)
# 处理请求
if stream:
return await handler.process_stream(
original_request_body=original_request_body,
original_headers=original_headers,
query_params=query_params,
path_params=context.path_params,
http_request=http_request,
client_content_encoding=context.client_content_encoding,
)
return await handler.process_sync(
original_request_body=original_request_body,
original_headers=original_headers,
query_params=query_params,
path_params=context.path_params,
client_content_encoding=context.client_content_encoding,
client_accept_encoding=context.client_accept_encoding,
)
except HTTPException:
raise
except (
ModelNotSupportedException,
BalanceInsufficientException,
InvalidRequestException,
) as e:
logger.debug("客户端请求错误: {}", e.error_type)
return self._error_response(
status_code=e.status_code,
error_type=("invalid_request_error" if e.status_code == 400 else "quota_exceeded"),
message=e.message,
)
except (
ProviderAuthException,
ProviderRateLimitException,
ProviderNotAvailableException,
ProviderTimeoutException,
UpstreamClientException,
) as e:
return await self._handle_provider_exception(
e,
db=db,
user=user,
api_key=api_key,
model=model,
stream=stream,
start_time=start_time,
original_headers=original_headers,
original_request_body=original_request_body,
client_ip=client_ip,
request_id=request_id,
)
except Exception as e:
return await self._handle_unexpected_exception(
e,
db=db,
user=user,
api_key=api_key,
model=model,
stream=stream,
start_time=start_time,
original_headers=original_headers,
original_request_body=original_request_body,
client_ip=client_ip,
request_id=request_id,
)
def _extract_message_count(self, payload: dict[str, Any]) -> int:
"""提取消息数量 - 子类可覆盖"""
if "input" not in payload:
return 0
input_data = payload["input"]
if isinstance(input_data, list):
return len(input_data)
if isinstance(input_data, dict) and "messages" in input_data:
return len(input_data.get("messages", []))
return 0
def _build_audit_metadata(
self,
payload: dict[str, Any],
path_params: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""构建审计日志元数据 - 子类可覆盖"""
model = payload.get("model")
if model is None and path_params:
model = path_params.get("model", "unknown")
model = model or "unknown"
stream = payload.get("stream", False)
messages_count = self._extract_message_count(payload)
return {
"action": f"{self.FORMAT_ID.lower()}_request",
"model": model,
"stream": bool(stream),
"max_tokens": payload.get("max_tokens"),
"messages_count": messages_count,
"temperature": payload.get("temperature"),
"top_p": payload.get("top_p"),
"tool_count": len(payload.get("tools") or []),
"instructions_present": bool(payload.get("instructions")),
}
# =========================================================================
# CLI Adapter 注册表
# =========================================================================
_CLI_ADAPTER_REGISTRY: dict[str, type[CliAdapterBase]] = {}
_CLI_ADAPTERS_LOADED = False
def register_cli_adapter(adapter_class: type[CliAdapterBase]) -> type[CliAdapterBase]:
"""
注册 CLI Adapter 类到注册表
用法:
@register_cli_adapter
class ClaudeCliAdapter(CliAdapterBase):
FORMAT_ID = "CLAUDE_CLI"
...
"""
format_id = adapter_class.FORMAT_ID
if format_id and format_id != "UNKNOWN":
_CLI_ADAPTER_REGISTRY[format_id.upper()] = adapter_class
return adapter_class
def _ensure_cli_adapters_loaded() -> None:
"""确保所有 CLI Adapter 已被加载(触发注册)"""
global _CLI_ADAPTERS_LOADED
if _CLI_ADAPTERS_LOADED:
return
try:
from src.api.handlers.claude_cli import adapter as _ # noqa: F401
except ImportError:
pass
try:
from src.api.handlers.openai_cli import adapter as _ # noqa: F401
except ImportError:
pass
try:
from src.api.handlers.gemini_cli import adapter as _ # noqa: F401
except ImportError:
pass
_CLI_ADAPTERS_LOADED = True
def get_cli_adapter_class(api_format: str) -> type[CliAdapterBase] | None:
"""根据 API format 获取 CLI Adapter 类"""
_ensure_cli_adapters_loaded()
return _CLI_ADAPTER_REGISTRY.get(api_format.upper()) if api_format else None
def get_cli_adapter_instance(api_format: str) -> CliAdapterBase | None:
"""根据 API format 获取 CLI Adapter 实例"""
adapter_class = get_cli_adapter_class(api_format)
if adapter_class:
return adapter_class()
return None
def list_registered_cli_formats() -> list[str]:
"""返回所有已注册的 CLI API 格式"""
_ensure_cli_adapters_loaded()
return list(_CLI_ADAPTER_REGISTRY.keys())

View File

@@ -0,0 +1,587 @@
"""CLI Handler - SSE 事件处理 + 格式转换 Mixin"""
from __future__ import annotations
import codecs
import json
from collections.abc import Iterator
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
from src.api.handlers.base.utils import get_format_converter_registry
from src.core.logger import logger
from src.core.usage_tokens import extract_cache_creation_tokens, extract_cache_read_tokens
from src.services.provider.behavior import get_provider_behavior
from src.utils.sse_parser import SSEEventParser
from .cli_sse_helpers import (
_format_converted_events_to_sse,
_parse_gemini_json_array_line,
_parse_sse_data_line,
_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: CliHandlerProtocol,
ctx: StreamContext,
event_name: str | None,
data_str: str,
record_chunk: bool = False,
) -> None:
"""
处理 SSE 事件
通用框架:解析 JSON、更新计数器
子类可覆盖 _process_event_data() 实现格式特定逻辑
Args:
ctx: 流上下文
event_name: 事件名称(如 message_start, content_block_delta 等)
data_str: 事件数据字符串JSON 格式)
record_chunk: 是否记录到 parsed_chunks不需要格式转换时应为 True
当为 True 时,同时更新 data_count
当为 False 时data_count 由 _record_converted_chunks 更新
"""
if not data_str:
return
if data_str == "[DONE]":
ctx.has_completion = True
return
try:
data = json.loads(data_str)
except json.JSONDecodeError:
return
if not isinstance(data, dict):
return
behavior = get_provider_behavior(
provider_type=ctx.provider_type,
endpoint_sig=ctx.provider_api_format,
)
envelope = behavior.envelope
if envelope:
data = envelope.unwrap_response(data)
if not isinstance(data, dict):
return
# 当不需要格式转换时,更新 data_count需要记录时再写入 parsed_chunks。
# 当需要格式转换时record_chunk=Falsedata_count 由 _record_converted_chunks 更新
if record_chunk:
ctx.data_count += 1
if ctx.record_parsed_chunks:
ctx.parsed_chunks.append(data)
else:
# 格式转换场景:保留提供商原始数据
if ctx.record_parsed_chunks:
ctx.provider_parsed_chunks.append(data)
event_type = event_name or data.get("type", "")
if envelope:
envelope.postprocess_unwrapped_response(model=ctx.model, data=data)
# 调用格式特定的处理逻辑
# 注意跨格式转换时_process_event_data 会自动选择正确的 Provider 解析器
self._process_event_data(ctx, event_type, data)
def _process_event_data(
self,
ctx: StreamContext,
event_type: str,
data: dict[str, Any],
) -> None:
"""
处理解析后的事件数据 - 子类应覆盖此方法
默认实现使用 ResponseParser 提取 usage
"""
# 提取 response_id
if not ctx.response_id:
response_obj = data.get("response")
if isinstance(response_obj, dict) and response_obj.get("id"):
ctx.response_id = response_obj["id"]
elif "id" in data:
ctx.response_id = data["id"]
# 使用解析器提取 usage
# Claude/CLI 流式响应的 usage 可能在首个 chunk 或最后一个 chunk 中
# 首个 chunk 可能部分为 0最后一个 chunk 包含完整值,因此取最大值确保正确计费
#
# 重要:当跨格式转换时,收到的数据是 Provider 格式,需要使用 Provider 格式的解析器
# 而不是客户端格式的解析器self.parser
parser = self.parser
if ctx.provider_api_format and ctx.provider_api_format != ctx.client_api_format:
# 跨格式转换:使用 Provider 格式的解析器
try:
provider_parser = get_parser_for_format(ctx.provider_api_format)
if provider_parser:
parser = provider_parser
except KeyError:
logger.debug(
f"[{getattr(ctx, 'request_id', 'unknown')}] 未找到 Provider 格式解析器: "
f"{ctx.provider_api_format}, 回退使用客户端格式解析器"
)
usage = parser.extract_usage_from_response(data)
if usage:
new_input = usage.get("input_tokens", 0)
new_output = usage.get("output_tokens", 0)
new_cached = usage.get("cache_read_tokens", 0)
new_cache_creation = usage.get("cache_creation_tokens", 0)
# 取最大值更新
if new_input > ctx.input_tokens:
ctx.input_tokens = new_input
if new_output > ctx.output_tokens:
ctx.output_tokens = new_output
if new_cached > ctx.cached_tokens:
ctx.cached_tokens = new_cached
if new_cache_creation > ctx.cache_creation_tokens:
ctx.cache_creation_tokens = new_cache_creation
# 保存最后一个非空 usage 作为 final_usage
if any([new_input, new_output, new_cached, new_cache_creation]):
ctx.final_usage = usage
# 提取文本内容(同样使用正确的解析器)
text = parser.extract_text_content(data)
if text:
ctx.append_text(text)
# 检查完成事件
if event_type in ("response.completed", "message_stop"):
ctx.has_completion = True
response_obj = data.get("response")
if isinstance(response_obj, dict):
ctx.final_response = response_obj
def _record_converted_chunks(
self,
ctx: StreamContext,
converted_events: list[dict[str, Any]],
) -> None:
"""
记录转换后的 chunk 数据到 parsed_chunks并更新统计信息
当需要格式转换时,记录的是转换后的数据(客户端实际收到的格式);
同时更新 data_count、has_completion 等统计信息。
重要:此方法也从转换后的事件中提取 usage 信息,作为 _process_event_data
从原始数据提取的补充。这确保即使原始 Provider 数据中没有 usage如 OpenAI
未设置 stream_options也能从转换后的格式中获取。
Args:
ctx: 流上下文
converted_events: 转换后的事件列表
"""
for evt in converted_events:
if isinstance(evt, dict):
ctx.data_count += 1
if ctx.record_parsed_chunks:
ctx.parsed_chunks.append(evt)
# 检测完成事件(根据客户端格式判断)
# OpenAI 格式: choices[].finish_reason
# Claude 格式: type == "message_stop" 或 stop_reason
event_type = evt.get("type", "")
if event_type == "message_stop":
ctx.has_completion = True
elif event_type == "response.completed":
ctx.has_completion = True
elif "choices" in evt:
choices = evt.get("choices", [])
for choice in choices:
if isinstance(choice, dict) and choice.get("finish_reason"):
ctx.has_completion = True
break
# 从转换后的事件中提取 usage补充 _process_event_data 的提取)
# Claude 格式: message_delta.usage 或 message_start.message.usage
# OpenAI 格式: chunk.usage
self._extract_usage_from_converted_event(ctx, evt, event_type)
def _extract_usage_from_converted_event(
self,
ctx: StreamContext,
evt: dict[str, Any],
event_type: str,
) -> None:
"""
从转换后的事件中提取 usage 信息
支持多种格式:
- Claude: message_delta.usage, message_start.message.usage
- OpenAI: chunk.usage
- Gemini: usageMetadata
Args:
ctx: 流上下文
evt: 转换后的事件
event_type: 事件类型
"""
usage: dict[str, Any] | None = None
# Claude 格式: message_delta 或 message_start
if event_type == "message_delta":
usage = evt.get("usage")
elif event_type == "message_start":
message = evt.get("message", {})
if isinstance(message, dict):
usage = message.get("usage")
# OpenAI Responses API (openai:cli) 格式: response.completed 中 usage 嵌套在 response 对象内
elif event_type == "response.completed":
resp_obj = evt.get("response")
if isinstance(resp_obj, dict):
usage = resp_obj.get("usage")
# 兼容: 部分实现可能在顶层也有 usage
if not usage:
usage = evt.get("usage")
# OpenAI Chat 格式: 直接在 chunk 中
elif "usage" in evt:
usage = evt.get("usage")
# Gemini 格式: usageMetadata
elif "usageMetadata" in evt:
meta = evt.get("usageMetadata", {})
if isinstance(meta, dict):
usage = {
"input_tokens": meta.get("promptTokenCount", 0),
"output_tokens": meta.get("candidatesTokenCount", 0),
"cache_read_tokens": meta.get("cachedContentTokenCount", 0),
"cache_creation_tokens": 0, # Gemini 目前不支持缓存创建
}
if usage and isinstance(usage, dict):
new_input = usage.get("input_tokens") or usage.get("prompt_tokens") or 0
new_output = usage.get("output_tokens") or usage.get("completion_tokens") or 0
new_cached = extract_cache_read_tokens(usage)
new_cache_creation = extract_cache_creation_tokens(usage)
# 取最大值更新(与 _process_event_data 相同的策略)
if new_input > ctx.input_tokens:
ctx.input_tokens = new_input
logger.debug("[{}] 从转换后事件更新 input_tokens: {}", ctx.request_id, new_input)
if new_output > ctx.output_tokens:
ctx.output_tokens = new_output
logger.debug("[{}] 从转换后事件更新 output_tokens: {}", ctx.request_id, new_output)
if new_cached > ctx.cached_tokens:
ctx.cached_tokens = new_cached
if new_cache_creation > ctx.cache_creation_tokens:
ctx.cache_creation_tokens = new_cache_creation
# 保存最后一个非空 usage
if any([new_input, new_output, new_cached, new_cache_creation]):
ctx.final_usage = usage
def _flush_buffer_with_conversion(
self: CliHandlerProtocol,
ctx: StreamContext,
buffer: bytes,
decoder: codecs.IncrementalDecoder,
sse_parser: SSEEventParser,
needs_conversion: bool,
) -> Iterator[bytes]:
"""flush 字节 buffer 残余数据 + SSE parser 内部缓冲区,并做格式转换。
正常流结束时调用(区别于异常路径的 _flush_remaining_sse_data
当 needs_conversion=True 时 yield 转换后的 SSE 行;
当 needs_conversion=False 时 yield 原始行透传。
"""
# 1) flush 字节 buffer 中的残余数据(最后一个 chunk 可能不以换行结尾)
if buffer:
try:
remaining = decoder.decode(buffer, True)
except Exception:
remaining = ""
for tail_line in remaining.split("\n"):
stripped = tail_line.rstrip("\r")
tail_events = sse_parser.feed_line(stripped)
if stripped == "":
for event in tail_events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=not needs_conversion,
)
continue
if needs_conversion and stripped:
converted_lines, converted_events = self._convert_sse_line(
ctx, stripped, tail_events
)
self._record_converted_chunks(ctx, converted_events)
for converted_line in converted_lines:
if converted_line:
yield (converted_line + "\n").encode("utf-8")
elif stripped:
# 非 conversion 模式:透传原始行
yield (stripped + "\n").encode("utf-8")
# 2) flush SSE parser 内部缓冲区中的残余事件
for event in sse_parser.flush():
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=not needs_conversion,
)
if needs_conversion:
data_str = event.get("data") or ""
if data_str and data_str != "[DONE]":
flush_line = f"data: {data_str}"
converted_lines, converted_events = self._convert_sse_line(ctx, flush_line, [])
self._record_converted_chunks(ctx, converted_events)
for converted_line in converted_lines:
if converted_line:
yield (converted_line + "\n").encode("utf-8")
def _finalize_stream_metadata(self, ctx: StreamContext) -> None:
"""
在记录统计前从 parsed_chunks 中提取额外的元数据 - 子类可覆盖
这是一个后处理钩子,在流传输完成后、记录 Usage 之前调用。
子类可以覆盖此方法从 ctx.parsed_chunks 中提取格式特定的元数据,
如 Gemini 的 modelVersion、token 统计等。
Args:
ctx: 流上下文,包含 parsed_chunks 和 response_metadata
"""
pass
def _needs_format_conversion(self, ctx: StreamContext) -> bool:
"""
[已废弃] 仅根据格式差异判断是否需要转换
警告:此方法只检查格式是否不同,不检查端点的 format_acceptance_config 配置!
正确的判断应使用候选筛选阶段的结果ctx.needs_conversion该结果由
is_format_compatible() 函数根据全局开关和端点配置计算得出。
此方法保留仅供调试和日志输出使用,流生成器中不应调用此方法。
当 Provider 的 API 格式与客户端请求的 API 格式不同时,需要转换响应。
例如:客户端请求 Claude 格式,但 Provider 返回 OpenAI 格式。
注意:
- CLAUDE 和 CLAUDE_CLI、GEMINI 和 GEMINI_CLI格式相同只是认证不同可透传
- OPENAI 和 OPENAI_CLI格式不同Chat Completions vs Responses API需要转换
"""
from src.core.api_format.metadata import can_passthrough_endpoint
from src.core.api_format.signature import normalize_signature_key
if not ctx.provider_api_format or not ctx.client_api_format:
logger.debug(
f"[{getattr(ctx, 'request_id', 'unknown')}] _needs_format_conversion: "
f"provider_api_format={ctx.provider_api_format!r}, client_api_format={ctx.client_api_format!r} -> False (missing)"
)
return False
provider_format = normalize_signature_key(str(ctx.provider_api_format))
client_format = normalize_signature_key(str(ctx.client_api_format))
# 1. 格式完全匹配 -> 不需要转换
if provider_format == client_format:
logger.debug(
f"[{getattr(ctx, 'request_id', 'unknown')}] _needs_format_conversion: "
f"provider={provider_format}, client={client_format} -> False (exact match)"
)
return False
# 2. 根据 data_format_id 判断是否可透传(可透传则不需要转换)
if can_passthrough_endpoint(client_format, provider_format):
logger.debug(
f"[{getattr(ctx, 'request_id', 'unknown')}] _needs_format_conversion: "
f"provider={provider_format}, client={client_format} -> False (passthroughable)"
)
return False
# 3. 其他情况 -> 需要转换
logger.debug(
f"[{getattr(ctx, 'request_id', 'unknown')}] _needs_format_conversion: "
f"provider={provider_format}, client={client_format} -> True"
)
return True
def _mark_first_output(self, ctx: StreamContext, state: dict[str, bool]) -> None:
"""
标记首次输出:记录 TTFB 并更新 streaming 状态
在第一次 yield 数据前调用,确保:
1. 首字时间 (TTFB) 已记录到 ctx
2. Usage 状态已更新为 streaming包含 provider/key/TTFB 信息)
Args:
ctx: 流上下文
state: 包含 first_yield 和 streaming_updated 的状态字典
"""
if state["first_yield"]:
ctx.record_first_byte_time(self.start_time)
state["first_yield"] = False
if not state["streaming_updated"]:
# 优先使用当前请求的 DB 会话同步更新,避免状态延迟或丢失
try:
from src.services.usage import UsageService
UsageService.update_usage_status(
db=self.db,
request_id=self.request_id,
status="streaming",
provider=ctx.provider_name,
target_model=ctx.mapped_model,
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
first_byte_time_ms=ctx.first_byte_time_ms,
api_format=ctx.api_format,
endpoint_api_format=ctx.provider_api_format or None,
has_format_conversion=ctx.has_format_conversion,
provider_request_headers=ctx.provider_request_headers or None,
provider_request_body=ctx.provider_request_body,
)
except Exception as e:
logger.warning("[{}] 同步更新 streaming 状态失败: {}", self.request_id, e)
# 回退到后台任务更新
self._update_usage_to_streaming_with_ctx(ctx)
state["streaming_updated"] = True
def _convert_sse_line(
self,
ctx: StreamContext,
line: str,
events: list, # noqa: ARG002 - 预留给上下文感知转换
) -> tuple[list[str], list[dict[str, Any]]]:
"""
将 SSE 行从 Provider 格式转换为客户端格式
Args:
ctx: 流上下文
line: 原始 SSE 行
events: 当前累积的事件列表(预留参数,用于未来上下文感知转换如合并相邻事件)
Returns:
(sse_lines, converted_events) 元组:
- sse_lines: 转换后的 SSE 行列表(一入多出),空列表表示跳过该行
- converted_events: 转换后的事件对象列表(用于记录到 parsed_chunks
"""
# 空行直接返回
if not line or line.strip() == "":
return ([line] if line else [], [])
client_format = (ctx.client_api_format or "").strip().lower()
# [DONE] 标记处理:只有 OpenAI Chat 兼容流需要Responses(openai:cli) 不需要
if line == "data: [DONE]":
if client_format == "openai:chat":
return [line], []
else:
# Claude/Gemini/Responses 客户端不需要 [DONE] 标记
return [], []
provider_format = (ctx.provider_api_format or "").strip().lower()
# 过滤上游控制行id/retry避免与目标格式混淆
if line.startswith(("id:", "retry:")):
return [], []
# 解析 SSE 行为 JSON 对象
data_obj, status = self._parse_sse_line_to_json(line, provider_format)
# 根据解析状态决定行为
if status == "empty" or status == "skip":
return [], []
if status == "invalid" or status == "passthrough":
return [line], []
behavior = get_provider_behavior(
provider_type=ctx.provider_type,
endpoint_sig=ctx.provider_api_format,
)
envelope = behavior.envelope
if envelope:
data_obj = envelope.unwrap_response(data_obj)
envelope.postprocess_unwrapped_response(model=ctx.model, data=data_obj)
# 初始化流式转换状态
if ctx.stream_conversion_state is None:
from src.core.api_format.conversion.stream_state import StreamState
# 使用客户端请求的模型ctx.model而非映射后的上游模型ctx.mapped_model
init_model = ctx.model or ""
logger.debug(
f"[{ctx.request_id}] StreamState init: ctx.model={ctx.model!r}, "
f"mapped_model={ctx.mapped_model!r}, using={init_model!r}"
)
ctx.stream_conversion_state = StreamState(
model=init_model,
message_id=ctx.response_id or ctx.request_id or "",
)
# 执行格式转换
try:
registry = get_format_converter_registry()
# status == "ok" 时 data_obj 必定是有效的 dict防御性检查
if data_obj is None:
return [], []
converted_events = registry.convert_stream_chunk(
data_obj,
provider_format,
client_format,
state=ctx.stream_conversion_state,
)
result = _format_converted_events_to_sse(converted_events, client_format)
if result:
ctx.stream_conversion_event_count += len(converted_events)
return result, converted_events
except Exception as e:
logger.warning("格式转换失败,透传原始数据: {}", e)
return [line], []
def _parse_sse_line_to_json(self, line: str, provider_format: str) -> tuple[Any | None, str]:
"""
解析 SSE 行为 JSON 对象
支持多种格式:
- 标准 SSE: "data: {...}"
- event+data 同行: "event: xxx data: {...}"
- Gemini JSON-array: 裸 JSON 行
Args:
line: 原始 SSE 行
provider_format: Provider API 格式
Returns:
(parsed_json, status) 元组:
- (obj, "ok") - 解析成功
- (None, "empty") - 内容为空,应跳过
- (None, "invalid") - JSON 解析失败,应透传原始行
- (None, "skip") - 应跳过(如纯 event 行)
- (None, "passthrough") - 无法识别,应透传原始行
"""
# 标准 SSE: data: {...}
if line.startswith("data:"):
return _parse_sse_data_line(line)
# event + data 同行: event: xxx data: {...}
if line.startswith("event:") and " data:" in line:
return _parse_sse_event_data_line(line)
# 纯 event 行不参与转换
if line.startswith("event:"):
return None, "skip"
# Gemini JSON-array 格式
if provider_format.startswith("gemini"):
return _parse_gemini_json_array_line(line)
# 其他格式:无法识别,透传
return None, "passthrough"

View File

@@ -0,0 +1,119 @@
"""
CLI Message Handler 通用基类
将 CLI 格式处理器的通用逻辑HTTP 请求、SSE 解析、统计记录)抽取到基类,
子类只需实现格式特定的事件解析逻辑。
设计目标:
1. 减少代码重复 - 原来每个 CLI Handler 900+ 行,抽取后子类只需 ~100 行
2. 统一错误处理 - 超时、空流、故障转移等逻辑集中管理
3. 简化新格式接入 - 只需实现 ResponseParser 和少量钩子方法
"""
from __future__ import annotations
from collections.abc import Callable
from typing import Any
from sqlalchemy.orm import Session
from src.api.handlers.base.base_handler import BaseMessageHandler
from src.api.handlers.base.cli_event_mixin import CliEventMixin
from src.api.handlers.base.cli_monitor_mixin import CliMonitorMixin
from src.api.handlers.base.cli_prefetch_mixin import CliPrefetchMixin
from src.api.handlers.base.cli_request_mixin import CliRequestMixin
from src.api.handlers.base.cli_stream_mixin import CliStreamMixin
from src.api.handlers.base.cli_sync_mixin import CliSyncMixin
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.request_builder import PassthroughRequestBuilder
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
__all__ = [
"CliMessageHandlerBase",
"StreamContext",
]
class CliMessageHandlerBase(
CliRequestMixin,
CliStreamMixin,
CliPrefetchMixin,
CliEventMixin,
CliMonitorMixin,
CliSyncMixin,
BaseMessageHandler,
):
"""
CLI 格式消息处理器基类
提供 CLI 格式(直接透传请求)的通用处理逻辑:
- 流式请求的 HTTP 连接管理
- SSE 事件解析框架
- 统计信息收集和记录
- 错误处理和故障转移
子类需要实现:
- get_response_parser(): 返回格式特定的响应解析器
- 可选覆盖 handle_sse_event() 自定义事件处理
"""
# 子类可覆盖的配置
FORMAT_ID: str = "UNKNOWN" # API 格式标识
DATA_TIMEOUT: int = 30 # 流数据超时时间(秒)
EMPTY_CHUNK_THRESHOLD: int = 10 # 空流检测的 chunk 阈值
def __init__(
self,
db: Session,
user: User,
api_key: ApiKey,
request_id: str,
client_ip: str,
user_agent: str,
start_time: float,
allowed_api_formats: list | None = None,
adapter_detector: None | (
Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
) = None,
perf_metrics: dict[str, Any] | None = None,
api_family: str | None = None,
endpoint_kind: str | None = None,
):
allowed = allowed_api_formats or [self.FORMAT_ID]
super().__init__(
db=db,
user=user,
api_key=api_key,
request_id=request_id,
client_ip=client_ip,
user_agent=user_agent,
start_time=start_time,
allowed_api_formats=allowed,
adapter_detector=adapter_detector,
perf_metrics=perf_metrics,
api_family=api_family,
endpoint_kind=endpoint_kind,
)
self._parser: ResponseParser | None = None
self._request_builder = PassthroughRequestBuilder()
@property
def parser(self) -> ResponseParser:
"""获取响应解析器(懒加载)"""
if self._parser is None:
self._parser = self.get_response_parser()
return self._parser
def get_response_parser(self) -> ResponseParser:
"""
获取格式特定的响应解析器
子类可覆盖此方法提供自定义解析器,
默认从解析器注册表获取
"""
return get_parser_for_format(self.FORMAT_ID)
# _update_usage_to_streaming 方法已移至基类 BaseMessageHandler

View File

@@ -0,0 +1,723 @@
"""CLI Handler - 监控/统计 Mixin"""
from __future__ import annotations
import asyncio
import os
import time
from collections.abc import AsyncGenerator
from typing import TYPE_CHECKING, Any
import httpx
from fastapi import Request
from src.api.handlers.base.base_handler import MessageTelemetry
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.utils import filter_proxy_response_headers
from src.core.error_utils import extract_client_error_message
from src.core.exceptions import (
ProviderAuthException,
ProviderRateLimitException,
ProviderTimeoutException,
ThinkingSignatureException,
)
from src.core.logger import logger
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
def _read_stream_idle_timeout_seconds() -> float:
raw_value = os.getenv("STREAM_IDLE_TIMEOUT_SECONDS", "30")
try:
parsed = float(raw_value)
except (TypeError, ValueError):
return 30.0
return parsed if parsed > 0 else 30.0
class CliMonitorMixin:
"""监控和统计相关方法的 Mixin"""
# CancelledError 归因时,断连检查参数(秒)
CANCEL_DISCONNECT_CHECK_TIMEOUT_SECONDS = 0.5
CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS = (0.1, 0.2)
# 流式传输过程中若长时间没有任何 chunk提前判定为 idle timeout避免一直等到 worker 超时
STREAM_IDLE_TIMEOUT_SECONDS = _read_stream_idle_timeout_seconds()
async def _probe_client_disconnect(
self,
http_request: Request,
*,
request_id: str,
) -> tuple[bool, bool]:
"""单次探测客户端是否断连。
Returns:
(is_disconnected, is_indeterminate)
"""
try:
disconnected = await asyncio.wait_for(
asyncio.shield(http_request.is_disconnected()),
timeout=self.CANCEL_DISCONNECT_CHECK_TIMEOUT_SECONDS,
)
return bool(disconnected), False
except (asyncio.CancelledError, asyncio.TimeoutError):
return False, True
except Exception as e:
logger.debug("ID:{} | cancel 断连检测失败: {}", request_id, e)
return False, True
async def _confirm_client_disconnect(
self,
http_request: Request,
*,
request_id: str,
) -> tuple[bool, bool]:
"""CancelledError 场景下做多次断连确认,降低误判。
Returns:
(is_client_disconnected, check_indeterminate)
"""
disconnected, uncertain = await self._probe_client_disconnect(
http_request,
request_id=request_id,
)
if disconnected:
return True, uncertain
for delay in self.CANCEL_DISCONNECT_RETRY_DELAYS_SECONDS:
try:
await asyncio.sleep(delay)
except asyncio.CancelledError:
# 重试期间协程再次被取消,无法继续探测,标记为不确定
return False, True
disconnected, step_uncertain = await self._probe_client_disconnect(
http_request,
request_id=request_id,
)
uncertain = uncertain or step_uncertain
if disconnected:
return True, uncertain
return False, uncertain
async def _create_monitored_stream(
self,
ctx: StreamContext,
stream_generator: AsyncGenerator[bytes],
http_request: Request | None = None,
) -> AsyncGenerator[bytes]:
"""
创建带监控的流生成器
支持两种断连检测方式:
1. 如果提供了 http_request使用后台任务主动检测客户端断连
2. 如果未提供,仅依赖 asyncio.CancelledError 被动检测
Args:
ctx: 流上下文
stream_generator: 底层流生成器
http_request: FastAPI Request 对象,用于检测客户端断连
"""
import time as time_module
last_chunk_time = time_module.time()
chunk_count = 0
stream_started = False
idle_timeout_triggered = False
idle_timeout = self.STREAM_IDLE_TIMEOUT_SECONDS
parent_task = asyncio.current_task()
idle_watch_task: asyncio.Task[None] | None = None
async def watch_stream_idle_timeout() -> None:
nonlocal idle_timeout_triggered
if parent_task is None:
return
poll_interval = min(1.0, max(0.1, idle_timeout / 5))
while not ctx.has_completion:
await asyncio.sleep(poll_interval)
if not stream_started:
continue
if (time_module.time() - last_chunk_time) <= idle_timeout:
continue
idle_timeout_triggered = True
parent_task.cancel()
return
try:
idle_watch_task = asyncio.create_task(watch_stream_idle_timeout())
if http_request is not None:
# 使用后台任务检测断连,完全不阻塞流式传输
disconnected = False
async def check_disconnect_background() -> None:
nonlocal disconnected
while not disconnected and not ctx.has_completion:
await asyncio.sleep(0.5)
try:
if await http_request.is_disconnected():
disconnected = True
break
except Exception as e:
# 检测失败时不中断流,继续传输
logger.debug("ID:{} | 断连检测异常: {}", ctx.request_id, e)
# 启动后台检查任务
check_task = asyncio.create_task(check_disconnect_background())
try:
async for chunk in stream_generator:
if disconnected:
# 如果响应已完成,客户端断开不算失败
if ctx.has_completion:
logger.info(
f"ID:{ctx.request_id} | Client disconnected after completion"
)
else:
logger.warning("ID:{} | Client disconnected", ctx.request_id)
ctx.status_code = 499
ctx.error_message = "client_disconnected"
break
stream_started = True
last_chunk_time = time_module.time()
chunk_count += 1
yield chunk
finally:
check_task.cancel()
try:
await check_task
except asyncio.CancelledError:
pass
if idle_watch_task is not None:
idle_watch_task.cancel()
try:
await idle_watch_task
except asyncio.CancelledError:
pass
else:
# 无 http_request仅被动监控
try:
async for chunk in stream_generator:
stream_started = True
last_chunk_time = time_module.time()
chunk_count += 1
yield chunk
finally:
if idle_watch_task is not None:
idle_watch_task.cancel()
try:
await idle_watch_task
except asyncio.CancelledError:
pass
except asyncio.CancelledError:
# 防御性清理:正常路径中 idle_watch_task 已由内部 finally 取消,
# 但若 CancelledError 在异常路径传播,确保不留孤儿 task。
if idle_watch_task is not None and not idle_watch_task.done():
idle_watch_task.cancel()
# 注意CancelledError 不等于"用户手动取消",它既可能是客户端断连触发,
# 也可能是服务端(重载/关停/内部取消)导致的协程取消。
# 这里尽量做一次"断连归因":仅当能确认客户端已断开时才记为 499 cancelled。
time_since_last_chunk = time_module.time() - last_chunk_time
if not ctx.has_completion:
ctx.ensure_estimated_output_tokens()
if not ctx.has_completion and idle_timeout_triggered:
ctx.status_code = 504
ctx.error_message = "stream_idle_timeout"
cancel_origin = "stream_idle_timeout"
logger.warning(
f"ID:{ctx.request_id} | Stream idle timeout: "
f"idle_timeout={idle_timeout:g}s, "
f"chunks={chunk_count}, "
f"has_completion={ctx.has_completion}, "
f"time_since_last_chunk={time_since_last_chunk:.2f}s, "
f"output_tokens={ctx.output_tokens}"
)
ctx.upstream_response = (
f"cancel_origin={cancel_origin}, "
f"chunks={chunk_count}, "
f"has_completion={ctx.has_completion}, "
f"time_since_last_chunk={time_since_last_chunk:.2f}s, "
f"output_tokens={ctx.output_tokens}, "
f"idle_timeout={idle_timeout:g}s"
)
raise
is_client_disconnected = False
disconnect_check_uncertain = False
if http_request is not None:
is_client_disconnected, disconnect_check_uncertain = (
await self._confirm_client_disconnect(
http_request,
request_id=ctx.request_id,
)
)
# 如果响应已完成,不标记为失败/取消
if not ctx.has_completion:
if is_client_disconnected:
ctx.status_code = 499
ctx.error_message = "client_disconnected"
cancel_origin = "client_disconnected"
logger.warning(
f"ID:{ctx.request_id} | Stream cancelled by client: "
f"chunks={chunk_count}, "
f"has_completion={ctx.has_completion}, "
f"time_since_last_chunk={time_since_last_chunk:.2f}s, "
f"output_tokens={ctx.output_tokens}"
)
elif disconnect_check_uncertain:
# 断连检查本身不稳定(超时/取消/异常)时,避免直接定性为 server_cancelled。
ctx.status_code = 503
ctx.error_message = "cancelled_unknown"
cancel_origin = "cancelled_unknown"
logger.warning(
f"ID:{ctx.request_id} | Stream cancelled with unknown origin: "
f"chunks={chunk_count}, "
f"has_completion={ctx.has_completion}, "
f"time_since_last_chunk={time_since_last_chunk:.2f}s, "
f"output_tokens={ctx.output_tokens}"
)
else:
# 服务端中断(例如重载/关停/内部取消) -- 不应伪装成客户端取消
ctx.status_code = 503
ctx.error_message = "server_cancelled"
cancel_origin = "server_cancelled"
logger.error(
f"ID:{ctx.request_id} | Stream interrupted by server: "
f"chunks={chunk_count}, "
f"has_completion={ctx.has_completion}, "
f"time_since_last_chunk={time_since_last_chunk:.2f}s, "
f"output_tokens={ctx.output_tokens}"
)
ctx.upstream_response = (
f"cancel_origin={cancel_origin}, "
f"chunks={chunk_count}, "
f"has_completion={ctx.has_completion}, "
f"time_since_last_chunk={time_since_last_chunk:.2f}s, "
f"output_tokens={ctx.output_tokens}"
)
raise
except httpx.TimeoutException as e:
if idle_watch_task is not None and not idle_watch_task.done():
idle_watch_task.cancel()
ctx.status_code = 504
ctx.error_message = str(e)
raise
except Exception as e:
if idle_watch_task is not None and not idle_watch_task.done():
idle_watch_task.cancel()
ctx.status_code = 500
ctx.error_message = str(e)
raise
async def _record_stream_stats(
self: CliHandlerProtocol,
ctx: StreamContext,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None:
"""在流完成后记录统计信息"""
try:
# 使用 self.start_time 作为时间基准,与首字时间保持一致
# 注意:不要把统计延迟算进响应时间里
response_time_ms = int((time.time() - self.start_time) * 1000)
await asyncio.sleep(0.1)
if not ctx.provider_name:
logger.warning("[{}] 流式请求失败,未选中提供商", ctx.request_id)
return
behavior = get_provider_behavior(
provider_type=ctx.provider_type,
endpoint_sig=ctx.provider_api_format,
)
envelope = behavior.envelope
if envelope:
envelope.on_http_status(
base_url=ctx.selected_base_url,
status_code=ctx.status_code,
)
# 获取新的 DB session
db_gen = get_db()
bg_db = next(db_gen)
try:
from src.models.database import ApiKey as ApiKeyModel
user = bg_db.query(User).filter(User.id == ctx.user_id).first()
api_key = bg_db.query(ApiKeyModel).filter(ApiKeyModel.id == ctx.api_key_id).first()
if not user or not api_key:
logger.warning(
"[{}] 无法记录统计: user={} api_key={}",
ctx.request_id,
user is not None,
api_key is not None,
)
return
bg_telemetry = MessageTelemetry(
bg_db, user, api_key, ctx.request_id, self.client_ip
)
if ctx.should_estimate_incomplete_tokens():
self._estimate_tokens_for_incomplete_stream(
ctx, ctx.provider_request_body or original_request_body
)
with ctx.managed_recorded_bodies(response_time_ms) as recorded_bodies:
# 根据状态码决定记录成功还是失败
# 499 = 客户端取消(不算系统失败);其他 4xx/5xx 视为失败
if ctx.status_code and ctx.status_code >= 400:
client_response_headers = ctx.client_response_headers or {
"content-type": "application/json"
}
if ctx.is_client_disconnected():
# 客户端取消:记录为 cancelled不算系统失败
request_metadata = self._merge_scheduling_metadata(
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
selected_key_id=ctx.key_id,
candidate_keys=ctx.candidate_keys,
pool_summary=ctx.pool_summary,
fallback_from_request=True,
)
await bg_telemetry.record_cancelled(
provider=ctx.provider_name or "unknown",
model=ctx.model,
response_time_ms=response_time_ms,
first_byte_time_ms=ctx.first_byte_time_ms,
status_code=ctx.status_code,
request_headers=original_headers,
request_body=original_request_body,
is_stream=True,
api_format=ctx.api_format,
api_family=self.api_family,
endpoint_kind=self.endpoint_kind,
provider_request_headers=ctx.provider_request_headers,
provider_request_body=ctx.provider_request_body,
input_tokens=ctx.input_tokens,
output_tokens=ctx.output_tokens,
cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens,
response_body=recorded_bodies.response_body,
client_response_body=recorded_bodies.client_response_body,
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=request_metadata,
)
logger.debug("{} 流式响应被客户端取消", self.FORMAT_ID)
logger.info(
"[CANCEL] {} | {} | {} | {}ms | {} | in:{} out:{} cache:{}",
self.request_id[:8],
ctx.model,
ctx.provider_name,
response_time_ms,
ctx.status_code,
ctx.input_tokens,
ctx.output_tokens,
ctx.cached_tokens,
)
else:
# 服务端/上游异常:记录为失败
request_metadata = self._merge_scheduling_metadata(
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
selected_key_id=ctx.key_id,
candidate_keys=ctx.candidate_keys,
pool_summary=ctx.pool_summary,
fallback_from_request=True,
)
await bg_telemetry.record_failure(
provider=ctx.provider_name or "unknown",
model=ctx.model,
response_time_ms=response_time_ms,
status_code=ctx.status_code,
error_message=ctx.error_message or f"HTTP {ctx.status_code}",
request_headers=original_headers,
request_body=original_request_body,
is_stream=True,
api_format=ctx.api_format,
api_family=self.api_family,
endpoint_kind=self.endpoint_kind,
provider_request_headers=ctx.provider_request_headers,
provider_request_body=ctx.provider_request_body,
# 预估 token 信息(来自 message_start 事件)
input_tokens=ctx.input_tokens,
output_tokens=ctx.output_tokens,
cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens,
response_body=recorded_bodies.response_body,
client_response_body=recorded_bodies.client_response_body,
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=request_metadata,
)
logger.debug("{} 流式响应中断", self.FORMAT_ID)
logger.info(
"[FAIL] {} | {} | {} | {}ms | {} | in:{} out:{} cache:{}",
self.request_id[:8],
ctx.model,
ctx.provider_name,
response_time_ms,
ctx.status_code,
ctx.input_tokens,
ctx.output_tokens,
ctx.cached_tokens,
)
else:
# 在记录统计前,允许子类从 parsed_chunks 中提取额外的元数据
self._finalize_stream_metadata(ctx)
# 流式格式转换汇总日志
if ctx.stream_conversion_event_count > 0:
logger.debug(
"[{}] 流式转换完成: {}->{}, total_events={}",
self.request_id[:8],
ctx.provider_api_format,
ctx.client_api_format,
ctx.stream_conversion_event_count,
)
# 流未正常完成(如上游截断/连接中断)且无 token 数据时,
# 从已收集的文本和请求体估算 tokens避免 usage 记录为 0
# 流式成功时,返回给客户端的是提供商响应头 + SSE 必需头
client_response_headers = filter_proxy_response_headers(
ctx.response_headers
)
client_response_headers.update(
{
"Cache-Control": "no-cache, no-transform",
"X-Accel-Buffering": "no",
"content-type": "text/event-stream",
}
)
logger.debug(
"[{}] 开始记录 Usage: provider={}, model={}, in={}, out={}",
ctx.request_id,
ctx.provider_name,
ctx.model,
ctx.input_tokens,
ctx.output_tokens,
)
request_metadata = self._merge_scheduling_metadata(
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
selected_key_id=ctx.key_id,
candidate_keys=ctx.candidate_keys,
pool_summary=ctx.pool_summary,
fallback_from_request=True,
)
total_cost = await bg_telemetry.record_success(
provider=ctx.provider_name,
model=ctx.model,
input_tokens=ctx.input_tokens,
output_tokens=ctx.output_tokens,
response_time_ms=response_time_ms,
first_byte_time_ms=ctx.first_byte_time_ms, # 传递首字时间
status_code=ctx.status_code,
request_headers=original_headers,
request_body=original_request_body,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
response_body=recorded_bodies.response_body,
client_response_body=recorded_bodies.client_response_body,
provider_request_body=ctx.provider_request_body,
cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens,
is_stream=True,
provider_request_headers=ctx.provider_request_headers,
api_format=ctx.api_format,
api_family=self.api_family,
endpoint_kind=self.endpoint_kind,
# 格式转换追踪
endpoint_api_format=ctx.provider_api_format or None,
has_format_conversion=ctx.has_format_conversion,
# Provider 侧追踪信息(用于记录真实成本)
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
# 模型映射信息
target_model=ctx.mapped_model,
# Provider 响应元数据(如 Gemini 的 modelVersion
response_metadata=(
ctx.response_metadata if ctx.response_metadata else None
),
request_metadata=request_metadata,
)
logger.debug(
"[{}] Usage 记录完成: cost=${:.6f}", ctx.request_id, total_cost
)
# 简洁的请求完成摘要(两行格式)
ttfb_part = (
f" | TTFB: {ctx.first_byte_time_ms}ms" if ctx.first_byte_time_ms else ""
)
logger.info(
"[OK] {} | {} | {}{}\n Total: {}ms | in:{} out:{}",
self.request_id[:8],
ctx.model,
ctx.provider_name,
ttfb_part,
response_time_ms,
ctx.input_tokens or 0,
ctx.output_tokens or 0,
)
# 更新候选记录的最终状态和延迟时间
# 注意RequestExecutor 会在流开始时过早地标记成功(只记录了连接建立的时间)
# 这里用流传输完成后的实际时间覆盖
if ctx.attempt_id:
from src.services.request.candidate import RequestCandidateService
# 计算候选自身的 TTFB
candidate_first_byte_time_ms: int | None = None
if ctx.first_byte_time_ms is not None:
candidate_first_byte_time_ms = (
RequestCandidateService.calculate_candidate_ttfb(
db=bg_db,
candidate_id=ctx.attempt_id,
request_start_time=self.start_time,
global_first_byte_time_ms=ctx.first_byte_time_ms,
)
)
# 根据状态码决定是成功还是失败
# 499 = 客户端断开连接,应标记为失败
# 503 = 服务不可用(如流中断),应标记为失败
if ctx.status_code and ctx.status_code >= 400:
# 请求链路追踪使用 upstream_response原始响应回退到 error_message友好消息
trace_error_message = (
ctx.upstream_response or ctx.error_message or f"HTTP {ctx.status_code}"
)
extra_data = {
"stream_completed": False,
"chunk_count": ctx.chunk_count,
"data_count": ctx.data_count,
}
if ctx.proxy_info:
extra_data["proxy"] = ctx.proxy_info
if candidate_first_byte_time_ms is not None:
extra_data["first_byte_time_ms"] = candidate_first_byte_time_ms
if ctx.is_client_disconnected():
RequestCandidateService.mark_candidate_cancelled(
db=bg_db,
candidate_id=ctx.attempt_id,
status_code=ctx.status_code,
latency_ms=response_time_ms,
extra_data=extra_data,
)
else:
RequestCandidateService.mark_candidate_failed(
db=bg_db,
candidate_id=ctx.attempt_id,
error_type="stream_error",
error_message=trace_error_message,
status_code=ctx.status_code,
latency_ms=response_time_ms,
extra_data=extra_data,
)
else:
extra_data = {
"stream_completed": True,
"chunk_count": ctx.chunk_count,
"data_count": ctx.data_count,
}
if ctx.proxy_info:
extra_data["proxy"] = ctx.proxy_info
if ctx.rectified:
extra_data["rectified"] = True
if candidate_first_byte_time_ms is not None:
extra_data["first_byte_time_ms"] = candidate_first_byte_time_ms
RequestCandidateService.mark_candidate_success(
db=bg_db,
candidate_id=ctx.attempt_id,
status_code=ctx.status_code,
latency_ms=response_time_ms,
extra_data=extra_data,
)
finally:
bg_db.close()
except Exception as e:
logger.exception("记录流式统计信息时出错")
finally:
# 遥测写入完成后主动释放大对象列表,降低高并发长流的内存滞留。
ctx.release_recorded_chunks()
async def _record_stream_failure(
self,
ctx: StreamContext,
error: Exception,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None:
"""记录流式请求失败"""
# 使用 self.start_time 作为时间基准,与首字时间保持一致
response_time_ms = int((time.time() - self.start_time) * 1000)
status_code = 503
if isinstance(error, ThinkingSignatureException):
status_code = 400
elif isinstance(error, ProviderAuthException):
status_code = 503
elif isinstance(error, ProviderRateLimitException):
status_code = 429
elif isinstance(error, ProviderTimeoutException):
status_code = 504
ctx.status_code = status_code
ctx.error_message = str(error)
# 失败时返回给客户端的是 JSON 错误响应
client_response_headers = {"content-type": "application/json"}
request_metadata = self._merge_scheduling_metadata(
{"perf": ctx.perf_metrics} if ctx.perf_metrics else None,
selected_key_id=ctx.key_id,
candidate_keys=ctx.candidate_keys,
pool_summary=ctx.pool_summary,
fallback_from_request=True,
)
try:
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=original_request_body,
is_stream=True,
api_format=ctx.api_format,
api_family=self.api_family,
endpoint_kind=self.endpoint_kind,
provider_request_headers=ctx.provider_request_headers,
provider_request_body=ctx.provider_request_body,
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=request_metadata,
)
finally:
# 失败路径同样可能持有 chunk 审计数据,及时释放。
ctx.release_recorded_chunks()

View File

@@ -0,0 +1,634 @@
"""CLI Handler - Prefetch 和错误检测 Mixin"""
from __future__ import annotations
import codecs
import json
import time
from collections.abc import AsyncGenerator
from typing import TYPE_CHECKING, Any
import httpx
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.utils import (
check_html_response,
check_prefetched_response_error,
ensure_stream_buffer_limit,
)
from src.config.constants import StreamDefaults
from src.config.settings import config
from src.core.exceptions import (
EmbeddedErrorException,
ProviderNotAvailableException,
ProviderTimeoutException,
)
from src.core.logger import logger
from src.services.provider.behavior import get_provider_behavior
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
class CliPrefetchMixin:
"""Prefetch 和错误检测相关方法的 Mixin"""
def _flush_remaining_sse_data(
self,
ctx: StreamContext,
buffer: bytes,
decoder: codecs.IncrementalDecoder,
sse_parser: SSEEventParser,
*,
record_chunk: bool = True,
) -> None:
"""
异常发生时 flush 残留的字节 buffer 和 SSE parser 内部缓冲区。
用于 StreamClosed / RemoteProtocolError 等场景:
连接断开可能恰好发生在最后一个 SSE 事件(如 response.completed
的 data 行已收到、但终止空行尚未到达之时。此方法确保这些事件仍能被处理,
从而正确捕获 usage 等关键信息。
"""
try:
# 1) flush 字节 buffer 中的残余行
if buffer:
remaining = decoder.decode(buffer, True)
for line in remaining.split("\n"):
stripped = line.rstrip("\r")
events = sse_parser.feed_line(stripped)
for event in events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=record_chunk,
)
# 2) flush SSE parser 内部累积的未完成事件
for event in sse_parser.flush():
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=record_chunk,
)
except Exception:
# best-effort: 不应因 flush 失败影响后续流程
pass
def _estimate_tokens_for_incomplete_stream(
self,
ctx: StreamContext,
request_body: dict[str, Any],
) -> None:
"""
流未正常完成(无 response.completed且 token 均为 0 时的兜底估算。
从已收集的输出文本和请求体粗略估算 token 数,确保 usage 记录不为 0。
估算采用 ~4 字符/token 的保守比例。
"""
# 输出 tokens从已收集的文本估算
if ctx.collected_text_length > 0:
ctx.output_tokens = max(1, ctx.collected_text_length // 4)
# 输入 tokens从请求体文本内容估算
try:
total_input_len = 0
instructions = request_body.get("instructions")
if isinstance(instructions, str):
total_input_len += len(instructions)
# OpenAI Responses API 使用 input 字段Claude 使用 messages
input_items = request_body.get("input") or request_body.get("messages") or []
if isinstance(input_items, list):
for item in input_items:
if isinstance(item, str):
total_input_len += len(item)
elif isinstance(item, dict):
content = item.get("content", "")
if isinstance(content, str):
total_input_len += len(content)
elif isinstance(content, list):
for block in content:
if isinstance(block, dict):
text = block.get("text", "")
if isinstance(text, str):
total_input_len += len(text)
if total_input_len > 0:
ctx.input_tokens = max(1, total_input_len // 4)
else:
# fallback: 整个请求体 JSON 大小
body_str = json.dumps(request_body, ensure_ascii=False)
ctx.input_tokens = max(1, len(body_str) // 4)
except Exception:
pass
if ctx.input_tokens > 0 or ctx.output_tokens > 0:
logger.warning(
"[{}] 流未正常完成 (has_completion=False, data_count={}), "
"使用估算 tokens: in={}, out={}",
ctx.request_id,
ctx.data_count,
ctx.input_tokens,
ctx.output_tokens,
)
async def _prefetch_and_check_embedded_error(
self: CliHandlerProtocol,
byte_iterator: Any,
provider: "Provider",
endpoint: "ProviderEndpoint",
ctx: StreamContext,
) -> list:
"""
预读流的前几行,检测嵌套错误
某些 Provider如 Gemini可能返回 HTTP 200但在响应体中包含错误信息。
这种情况需要在流开始输出之前检测,以便触发重试逻辑。
同时检测 HTML 响应(通常是 base_url 配置错误导致返回网页)。
首次读取时会应用 TTFB首字节超时检测超时则触发故障转移。
Args:
byte_iterator: 字节流迭代器
provider: Provider 对象
endpoint: Endpoint 对象
ctx: 流上下文
Returns:
预读的字节块列表(需要在后续流中先输出)
Raises:
EmbeddedErrorException: 如果检测到嵌套错误
ProviderNotAvailableException: 如果检测到 HTML 响应(配置错误)
ProviderTimeoutException: 如果首字节超时TTFB timeout
"""
prefetched_chunks: list = []
max_prefetch_lines = config.stream_prefetch_lines # 最多预读行数来检测错误
max_prefetch_bytes = StreamDefaults.MAX_PREFETCH_BYTES # 避免无换行响应导致 buffer 增长
total_prefetched_bytes = 0
buffer = b""
line_count = 0
should_stop = False
# 使用增量解码器处理跨 chunk 的 UTF-8 字符
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
try:
# 获取对应格式的解析器
provider_format = ctx.provider_api_format
if provider_format:
try:
provider_parser = get_parser_for_format(provider_format)
except KeyError:
provider_parser = self.parser
else:
provider_parser = self.parser
# 使用共享的 TTFB 超时函数读取首字节
# 优先使用 Provider 配置,否则使用全局配置
ttfb_timeout = provider.stream_first_byte_timeout or config.stream_first_byte_timeout
first_chunk, aiter = await read_first_chunk_with_ttfb_timeout(
byte_iterator,
timeout=ttfb_timeout,
request_id=self.request_id,
provider_name=str(provider.name),
)
prefetched_chunks.append(first_chunk)
total_prefetched_bytes += len(first_chunk)
buffer += first_chunk
ensure_stream_buffer_limit(
buffer,
request_id=self.request_id,
provider_name=str(provider.name),
)
# 继续读取剩余的预读数据
async for chunk in aiter:
prefetched_chunks.append(chunk)
total_prefetched_bytes += len(chunk)
buffer += chunk
ensure_stream_buffer_limit(
buffer,
request_id=self.request_id,
provider_name=str(provider.name),
)
# 尝试按行解析缓冲区SSE 格式)
while b"\n" in buffer:
line_bytes, buffer = buffer.split(b"\n", 1)
try:
# 使用增量解码器,可以正确处理跨 chunk 的多字节字符
line = decoder.decode(line_bytes + b"\n", False).rstrip("\n")
except Exception as e:
logger.warning(
f"[{self.request_id}] 预读时 UTF-8 解码失败: {e}, "
f"bytes={line_bytes[:50]!r}"
)
continue
line_count += 1
normalized_line = line.rstrip("\r")
# 检测 HTML 响应base_url 配置错误的常见症状)
if check_html_response(normalized_line):
logger.error(
f" [{self.request_id}] 检测到 HTML 响应,可能是 base_url 配置错误: "
f"Provider={provider.name}, Endpoint={endpoint.id[:8]}..., "
f"base_url={endpoint.base_url}"
)
raise ProviderNotAvailableException(
"上游服务返回了非预期的响应格式",
provider_name=str(provider.name),
upstream_status=200,
upstream_response=(
normalized_line[:500] if normalized_line else "(empty)"
),
)
if not normalized_line or normalized_line.startswith(":"):
# 空行或注释行,继续预读
if line_count >= max_prefetch_lines:
break
continue
# 尝试解析 SSE 数据
data_str = normalized_line
if normalized_line.startswith("data: "):
data_str = normalized_line[6:]
if data_str == "[DONE]":
should_stop = True
break
try:
data = json.loads(data_str)
except json.JSONDecodeError:
# 不是有效 JSON可能是部分数据继续
if line_count >= max_prefetch_lines:
break
continue
# 使用解析器检查是否为错误响应
if isinstance(data, dict) and provider_parser.is_error_response(data):
# 提取错误信息
parsed = provider_parser.parse_response(data, 200)
logger.warning(
f" [{self.request_id}] 检测到嵌套错误: "
f"Provider={provider.name}, "
f"error_type={parsed.error_type}, "
f"message={parsed.error_message}"
)
raise EmbeddedErrorException(
provider_name=str(provider.name),
error_code=(
int(parsed.error_type)
if parsed.error_type and parsed.error_type.isdigit()
else None
),
error_message=parsed.error_message,
error_status=parsed.error_type,
)
# 预读到有效数据,没有错误,停止预读
should_stop = True
break
# 达到预读字节上限,停止继续预读(避免无换行响应导致内存增长)
if not should_stop and total_prefetched_bytes >= max_prefetch_bytes:
logger.debug(
f" [{self.request_id}] 预读达到字节上限,停止继续预读: "
f"Provider={provider.name}, bytes={total_prefetched_bytes}, "
f"max_bytes={max_prefetch_bytes}"
)
break
if should_stop or line_count >= max_prefetch_lines:
break
# 预读结束后,检查是否为非 SSE 格式的 HTML/JSON 响应
# 处理某些代理返回的纯 JSON 错误(可能无换行/多行 JSON以及 HTML 页面base_url 配置错误)
if not should_stop and prefetched_chunks:
check_prefetched_response_error(
prefetched_chunks=prefetched_chunks,
parser=provider_parser,
request_id=self.request_id,
provider_name=str(provider.name),
endpoint_id=endpoint.id,
base_url=endpoint.base_url,
)
except (EmbeddedErrorException, ProviderTimeoutException, ProviderNotAvailableException):
# 重新抛出可重试的 Provider 异常,触发故障转移
raise
except OSError as e:
# 网络 I/O 异常:记录警告,可能需要重试
logger.warning(
" [{}] 预读流时发生网络异常: {}: {}", self.request_id, type(e).__name__, e
)
except Exception as e:
# 未预期的严重异常:记录错误并重新抛出,避免掩盖问题
logger.error(
f" [{self.request_id}] 预读流时发生严重异常: {type(e).__name__}: {e}",
exc_info=True,
)
raise
return prefetched_chunks
async def _create_response_stream_with_prefetch(
self,
ctx: StreamContext,
byte_iterator: Any,
response_ctx: Any,
prefetched_chunks: list,
) -> AsyncGenerator[bytes]:
"""创建响应流生成器(带预读数据,使用字节流)"""
try:
sse_parser = SSEEventParser()
last_data_time = time.time()
buffer = b""
output_state = {"first_yield": True, "streaming_updated": False}
_sample_lines: list[str] = [] # 采集前几行原始内容,用于空流诊断
_MAX_SAMPLE_LINES = 5
# 使用增量解码器处理跨 chunk 的 UTF-8 字符
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
# 使用已设置的 ctx.needs_conversion由候选筛选阶段根据端点配置判断
# 不再调用 _needs_format_conversion它只检查格式差异不检查端点配置
needs_conversion = ctx.needs_conversion
behavior = get_provider_behavior(
provider_type=ctx.provider_type,
endpoint_sig=ctx.provider_api_format,
)
envelope = behavior.envelope
if envelope and envelope.force_stream_rewrite():
needs_conversion = True
ctx.needs_conversion = True
# Kiro 特殊处理AWS Event Stream 二进制流需要重写为 SSE
ctx_provider_type = str(ctx.provider_type or "").strip().lower()
if ctx_provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
from src.services.provider.adapters.kiro.eventstream_rewriter import (
apply_kiro_stream_rewrite,
)
byte_iterator = apply_kiro_stream_rewrite(
byte_iterator,
model=str(ctx.model or ""),
input_tokens=int(ctx.input_tokens or 0),
prefetched_chunks=list(prefetched_chunks) if prefetched_chunks else None,
)
prefetched_chunks = []
# Kiro 重写后输出的是 Claude SSE 格式
# 客户端也是 Claude CLI不需要再进行格式转换
needs_conversion = False
ctx.needs_conversion = False
# 先处理预读的字节块
for chunk in prefetched_chunks:
buffer += chunk
ensure_stream_buffer_limit(
buffer,
request_id=self.request_id,
provider_name=ctx.provider_name,
)
# 处理缓冲区中的完整行
while b"\n" in buffer:
line_bytes, buffer = buffer.split(b"\n", 1)
try:
# 使用增量解码器,可以正确处理跨 chunk 的多字节字符
line = decoder.decode(line_bytes + b"\n", False).rstrip("\n")
except Exception as e:
logger.warning(
f"[{self.request_id}] UTF-8 解码失败: {e}, "
f"bytes={line_bytes[:50]!r}"
)
continue
normalized_line = line.rstrip("\r")
events = sse_parser.feed_line(normalized_line)
if normalized_line == "":
for event in events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=not needs_conversion,
)
self._mark_first_output(ctx, output_state)
yield b"\n"
continue
ctx.chunk_count += 1
if len(_sample_lines) < _MAX_SAMPLE_LINES:
_sample_lines.append(normalized_line[:200])
# 格式转换或直接透传
if needs_conversion:
converted_lines, converted_events = self._convert_sse_line(
ctx, line, events
)
# 记录转换后的数据到 parsed_chunks
self._record_converted_chunks(ctx, converted_events)
for converted_line in converted_lines:
if converted_line:
self._mark_first_output(ctx, output_state)
yield (converted_line + "\n").encode("utf-8")
else:
self._mark_first_output(ctx, output_state)
yield (line + "\n").encode("utf-8")
for event in events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=not needs_conversion,
)
if ctx.data_count > 0:
last_data_time = time.time()
# 继续处理剩余的流数据(使用同一个迭代器)
async for chunk in byte_iterator:
buffer += chunk
ensure_stream_buffer_limit(
buffer,
request_id=self.request_id,
provider_name=ctx.provider_name,
)
# 处理缓冲区中的完整行
while b"\n" in buffer:
line_bytes, buffer = buffer.split(b"\n", 1)
try:
# 使用增量解码器,可以正确处理跨 chunk 的多字节字符
line = decoder.decode(line_bytes + b"\n", False).rstrip("\n")
except Exception as e:
logger.warning(
f"[{self.request_id}] UTF-8 解码失败: {e}, "
f"bytes={line_bytes[:50]!r}"
)
continue
normalized_line = line.rstrip("\r")
events = sse_parser.feed_line(normalized_line)
if normalized_line == "":
for event in events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=not needs_conversion,
)
self._mark_first_output(ctx, output_state)
yield b"\n"
continue
ctx.chunk_count += 1
if len(_sample_lines) < _MAX_SAMPLE_LINES:
_sample_lines.append(normalized_line[:200])
# 空流检测:超过阈值且无数据,发送错误事件并结束
if ctx.chunk_count > self.EMPTY_CHUNK_THRESHOLD and ctx.data_count == 0:
elapsed = time.time() - last_data_time
if elapsed > self.DATA_TIMEOUT:
logger.warning("Provider '{}' 流超时且无数据", ctx.provider_name)
# 设置错误状态用于后续记录
ctx.status_code = 504
ctx.error_message = "流式响应超时,未收到有效数据"
ctx.upstream_response = f"流超时: Provider={ctx.provider_name}, elapsed={elapsed:.1f}s, chunk_count={ctx.chunk_count}, data_count=0"
error_event = {
"type": "error",
"error": {
"type": "empty_stream_timeout",
"message": ctx.error_message,
},
}
self._mark_first_output(ctx, output_state)
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
return
# 格式转换或直接透传
if needs_conversion:
converted_lines, converted_events = self._convert_sse_line(
ctx, line, events
)
# 记录转换后的数据到 parsed_chunks
self._record_converted_chunks(ctx, converted_events)
for converted_line in converted_lines:
if converted_line:
self._mark_first_output(ctx, output_state)
yield (converted_line + "\n").encode("utf-8")
else:
self._mark_first_output(ctx, output_state)
yield (line + "\n").encode("utf-8")
for event in events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=not needs_conversion,
)
if ctx.data_count > 0:
last_data_time = time.time()
# flush 字节 buffer 残余数据 + SSE parser 内部缓冲区
for chunk in self._flush_buffer_with_conversion(
ctx, buffer, decoder, sse_parser, needs_conversion
):
yield chunk
# 检查是否收到数据
if ctx.data_count == 0:
# 空流通常意味着配置错误(如 base_url 指向了网页而非 API
sample_info = f", 前几行内容: {_sample_lines!r}" if _sample_lines else ""
logger.error(
f"Provider '{ctx.provider_name}' 返回空流式响应 (收到 {ctx.chunk_count} 个非数据行), "
f"可能是 endpoint base_url 配置错误{sample_info}"
)
# 设置错误状态用于后续记录
ctx.status_code = 503
ctx.error_message = "上游服务返回了空的流式响应"
ctx.upstream_response = f"空流式响应: Provider={ctx.provider_name}, chunk_count={ctx.chunk_count}, data_count=0, 可能是 base_url 配置错误"
error_event = {
"type": "error",
"error": {
"type": "empty_response",
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
logger.debug("流式数据转发完成")
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
client_fmt = (ctx.client_api_format or "").strip().lower()
if needs_conversion and client_fmt == "openai:chat":
yield b"data: [DONE]\n\n"
except GeneratorExit:
raise
except httpx.StreamClosed:
# 连接关闭前 flush 残余数据,尝试捕获尾部事件(如 response.completed 中的 usage
self._flush_remaining_sse_data(
ctx, buffer, decoder, sse_parser, record_chunk=not needs_conversion
)
if ctx.data_count == 0:
logger.warning("Provider '{}' 流连接关闭且无数据", ctx.provider_name)
# 设置错误状态用于后续记录
ctx.status_code = 503
ctx.error_message = "上游服务连接关闭且未返回数据"
ctx.upstream_response = f"流连接关闭: Provider={ctx.provider_name}, chunk_count={ctx.chunk_count}, data_count=0"
error_event = {
"type": "error",
"error": {
"type": "stream_closed",
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
except httpx.RemoteProtocolError:
# 连接异常关闭前 flush 残余数据,尝试捕获尾部事件(如 response.completed 中的 usage
self._flush_remaining_sse_data(
ctx, buffer, decoder, sse_parser, record_chunk=not needs_conversion
)
if ctx.data_count > 0:
error_event = {
"type": "error",
"error": {
"type": "connection_error",
"message": "上游连接意外关闭,部分响应已成功传输",
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
raise
except httpx.ReadError:
# 代理/上游连接读取失败(如 aether-proxy 中断),与 RemoteProtocolError 处理逻辑一致
self._flush_remaining_sse_data(
ctx, buffer, decoder, sse_parser, record_chunk=not needs_conversion
)
if ctx.data_count > 0:
error_event = {
"type": "error",
"error": {
"type": "connection_error",
"message": "代理或上游连接读取失败,部分响应已成功传输",
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
raise
finally:
try:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass

View File

@@ -0,0 +1,293 @@
"""
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.cli_request_mixin import CliUpstreamRequestResult
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]
primary_api_format: 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 = ...,
) -> bool: ...
def _build_request_metadata(
self,
http_request: Any | None = ...,
) -> dict[str, Any] | None: ...
def _merge_scheduling_metadata(
self,
request_metadata: dict[str, Any] | None,
*,
exec_result: Any | None = ...,
selected_key_id: str | None = ...,
candidate_keys: list[Any] | None = ...,
pool_summary: dict[str, Any] | None = ...,
fallback_from_request: bool = ...,
) -> 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: ...
async 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 = ...,
output_limit: int | None = ...,
) -> tuple[dict[str, Any], str]: ...
async def _build_upstream_request(
self,
*,
provider: Any,
endpoint: Any,
key: Any,
request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None,
client_api_format: str,
provider_api_format: str,
fallback_model: str,
mapped_model: str | None,
client_is_stream: bool,
needs_conversion: bool = ...,
output_limit: int | None = ...,
) -> CliUpstreamRequestResult: ...
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 _flush_buffer_with_conversion(
self,
ctx: StreamContext,
buffer: bytes,
decoder: Any,
sse_parser: Any,
needs_conversion: bool,
) -> Any: ... # Iterator[bytes]
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

@@ -0,0 +1,505 @@
"""CLI Handler - 请求准备 Mixin"""
from __future__ import annotations
from dataclasses import dataclass
from typing import (
TYPE_CHECKING,
Any,
)
from src.api.handlers.base.request_builder import get_provider_auth
from src.api.handlers.base.utils import get_format_converter_registry
from src.core.api_format.headers import set_accept_if_absent
from src.core.logger import logger
from src.services.provider.behavior import get_provider_behavior
from src.services.provider.prompt_cache import maybe_patch_request_with_prompt_cache_key
from src.services.provider.stream_policy import (
enforce_stream_mode_for_upstream,
get_upstream_stream_policy,
resolve_upstream_is_stream,
)
from src.services.provider.transport import build_provider_url
from src.services.provider.upstream_headers import build_upstream_extra_headers
if TYPE_CHECKING:
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
from src.core.api_format import EndpointDefinition
@dataclass(slots=True)
class CliUpstreamRequestResult:
"""Final outbound request artifacts for a selected Provider candidate."""
payload: dict[str, Any]
headers: dict[str, str]
url: str
url_model: str
envelope: Any
upstream_is_stream: bool
tls_profile: str | None = None
selected_base_url: str | None = None
class CliRequestMixin:
"""请求准备相关方法的 Mixin"""
async def _get_mapped_model(
self,
source_model: str,
provider_id: str,
) -> str | None:
"""
获取模型映射后的实际模型名
查找逻辑:
1. 直接通过 GlobalModel.name 匹配
2. 查找该 Provider 的 Model 实现
3. 使用 provider_model_name / provider_model_mappings 选择最终名称
Args:
source_model: 用户请求的模型名(必须是 GlobalModel.name
provider_id: Provider ID
Returns:
映射后的 Provider 模型名,如果没有找到映射则返回 None
"""
from src.services.model.mapper import ModelMapperMiddleware
mapper = ModelMapperMiddleware(self.db)
mapping = await mapper.get_mapping(source_model, provider_id)
logger.debug(
f"[CLI] _get_mapped_model: source={source_model}, provider={provider_id[:8]}..., mapping={mapping}"
)
if mapping and mapping.model:
# 使用 select_provider_model_name 支持模型映射功能
# 传入 api_key.id 作为 affinity_key实现相同用户稳定选择同一映射
# 传入 api_format 用于过滤适用的映射作用域
affinity_key = self.api_key.id if self.api_key else None
mapped_name = mapping.model.select_provider_model_name(
affinity_key, api_format=self.FORMAT_ID
)
logger.debug(
f"[CLI] 模型映射: {source_model} -> {mapped_name} (provider={provider_id[:8]}...)"
)
return mapped_name
logger.debug("[CLI] 无模型映射,使用原始名称: {}", source_model)
return None
def extract_model_from_request(
self: CliHandlerProtocol,
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002 - 子类使用
) -> str:
"""
从请求中提取模型名 - 子类可覆盖
不同 API 格式的 model 位置不同:
- OpenAI/Claude: 在请求体中 request_body["model"]
- Gemini: 在 URL 路径中 path_params["model"]
子类应覆盖此方法实现各自的提取逻辑。
Args:
request_body: 请求体
path_params: URL 路径参数
Returns:
模型名,如果无法提取则返回 "unknown"
"""
# 默认实现:从请求体获取
model = request_body.get("model")
return str(model) if model else "unknown"
def apply_mapped_model(
self,
request_body: dict[str, Any],
mapped_model: str, # noqa: ARG002 - 子类使用
) -> dict[str, Any]:
"""
将映射后的模型名应用到请求体
基类默认实现:不修改请求体,保持原样透传。
子类应覆盖此方法实现各自的模型名替换逻辑。
Args:
request_body: 原始请求体
mapped_model: 映射后的模型名(子类使用)
Returns:
请求体(默认不修改)
"""
# 基类不修改请求体,子类覆盖此方法实现特定格式的处理
return request_body
def prepare_provider_request_body(
self,
request_body: dict[str, Any],
) -> dict[str, Any]:
"""
准备发送给 Provider 的请求体 - 子类可覆盖
在模型映射之后、发送请求之前调用,用于移除不需要发送给上游的字段。
例如 Gemini API 需要移除请求体中的 model 字段(因为 model 在 URL 路径中)。
Args:
request_body: 经过模型映射处理后的请求体
Returns:
准备好的请求体
"""
return request_body
def finalize_provider_request(
self,
request_body: dict[str, Any],
*,
mapped_model: str | None,
provider_api_format: str | None,
) -> dict[str, Any]:
"""
格式转换完成后、envelope 之前的模型感知后处理钩子 - 子类可覆盖
用于根据目标模型的特性对请求体做最终调整,例如:
- 图像生成模型需要移除不兼容的 tools/system_instruction 并注入 imageConfig
- 特定模型需要注入/移除某些字段
- Gemini 格式:清理无效 parts 和合并连续同角色 contents
此方法在流式和非流式路径中均会被调用,且 mapped_model 已确定。
Args:
request_body: 已完成格式转换的请求体
mapped_model: 映射后的目标模型名
provider_api_format: Provider 侧 API 格式标识
Returns:
调整后的请求体
"""
# Gemini 格式请求:清理无效 parts 和合并连续同角色 contents
# 跨格式转换(如 Claude -> Gemini可能产生 thinking 等无法表示的块,
# 导致 parts 为空或缺少有效 data-oneof 字段,被 Google API 拒绝。
if provider_api_format and "gemini" in str(provider_api_format).lower():
contents = request_body.get("contents")
if isinstance(contents, list):
from src.core.api_format.conversion.normalizers.gemini import (
compact_gemini_contents,
)
request_body["contents"] = compact_gemini_contents(contents)
return request_body
async def _build_upstream_request(
self: CliHandlerProtocol,
*,
provider: Any,
endpoint: Any,
key: Any,
request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None,
client_api_format: str,
provider_api_format: str,
fallback_model: str,
mapped_model: str | None,
client_is_stream: bool,
needs_conversion: bool = False,
output_limit: int | None = None,
) -> CliUpstreamRequestResult:
"""Build the final outbound URL/body/headers for the selected upstream."""
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
behavior = get_provider_behavior(
provider_type=provider_type,
endpoint_sig=provider_api_format,
)
envelope = behavior.envelope
target_variant = behavior.same_format_variant
conversion_variant = behavior.cross_format_variant
upstream_policy = get_upstream_stream_policy(
endpoint,
provider_type=provider_type,
endpoint_sig=provider_api_format,
)
upstream_is_stream = resolve_upstream_is_stream(
client_is_stream=client_is_stream,
policy=upstream_policy,
)
envelope_tls_profile: str | None = None
if envelope and hasattr(envelope, "prepare_context"):
envelope_tls_profile = envelope.prepare_context(
provider_config=getattr(provider, "config", None),
key_id=str(getattr(key, "id", "") or ""),
user_api_key_id=str(getattr(self.api_key, "id", "") or ""),
is_stream=upstream_is_stream,
provider_id=str(getattr(provider, "id", "") or ""),
key=key,
)
if needs_conversion and provider_api_format:
request_body, url_model = await self._convert_request_for_cross_format(
request_body,
client_api_format,
provider_api_format,
mapped_model,
fallback_model,
is_stream=upstream_is_stream,
target_variant=conversion_variant,
output_limit=output_limit,
)
else:
request_body = self.prepare_provider_request_body(request_body)
url_model = (
self.get_model_for_url(request_body, mapped_model) or mapped_model or fallback_model
)
if target_variant and provider_api_format:
registry = get_format_converter_registry()
request_body = registry.convert_request(
request_body,
provider_api_format,
provider_api_format,
target_variant=target_variant,
)
request_body = self.finalize_provider_request(
request_body,
mapped_model=mapped_model,
provider_api_format=provider_api_format,
)
if provider_api_format:
enforce_stream_mode_for_upstream(
request_body,
provider_api_format=provider_api_format,
upstream_is_stream=upstream_is_stream,
)
request_body = maybe_patch_request_with_prompt_cache_key(
request_body,
provider_api_format=provider_api_format,
provider_type=provider_type,
base_url=getattr(endpoint, "base_url", None),
user_api_key_id=str(getattr(self.api_key, "id", "") or ""),
request_headers=original_headers,
)
auth_info = await get_provider_auth(endpoint, key)
if envelope:
request_body, url_model = envelope.wrap_request(
request_body,
model=url_model or fallback_model or "",
url_model=url_model,
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
)
if hasattr(envelope, "post_wrap_request"):
await envelope.post_wrap_request(request_body)
extra_headers: dict[str, str] = {}
if envelope:
extra_headers.update(envelope.extra_headers() or {})
hook_headers = build_upstream_extra_headers(
provider_type=provider_type,
endpoint_sig=provider_api_format,
request_body=request_body,
original_headers=original_headers,
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
)
if hook_headers:
extra_headers.update(hook_headers)
provider_payload, provider_headers = self._request_builder.build(
request_body,
original_headers,
endpoint,
key,
is_stream=upstream_is_stream,
extra_headers=extra_headers if extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=envelope,
provider_api_format=provider_api_format,
)
if upstream_is_stream:
set_accept_if_absent(provider_headers)
url = build_provider_url(
endpoint,
query_params=query_params,
path_params={"model": url_model},
is_stream=upstream_is_stream,
key=key,
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
)
selected_base_url = envelope.capture_selected_base_url() if envelope else None
return CliUpstreamRequestResult(
payload=provider_payload,
headers=provider_headers,
url=str(url),
url_model=str(url_model or fallback_model or ""),
envelope=envelope,
upstream_is_stream=upstream_is_stream,
tls_profile=envelope_tls_profile,
selected_base_url=selected_base_url,
)
@staticmethod
def _get_format_metadata(format_id: str) -> "EndpointDefinition | None":
"""获取 endpoint 元数据(解析失败返回 None"""
from src.core.api_format.metadata import resolve_endpoint_definition
return resolve_endpoint_definition(format_id)
def _finalize_converted_request(
self,
request_body: dict[str, Any],
client_api_format: str,
provider_api_format: str,
mapped_model: str | None,
fallback_model: str,
is_stream: bool,
) -> None:
"""
跨格式转换后统一设置并清理 model/stream 字段(原地修改)
处理逻辑:
1. 根据目标格式决定是否在 body 中设置 model
2. 若客户端格式不含 stream 字段但 Provider 需要,则显式设置
3. 移除目标格式不允许在 body 中携带的字段(如 Gemini 的 model/stream
Args:
request_body: 转换后的请求体(会被原地修改)
client_api_format: 客户端 API 格式
provider_api_format: Provider API 格式
mapped_model: 映射后的模型名
fallback_model: 备用模型名
is_stream: 是否流式请求
"""
client_meta = self._get_format_metadata(client_api_format)
provider_meta = self._get_format_metadata(provider_api_format)
# 默认model_in_body=True, stream_in_body=True如 OpenAI/Claude
client_uses_stream = client_meta.stream_in_body if client_meta else True
provider_model_in_body = provider_meta.model_in_body if provider_meta else True
provider_stream_in_body = provider_meta.stream_in_body if provider_meta else True
# 设置 model仅当 Provider 允许且 body 中需要)
if provider_model_in_body:
request_body["model"] = mapped_model or fallback_model
else:
request_body.pop("model", None)
# 设置 stream客户端不带但 Provider 需要时显式设置Provider 不需要时移除)
if provider_stream_in_body:
if not client_uses_stream:
request_body["stream"] = is_stream
else:
request_body.pop("stream", None)
# OpenAI Chat Completions: request usage in streaming mode.
provider_fmt = str(provider_api_format or "").strip().lower()
if is_stream and provider_fmt == "openai:chat":
stream_options = request_body.get("stream_options")
if not isinstance(stream_options, dict):
stream_options = {}
stream_options["include_usage"] = True
request_body["stream_options"] = stream_options
async 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 = None,
output_limit: int | None = None,
) -> tuple[dict[str, Any], str]:
"""
跨格式请求转换的公共逻辑
将客户端格式的请求体转换为 Provider 格式,并处理 model/stream 字段的补齐和清理。
Args:
request_body: 原始请求体(会被修改)
client_api_format: 客户端 API 格式
provider_api_format: Provider API 格式
mapped_model: 映射后的模型名
fallback_model: 备用模型名(通常是原始请求的 model
is_stream: 是否流式请求
target_variant: 目标变体(如 "codex"),用于同格式但有细微差异的上游
output_limit: GlobalModel 配置的模型输出上限
Returns:
(转换后的请求体, 用于 URL 的模型名)
"""
registry = get_format_converter_registry()
converted_body = await registry.convert_request_async(
request_body,
str(client_api_format),
str(provider_api_format),
target_variant=target_variant,
output_limit=output_limit,
)
# 先计算 URL 模型(在清理 body 中的 model 字段之前)
url_model = (
self.get_model_for_url(converted_body, mapped_model) or mapped_model or fallback_model
)
# 统一设置并清理 model/stream 字段
self._finalize_converted_request(
converted_body,
str(client_api_format),
str(provider_api_format),
mapped_model,
fallback_model,
is_stream,
)
return converted_body, url_model
def get_model_for_url(
self,
request_body: dict[str, Any],
mapped_model: str | None,
) -> str | None:
"""
获取用于 URL 路径的模型名
某些 API 格式(如 Gemini需要将 model 放入 URL 路径中。
子类应覆盖此方法返回正确的值。
Args:
request_body: 请求体
mapped_model: 映射后的模型名(如果有)
Returns:
用于 URL 路径的模型名,默认优先使用映射后的名称
"""
return mapped_model or request_body.get("model")
def _extract_response_metadata(
self,
response: dict[str, Any],
) -> dict[str, Any]:
"""
从响应中提取 Provider 特有的元数据 - 子类可覆盖
例如 Gemini 返回的 modelVersion 字段。
这些元数据会存储到 Usage.request_metadata 中。
Args:
response: Provider 返回的响应
Returns:
元数据字典,默认为空
"""
return {}

View File

@@ -0,0 +1,106 @@
"""SSE 解析辅助函数"""
from __future__ import annotations
import json
from typing import Any
from src.core.logger import logger
def _parse_sse_data_line(line: str) -> tuple[Any | None, str]:
"""
解析标准 SSE data 行
Args:
line: 以 "data:" 开头的 SSE 行
Returns:
(parsed_json, status) 元组:
- (parsed_dict, "ok") - 解析成功
- (None, "empty") - 内容为空
- (None, "invalid") - JSON 解析失败,调用方应透传原始行
"""
data_content = line[5:].strip()
if not data_content:
return None, "empty"
try:
return json.loads(data_content), "ok"
except json.JSONDecodeError:
return None, "invalid"
def _parse_sse_event_data_line(line: str) -> tuple[Any | None, str]:
"""
解析 event + data 同行格式(如 "event: xxx data: {...}"
Args:
line: 以 "event:" 开头且包含 " data:" 的 SSE 行
Returns:
(parsed_json, status) 元组
"""
_event_part, data_part = line.split(" data:", 1)
data_content = data_part.strip()
try:
return json.loads(data_content), "ok"
except json.JSONDecodeError:
return None, "invalid"
def _parse_gemini_json_array_line(line: str) -> tuple[Any | None, str]:
"""
解析 Gemini JSON-array 格式的裸 JSON 行
Gemini 流式响应可能是 JSON 数组格式,每行是数组元素。
Args:
line: 原始行(可能是 "[", "]", ",", 或 JSON 对象)
Returns:
(parsed_json, status) 元组
"""
stripped = line.strip()
if stripped in ("", "[", "]", ","):
return None, "skip"
candidate = stripped.lstrip(",").rstrip(",").strip()
try:
return json.loads(candidate), "ok"
except json.JSONDecodeError:
logger.debug("Gemini JSON-array line skip: {}", stripped[:50])
return None, "invalid"
def _format_converted_events_to_sse(
converted_events: list[dict[str, Any]],
client_format: str,
) -> list[str]:
"""
将转换后的事件格式化为 SSE 行
Args:
converted_events: 转换后的事件列表
client_format: 客户端 API 格式
Returns:
SSE 行列表(每个元素是完整的 SSE 事件,包含尾部空行)
"""
result: list[str] = []
client_fmt = str(client_format or "").strip().lower()
needs_event_line = client_fmt.startswith("claude:") or client_fmt == "openai:cli"
for evt in converted_events:
payload = json.dumps(evt, ensure_ascii=False)
if needs_event_line:
evt_type = evt.get("type") if isinstance(evt, dict) else None
if isinstance(evt_type, str) and evt_type:
# Claude 格式event + data + 空行
result.append(f"event: {evt_type}\ndata: {payload}\n")
else:
result.append(f"data: {payload}\n")
else:
# OpenAI 格式data + 空行
result.append(f"data: {payload}\n")
return result

View File

@@ -0,0 +1,946 @@
"""CLI Handler - 流式处理核心 Mixin"""
from __future__ import annotations
import asyncio
import codecs
import json
import time
from collections.abc import AsyncGenerator
from typing import TYPE_CHECKING, Any
import httpx
from fastapi import BackgroundTasks, Request
from fastapi.responses import StreamingResponse
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.utils import (
build_sse_headers,
ensure_stream_buffer_limit,
filter_proxy_response_headers,
get_format_converter_registry,
resolve_client_content_encoding,
)
from src.config.settings import config
from src.core.api_format.conversion.stream_bridge import (
iter_internal_response_as_stream_events,
)
from src.core.exceptions import (
EmbeddedErrorException,
ProviderNotAvailableException,
)
from src.core.logger import logger
from src.services.provider.behavior import get_provider_behavior
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
build_execution_plan_body,
is_remote_execution_runtime_contract_eligible,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.system.config import SystemConfigService
from src.services.task.request_state import MutableRequestBodyState
from src.utils.sse_parser import SSEEventParser
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
class CliStreamMixin:
"""流式处理核心方法的 Mixin"""
def _streamify_sync_response(
self: CliHandlerProtocol,
*,
ctx: StreamContext,
response_json: dict[str, Any],
client_api_format: str,
provider_api_format: str,
) -> AsyncGenerator[bytes]:
registry = get_format_converter_registry()
src_norm = registry.get_normalizer(provider_api_format) if provider_api_format else None
if src_norm is None:
raise RuntimeError(f"未注册 Normalizer: {provider_api_format}")
internal_resp = src_norm.response_to_internal(
response_json if isinstance(response_json, dict) else {}
)
internal_resp.model = str(ctx.model or internal_resp.model or "")
if internal_resp.id:
ctx.response_id = internal_resp.id
if internal_resp.usage:
ctx.input_tokens = int(internal_resp.usage.input_tokens or 0)
ctx.output_tokens = int(internal_resp.usage.output_tokens or 0)
ctx.cached_tokens = int(internal_resp.usage.cache_read_tokens or 0)
ctx.cache_creation_tokens = int(internal_resp.usage.cache_write_tokens or 0)
from src.core.api_format.conversion.stream_state import StreamState
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}")
state = StreamState(
model=str(ctx.model or ""),
message_id=str(ctx.response_id or ctx.request_id or self.request_id or ""),
)
output_state = {"first_yield": True, "streaming_updated": False}
async def _streamified() -> AsyncGenerator[bytes]:
for ev in iter_internal_response_as_stream_events(internal_resp):
converted_events = tgt_norm.stream_event_from_internal(ev, state)
if not converted_events:
continue
self._record_converted_chunks(ctx, converted_events)
for sse_line in _format_converted_events_to_sse(
converted_events, client_api_format
):
if not sse_line:
continue
ctx.chunk_count += 1
self._mark_first_output(ctx, output_state)
yield (sse_line + "\n").encode("utf-8")
if str(client_api_format or "").strip().lower() == "openai:chat":
ctx.chunk_count += 1
self._mark_first_output(ctx, output_state)
yield b"data: [DONE]\n\n"
ctx.has_completion = True
return _streamified()
async def process_stream(
self: CliHandlerProtocol,
original_request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
http_request: Request | None = None,
client_content_encoding: str | None = None,
) -> StreamingResponse:
"""
处理流式请求
通用流程:
1. 创建流上下文
2. 定义请求函数(供 TaskService/FailoverEngine 调用)
3. 执行请求并返回 StreamingResponse
4. 后台任务记录统计信息
Args:
original_request_body: 原始请求体
original_headers: 原始请求头
query_params: 查询参数
path_params: 路径参数
http_request: FastAPI Request 对象,用于检测客户端断连
"""
logger.debug("开始流式响应处理 ({})", self.FORMAT_ID)
effective_client_content_encoding = resolve_client_content_encoding(
original_headers,
client_content_encoding,
)
request_state = MutableRequestBodyState(original_request_body)
# 使用子类实现的方法提取 model不同 API 格式的 model 位置不同)
# 注意:使用 original_request_body因为整流只修改 messages不影响 model 字段
model = self.extract_model_from_request(original_request_body, path_params)
client_api_format = self.primary_api_format
# 提前创建 pending 记录,让前端可以立即看到"处理中"
pending_usage_created = self._create_pending_usage(
model=model,
is_stream=True,
request_type="chat",
api_format=client_api_format,
request_headers=original_headers,
request_body=original_request_body,
)
# 创建流上下文
ctx = StreamContext(
model=model,
api_format=client_api_format,
api_family=self.api_family,
endpoint_kind=self.endpoint_kind,
request_id=self.request_id,
user_id=self.user.id,
api_key_id=self.api_key.id,
)
# 仅在 FULL 级别才需要保留 parsed_chunks避免长流式响应导致的内存占用
ctx.record_parsed_chunks = SystemConfigService.should_log_body(self.db)
request_metadata = self._build_request_metadata(http_request)
if request_metadata and isinstance(request_metadata.get("perf"), dict):
ctx.perf_sampled = True
ctx.perf_metrics.update(request_metadata["perf"])
# 定义请求函数
async def stream_request_func(
provider: "Provider",
endpoint: "ProviderEndpoint",
key: "ProviderAPIKey",
candidate: ProviderCandidate,
) -> AsyncGenerator[bytes]:
return await self._execute_stream_request(
ctx,
provider,
endpoint,
key,
request_state.build_attempt_body(),
original_headers,
query_params,
candidate,
http_request, # 传递 http_request 用于断连检测
effective_client_content_encoding,
)
try:
# 解析能力需求
capability_requirements = self._resolve_capability_requirements(
model_name=ctx.model,
request_headers=original_headers,
request_body=original_request_body,
)
preferred_key_ids = await self._resolve_preferred_key_ids(
model_name=ctx.model,
request_body=original_request_body,
)
# 统一入口:总是通过 TaskService
from src.services.task import TaskService
from src.services.task.core.context import TaskMode
exec_result = await TaskService(self.db, self.redis).execute(
task_type="cli",
task_mode=TaskMode.SYNC,
api_format=ctx.api_format,
model_name=ctx.model,
user_api_key=self.api_key,
request_func=stream_request_func,
request_id=self.request_id,
is_stream=True,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body_state=request_state,
request_headers=original_headers,
request_body=original_request_body,
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
create_pending_usage=not pending_usage_created,
)
stream_generator = exec_result.response
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 信息已设置,用于 streaming 状态更新)
ctx.attempt_id = attempt_id
if not ctx.provider_name:
ctx.provider_name = provider_name
if not ctx.provider_id:
ctx.provider_id = provider_id
if not ctx.endpoint_id:
ctx.endpoint_id = endpoint_id
if not ctx.key_id:
ctx.key_id = key_id
if getattr(exec_result, "pool_summary", None):
ctx.pool_summary = exec_result.pool_summary
scheduling_metadata = (
self._merge_scheduling_metadata(
{},
exec_result=exec_result,
selected_key_id=key_id,
fallback_from_request=False,
)
or {}
)
candidate_keys = scheduling_metadata.get("candidate_keys")
if isinstance(candidate_keys, list):
ctx.candidate_keys = candidate_keys
scheduling_audit = scheduling_metadata.get("scheduling_audit")
if isinstance(scheduling_audit, dict):
ctx.scheduling_audit = scheduling_audit
# 同步整流状态(如果请求体被整流过)
ctx.rectified = request_state.is_rectified()
# 创建后台任务记录统计
background_tasks = BackgroundTasks()
background_tasks.add_task(
self._record_stream_stats,
ctx,
original_headers,
original_request_body,
)
# 创建监控流(传递 http_request 用于断连检测)
monitored_stream = self._create_monitored_stream(ctx, stream_generator, http_request)
# 透传提供商的响应头给客户端
# 同时添加必要的 SSE 头以确保流式传输正常工作
client_headers = filter_proxy_response_headers(ctx.response_headers)
# 添加/覆盖 SSE 必需的头
client_headers.update(build_sse_headers())
client_headers["content-type"] = "text/event-stream"
ctx.client_response_headers = client_headers
return StreamingResponse(
monitored_stream,
media_type="text/event-stream",
headers=client_headers,
background=background_tasks,
)
except Exception as e:
from src.core.exceptions import ThinkingSignatureException
if isinstance(e, ThinkingSignatureException):
# Thinking 签名错误TaskService 层已处理整流重试但仍失败
# 记录 original_request_body客户端原始请求便于排查问题根因
self._log_request_error("流式请求失败(签名错误)", e)
else:
self._log_request_error("流式请求失败", e)
await self._record_stream_failure(ctx, e, original_headers, original_request_body)
raise
async def _execute_stream_request(
self,
ctx: StreamContext,
provider: "Provider",
endpoint: "ProviderEndpoint",
key: "ProviderAPIKey",
working_request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
candidate: ProviderCandidate | None = None,
http_request: Request | None = None,
client_content_encoding: str | None = None,
) -> AsyncGenerator[bytes]:
"""执行流式请求并返回流生成器"""
# 重置上下文状态(重试时清除之前的数据,避免累积)
ctx.release_recorded_chunks()
ctx.chunk_count = 0
ctx.data_count = 0
ctx.has_completion = False
ctx._collected_text_parts = [] # 重置文本收集
ctx.input_tokens = 0
ctx.output_tokens = 0
ctx.cached_tokens = 0
ctx.cache_creation_tokens = 0
ctx.final_usage = None
ctx.final_response = None
ctx.response_id = None
ctx.response_metadata = {} # 重置 Provider 响应元数据
ctx.selected_base_url = None # 重置本次请求选用的 base_url重试时避免污染
# 记录 Provider 信息
ctx.provider_name = str(provider.name)
ctx.provider_id = str(provider.id)
ctx.provider_type = str(getattr(provider, "provider_type", "") or "")
ctx.endpoint_id = str(endpoint.id)
ctx.key_id = str(key.id)
# 记录格式转换信息
ctx.provider_api_format = str(endpoint.api_format) if endpoint.api_format else ""
ctx.client_api_format = ctx.api_format # 已在 process_stream 中设置
# 获取模型映射(优先使用映射匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.mapping_matched_model if candidate else None
if not mapped_model:
mapped_model = await self._get_mapped_model(
source_model=ctx.model,
provider_id=str(provider.id),
)
# `working_request_body` is already isolated per attempt.
request_body = working_request_body
if mapped_model:
ctx.mapped_model = mapped_model # 保存映射后的模型名,用于 Usage 记录
request_body = self.apply_mapped_model(request_body, mapped_model)
client_api_format = (
ctx.client_api_format.value
if hasattr(ctx.client_api_format, "value")
else str(ctx.client_api_format)
)
provider_api_format = str(ctx.provider_api_format or "")
needs_conversion = (
bool(getattr(candidate, "needs_conversion", False)) if candidate else False
)
ctx.needs_conversion = needs_conversion
upstream_request = await self._build_upstream_request(
provider=provider,
endpoint=endpoint,
key=key,
request_body=request_body,
original_headers=original_headers,
query_params=query_params,
client_api_format=client_api_format,
provider_api_format=provider_api_format,
fallback_model=ctx.model,
mapped_model=mapped_model,
client_is_stream=True,
needs_conversion=needs_conversion,
output_limit=candidate.output_limit if candidate else None,
)
provider_headers = upstream_request.headers
provider_payload = upstream_request.payload
url = upstream_request.url
envelope = upstream_request.envelope
upstream_is_stream = upstream_request.upstream_is_stream
envelope_tls_profile = upstream_request.tls_profile
# 保存发送给 Provider 的请求信息(用于调试和统计)
ctx.provider_request_headers = provider_headers
ctx.provider_request_body = provider_payload
ctx.selected_base_url = upstream_request.selected_base_url
# 解析有效代理Key 级别优先于 Provider 级别)
from src.services.proxy_node.resolver import build_proxy_url_async as _bpua
from src.services.proxy_node.resolver import get_proxy_label as _gpl
from src.services.proxy_node.resolver import resolve_delegate_config_async as _rda
from src.services.proxy_node.resolver import resolve_effective_proxy as _rep
from src.services.proxy_node.resolver import resolve_proxy_info_async as _rpi_async
effective_proxy = _rep(provider.proxy, getattr(key, "proxy", None))
ctx.proxy_info = await _rpi_async(effective_proxy)
delegate_cfg = await _rda(effective_proxy)
is_tunnel_delegate = bool(delegate_cfg and delegate_cfg.get("tunnel"))
proxy_url: str | None = None
if effective_proxy and not is_tunnel_delegate:
proxy_url = await _bpua(effective_proxy)
proxy_snapshot = ExecutionProxySnapshot.from_proxy_info(
ctx.proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if is_tunnel_delegate else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None
if is_tunnel_delegate
else None
),
)
# If upstream is forced to non-stream mode, we execute a sync request and then
# simulate streaming to the client (sync -> stream bridge).
if not upstream_is_stream:
request_timeout_sync = provider.request_timeout or config.http_request_timeout
rust_plan = ExecutionPlan(
request_id=str(self.request_id or ""),
candidate_id=str(
getattr(candidate, "request_candidate_id", "")
or getattr(candidate, "id", "")
or ""
)
or None,
provider_name=str(provider.name),
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
method="POST",
url=url,
headers=dict(provider_headers),
body=build_execution_plan_body(
provider_payload,
content_type=str(provider_headers.get("content-type") or "").strip() or None,
),
stream=False,
provider_api_format=provider_api_format,
client_api_format=client_api_format,
model_name=str(ctx.model or ""),
content_type=str(provider_headers.get("content-type") or "").strip() or None,
content_encoding=client_content_encoding,
proxy=proxy_snapshot,
tls_profile=envelope_tls_profile,
timeouts=ExecutionPlanTimeouts(
connect_ms=int(config.http_connect_timeout * 1000),
read_ms=int(config.http_read_timeout * 1000),
write_ms=int(config.http_write_timeout * 1000),
pool_ms=int(config.http_pool_timeout * 1000),
total_ms=int(request_timeout_sync * 1000),
),
)
if not is_remote_execution_runtime_contract_eligible(rust_plan):
raise ProviderNotAvailableException(
"CLI 请求暂不支持当前 Rust executor 契约",
provider_name=str(provider.name),
upstream_response="remote_contract_ineligible",
)
try:
rust_result = await ExecutionRuntimeClient().execute_sync_json(rust_plan)
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError) as exc:
logger.warning(
"[{}] CLI Rust executor(sync->stream) unavailable: {}",
self.request_id,
exc,
)
raise ProviderNotAvailableException(
"执行器暂时不可用,请稍后重试",
provider_name=str(provider.name),
upstream_response=str(exc),
) from exc
ctx.status_code = rust_result.status_code
ctx.response_headers = dict(rust_result.headers)
ctx.set_proxy_timing(ctx.response_headers)
if envelope:
envelope.on_http_status(
base_url=ctx.selected_base_url,
status_code=ctx.status_code,
)
request = httpx.Request("POST", url, headers=provider_headers)
synthetic_content = rust_result.response_body_bytes
if synthetic_content is None:
synthetic_content = json.dumps(
rust_result.response_json or {},
ensure_ascii=False,
).encode("utf-8")
synthetic_response = httpx.Response(
ctx.status_code,
request=request,
headers=ctx.response_headers,
content=synthetic_content,
)
if ctx.status_code >= 400:
error = httpx.HTTPStatusError(
f"Upstream status error: {ctx.status_code}",
request=request,
response=synthetic_response,
)
error_body = ""
try:
if envelope and hasattr(envelope, "extract_error_text"):
error_body = await envelope.extract_error_text(synthetic_response)
else:
error_body = synthetic_response.text[:4000] if synthetic_response.text else ""
except Exception:
error_body = synthetic_response.text[:4000] if synthetic_response.text else ""
error.upstream_response = error_body # type: ignore[attr-defined]
raise error
response_json = rust_result.response_json or {}
if envelope:
response_json = envelope.unwrap_response(response_json)
envelope.postprocess_unwrapped_response(model=ctx.model, data=response_json)
if isinstance(response_json, dict) and provider_api_format:
parser = get_parser_for_format(provider_api_format)
if parser.is_error_response(response_json):
parsed = parser.parse_response(response_json, 200)
raise EmbeddedErrorException(
provider_name=str(provider.name),
error_code=parsed.embedded_status_code,
error_message=parsed.error_message,
error_status=parsed.error_type,
)
if isinstance(response_json, dict):
ctx.response_metadata = self._extract_response_metadata(response_json)
return self._streamify_sync_response(
ctx=ctx,
response_json=response_json if isinstance(response_json, dict) else {},
client_api_format=client_api_format,
provider_api_format=provider_api_format,
)
# 流式请求使用 stream_first_byte_timeout 作为首字节超时
# 优先使用 Provider 配置,否则使用全局配置
request_timeout = provider.stream_first_byte_timeout or config.stream_first_byte_timeout
_proxy_label = _gpl(ctx.proxy_info)
logger.debug(
f" └─ [{self.request_id}] 发送流式请求: "
f"Provider={provider.name}, Endpoint={endpoint.id[:8] if endpoint.id else 'N/A'}..., "
f"Key=***{key.api_key[-4:] if key.api_key else 'N/A'}, "
f"原始模型={ctx.model}, 映射后={mapped_model or '无映射'}, URL模型={upstream_request.url_model}, "
f"timeout={request_timeout}s, 代理={_proxy_label}"
)
rust_plan = ExecutionPlan(
request_id=str(self.request_id or ""),
candidate_id=str(
getattr(candidate, "request_candidate_id", "") or getattr(candidate, "id", "") or ""
)
or None,
provider_name=str(provider.name),
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
method="POST",
url=url,
headers=dict(provider_headers),
body=build_execution_plan_body(
provider_payload,
content_type=str(provider_headers.get("content-type") or "").strip() or None,
),
stream=True,
provider_api_format=provider_api_format,
client_api_format=client_api_format,
model_name=str(ctx.model or ""),
content_type=str(provider_headers.get("content-type") or "").strip() or None,
content_encoding=client_content_encoding,
proxy=proxy_snapshot,
tls_profile=envelope_tls_profile,
timeouts=ExecutionPlanTimeouts(
connect_ms=int(config.http_connect_timeout * 1000),
read_ms=int(config.http_read_timeout * 1000),
write_ms=int(config.http_write_timeout * 1000),
pool_ms=int(config.http_pool_timeout * 1000),
total_ms=None,
),
)
if not is_remote_execution_runtime_contract_eligible(rust_plan):
raise ProviderNotAvailableException(
"CLI 请求暂不支持当前 Rust executor 契约",
provider_name=str(provider.name),
upstream_response="remote_contract_ineligible",
)
try:
rust_stream = await ExecutionRuntimeClient().execute_stream(rust_plan)
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError) as exc:
logger.warning(
"[{}] CLI Rust executor stream unavailable: {}",
self.request_id,
exc,
)
raise ProviderNotAvailableException(
"执行器暂时不可用,请稍后重试",
provider_name=str(provider.name),
upstream_response=str(exc),
) from exc
ctx.status_code = rust_stream.status_code
ctx.response_headers = dict(rust_stream.headers)
ctx.set_proxy_timing(ctx.response_headers)
if envelope:
envelope.on_http_status(
base_url=ctx.selected_base_url,
status_code=ctx.status_code,
)
try:
if ctx.status_code >= 400:
error_chunks: list[bytes] = []
async for chunk in rust_stream.byte_iterator:
if chunk:
error_chunks.append(chunk)
if sum(len(item) for item in error_chunks) >= 4000:
break
error_body = b"".join(error_chunks)[:4000].decode(
"utf-8",
errors="replace",
)
request = httpx.Request("POST", url, headers=provider_headers)
response = httpx.Response(
ctx.status_code,
request=request,
headers=rust_stream.headers,
content=b"".join(error_chunks),
)
error = httpx.HTTPStatusError(
f"Upstream status error: {ctx.status_code}",
request=request,
response=response,
)
error.upstream_response = error_body # type: ignore[attr-defined]
raise error
prefetched_chunks = await self._prefetch_and_check_embedded_error(
rust_stream.byte_iterator,
provider,
endpoint,
ctx,
)
except Exception:
await rust_stream.response_ctx.__aexit__(None, None, None)
raise
return self._create_response_stream_with_prefetch(
ctx,
rust_stream.byte_iterator,
rust_stream.response_ctx,
prefetched_chunks,
)
@staticmethod
def _fire_stream_timeout_policy(ctx: StreamContext) -> None:
"""Fire-and-forget: record stream timeout for pool health policy."""
if not ctx.provider_id or not ctx.key_id:
return
try:
from src.services.provider.adapters.claude_code.context import (
get_claude_code_request_context,
)
cc_ctx = get_claude_code_request_context()
pool_cfg = cc_ctx.pool_config if cc_ctx else None
if not pool_cfg:
return
from src.services.provider.pool.health_policy import apply_stream_timeout_policy
task = asyncio.create_task(
apply_stream_timeout_policy(
provider_id=ctx.provider_id,
key_id=ctx.key_id,
config=pool_cfg,
)
)
task.add_done_callback(lambda t: t.exception() if not t.cancelled() else None)
except Exception as exc:
logger.debug("Stream timeout policy trigger failed: {}", exc)
async def _create_response_stream(
self,
ctx: StreamContext,
stream_response: httpx.Response,
response_ctx: Any,
) -> AsyncGenerator[bytes]:
"""创建响应流生成器(使用字节流)"""
try:
sse_parser = SSEEventParser()
last_data_time = time.time()
buffer = b""
output_state = {"first_yield": True, "streaming_updated": False}
_sample_lines: list[str] = [] # 采集前几行原始内容,用于空流诊断
_MAX_SAMPLE_LINES = 5
# 使用增量解码器处理跨 chunk 的 UTF-8 字符
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
# 使用已设置的 ctx.needs_conversion由候选筛选阶段根据端点配置判断
# 不再调用 _needs_format_conversion它只检查格式差异不检查端点配置
needs_conversion = ctx.needs_conversion
behavior = get_provider_behavior(
provider_type=ctx.provider_type,
endpoint_sig=ctx.provider_api_format,
)
envelope = behavior.envelope
if envelope and envelope.force_stream_rewrite():
needs_conversion = True
ctx.needs_conversion = True
# Kiro 特殊处理AWS Event Stream 二进制流需要重写为 SSE
ctx_provider_type = str(ctx.provider_type or "").strip().lower()
if ctx_provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
from src.services.provider.adapters.kiro.eventstream_rewriter import (
apply_kiro_stream_rewrite,
)
chunk_source: AsyncGenerator[bytes, None] = apply_kiro_stream_rewrite(
stream_response.aiter_bytes(),
model=str(ctx.model or ""),
input_tokens=int(ctx.input_tokens or 0),
)
# Kiro 重写后输出的是 Claude SSE 格式,不需要再进行格式转换
needs_conversion = False
ctx.needs_conversion = False
else:
chunk_source = stream_response.aiter_bytes()
async for chunk in chunk_source:
buffer += chunk
ensure_stream_buffer_limit(
buffer,
request_id=self.request_id,
provider_name=ctx.provider_name,
)
# 处理缓冲区中的完整行
while b"\n" in buffer:
line_bytes, buffer = buffer.split(b"\n", 1)
try:
# 使用增量解码器,可以正确处理跨 chunk 的多字节字符
line = decoder.decode(line_bytes + b"\n", False).rstrip("\n")
except Exception as e:
logger.warning(
"[{}] UTF-8 解码失败: {}, bytes={!r}",
self.request_id,
e,
line_bytes[:50],
)
continue
normalized_line = line.rstrip("\r")
events = sse_parser.feed_line(normalized_line)
if normalized_line == "":
for event in events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=not needs_conversion,
)
self._mark_first_output(ctx, output_state)
yield b"\n"
continue
ctx.chunk_count += 1
if len(_sample_lines) < _MAX_SAMPLE_LINES:
_sample_lines.append(normalized_line[:200])
# 空流检测:超过阈值且无数据,发送错误事件并结束
if ctx.chunk_count > self.EMPTY_CHUNK_THRESHOLD and ctx.data_count == 0:
elapsed = time.time() - last_data_time
if elapsed > self.DATA_TIMEOUT:
logger.warning("Provider '{}' 流超时且无数据", ctx.provider_name)
ctx.status_code = 504
ctx.error_message = "流式响应超时,未收到有效数据"
ctx.upstream_response = (
f"流超时: Provider={ctx.provider_name}, "
f"elapsed={elapsed:.1f}s, "
f"chunk_count={ctx.chunk_count}, data_count=0"
)
self._fire_stream_timeout_policy(ctx)
error_event = {
"type": "error",
"error": {
"type": "empty_stream_timeout",
"message": ctx.error_message,
},
}
self._mark_first_output(ctx, output_state)
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
return # 结束生成器
# 格式转换或直接透传
if needs_conversion:
converted_lines, converted_events = self._convert_sse_line(
ctx, line, events
)
# 记录转换后的数据到 parsed_chunks
self._record_converted_chunks(ctx, converted_events)
for converted_line in converted_lines:
if converted_line:
self._mark_first_output(ctx, output_state)
yield (converted_line + "\n").encode("utf-8")
else:
self._mark_first_output(ctx, output_state)
yield (line + "\n").encode("utf-8")
for event in events:
self._handle_sse_event(
ctx,
event.get("event"),
event.get("data") or "",
record_chunk=not needs_conversion,
)
if ctx.data_count > 0:
last_data_time = time.time()
# flush 字节 buffer 残余数据 + SSE parser 内部缓冲区
for chunk in self._flush_buffer_with_conversion(
ctx, buffer, decoder, sse_parser, needs_conversion
):
yield chunk
# 检查是否收到数据
if ctx.data_count == 0:
sample_info = f", 前几行内容: {_sample_lines!r}" if _sample_lines else ""
logger.warning(
"Provider '{}' 返回空流式响应{}",
ctx.provider_name,
sample_info,
)
ctx.status_code = 503
ctx.error_message = "上游服务返回了空的流式响应"
ctx.upstream_response = (
f"空流式响应: Provider={ctx.provider_name}, "
f"chunk_count={ctx.chunk_count}, data_count=0"
)
error_event = {
"type": "error",
"error": {
"type": "empty_response",
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
logger.debug("流式数据转发完成")
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
client_fmt = (ctx.client_api_format or "").strip().lower()
if needs_conversion and client_fmt == "openai:chat":
yield b"data: [DONE]\n\n"
except GeneratorExit:
raise
except httpx.StreamClosed:
# 连接关闭前 flush 残余数据,尝试捕获尾部事件(如 response.completed 中的 usage
self._flush_remaining_sse_data(
ctx, buffer, decoder, sse_parser, record_chunk=not needs_conversion
)
if ctx.data_count == 0:
# 流已开始,发送错误事件而不是抛出异常
logger.warning("Provider '{}' 流连接关闭且无数据", ctx.provider_name)
# 设置错误状态用于后续记录
ctx.status_code = 503
ctx.error_message = "上游服务连接关闭且未返回数据"
ctx.upstream_response = f"流连接关闭: Provider={ctx.provider_name}, chunk_count={ctx.chunk_count}, data_count=0"
error_event = {
"type": "error",
"error": {
"type": "stream_closed",
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
except httpx.RemoteProtocolError:
# 连接异常关闭前 flush 残余数据,尝试捕获尾部事件(如 response.completed 中的 usage
self._flush_remaining_sse_data(
ctx, buffer, decoder, sse_parser, record_chunk=not needs_conversion
)
if ctx.data_count > 0:
error_event = {
"type": "error",
"error": {
"type": "connection_error",
"message": "上游连接意外关闭,部分响应已成功传输",
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
raise
except httpx.ReadError:
# 代理/上游连接读取失败(如 aether-proxy 中断),与 RemoteProtocolError 处理逻辑一致
self._flush_remaining_sse_data(
ctx, buffer, decoder, sse_parser, record_chunk=not needs_conversion
)
if ctx.data_count > 0:
error_event = {
"type": "error",
"error": {
"type": "connection_error",
"message": "代理或上游连接读取失败,部分响应已成功传输",
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
raise
finally:
try:
await response_ctx.__aexit__(None, None, None)
except Exception:
pass

View File

@@ -0,0 +1,634 @@
"""CLI Handler - 同步处理 Mixin"""
from __future__ import annotations
import json
import time
from typing import TYPE_CHECKING, Any
import httpx
from fastapi.responses import JSONResponse
from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.stream_context import extract_proxy_timing, is_format_converted
from src.api.handlers.base.upstream_stream_bridge import (
aggregate_upstream_stream_to_internal_response,
)
from src.api.handlers.base.utils import (
build_json_response_for_client,
filter_proxy_response_headers,
get_format_converter_registry,
resolve_client_accept_encoding,
resolve_client_content_encoding,
)
from src.config.settings import config
from src.core.error_utils import extract_client_error_message
from src.core.exceptions import (
ProviderAuthException,
ProviderNotAvailableException,
ProviderRateLimitException,
ProviderTimeoutException,
ThinkingSignatureException,
)
from src.core.logger import logger
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
build_execution_plan_body,
is_remote_execution_runtime_contract_eligible,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.task.request_state import MutableRequestBodyState
if TYPE_CHECKING:
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
class CliSyncMixin:
"""同步处理相关方法的 Mixin"""
async def _aggregate_upstream_stream_sync_response(
self: CliHandlerProtocol,
*,
body_bytes: bytes,
provider_api_format: str,
client_api_format: str,
provider_name: str,
provider_type: str,
model: str,
request_id: str,
envelope: Any,
) -> dict[str, Any]:
registry = get_format_converter_registry()
provider_parser = (
get_parser_for_format(provider_api_format) if provider_api_format else None
)
async def _byte_iter() -> Any:
yield body_bytes
byte_iter = _byte_iter()
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=provider_name,
model=model,
request_id=request_id,
envelope=envelope,
provider_parser=provider_parser,
)
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,
)
return response_json if isinstance(response_json, dict) else {}
async def process_sync(
self: CliHandlerProtocol,
original_request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
client_content_encoding: str | None = None,
client_accept_encoding: str | None = None,
) -> JSONResponse:
"""
处理非流式请求
通用流程:
1. 构建请求
2. 通过 TaskService/FailoverEngine 执行
3. 解析响应并记录统计
"""
logger.debug("开始非流式响应处理 ({})", self.FORMAT_ID)
effective_client_content_encoding = resolve_client_content_encoding(
original_headers,
client_content_encoding,
)
effective_client_accept_encoding = resolve_client_accept_encoding(
original_headers,
client_accept_encoding,
)
# 使用子类实现的方法提取 model不同 API 格式的 model 位置不同)
model = self.extract_model_from_request(original_request_body, path_params)
api_format = self.primary_api_format
sync_start_time = time.time()
# 提前创建 pending 记录,让前端可以立即看到"处理中"
pending_usage_created = self._create_pending_usage(
model=model,
is_stream=False,
request_type="chat",
api_format=api_format,
request_headers=original_headers,
request_body=original_request_body,
)
provider_name = None
response_json = None
status_code = 200
response_headers = {}
provider_api_format = "" # 用于追踪 Provider 的 API 格式
provider_request_headers = {} # 发送给 Provider 的请求头
provider_request_body = None # 实际发送给 Provider 的请求体
provider_id = None # Provider ID用于失败记录
endpoint_id = None # Endpoint ID用于失败记录
key_id = None # Key ID用于失败记录
exec_result = None
mapped_model_result = None # 映射后的目标模型名(用于 Usage 记录)
response_metadata_result: dict[str, Any] = {} # Provider 响应元数据
needs_conversion = False # 是否需要格式转换(由 candidate 决定)
sync_proxy_info: dict[str, Any] | None = None # 代理信息
request_state = MutableRequestBodyState(original_request_body)
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, provider_api_format, provider_request_headers, provider_request_body, mapped_model_result, response_metadata_result, needs_conversion, sync_proxy_info
provider_name = str(provider.name)
provider_api_format = str(endpoint.api_format) if endpoint.api_format else ""
# 获取模型映射(优先使用映射匹配到的模型,其次是 Provider 级别的映射)
mapped_model = candidate.mapping_matched_model if candidate else None
if not mapped_model:
mapped_model = await self._get_mapped_model(
source_model=model,
provider_id=str(provider.id),
)
request_body = request_state.build_attempt_body()
if mapped_model:
mapped_model_result = mapped_model # 保存映射后的模型名,用于 Usage 记录
request_body = self.apply_mapped_model(request_body, mapped_model)
client_api_format = (
api_format.value if hasattr(api_format, "value") else str(api_format)
)
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
upstream_request = await self._build_upstream_request(
provider=provider,
endpoint=endpoint,
key=key,
request_body=request_body,
original_headers=original_headers,
query_params=query_params,
client_api_format=client_api_format,
provider_api_format=provider_api_format,
fallback_model=model,
mapped_model=mapped_model,
client_is_stream=False,
needs_conversion=needs_conversion,
output_limit=candidate.output_limit if candidate else None,
)
provider_headers = upstream_request.headers
provider_payload = upstream_request.payload
provider_request_headers = provider_headers
provider_request_body = provider_payload
url = upstream_request.url
envelope = upstream_request.envelope
upstream_is_stream = upstream_request.upstream_is_stream
envelope_tls_profile = upstream_request.tls_profile
selected_base_url_cached = upstream_request.selected_base_url
# 解析有效代理Key 级别优先于 Provider 级别)
from src.services.proxy_node.resolver import (
get_proxy_label,
resolve_effective_proxy,
resolve_proxy_info_async,
)
_effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
sync_proxy_info = await resolve_proxy_info_async(_effective_proxy)
_proxy_label = get_proxy_label(sync_proxy_info)
logger.info(
f" └─ [{self.request_id}] 发送{'上游流式(聚合)' if upstream_is_stream else '非流式'}请求: "
f"Provider={provider.name}, Endpoint={endpoint.id[:8] if endpoint.id else 'N/A'}..., "
f"Key=***{key.api_key[-4:] if key.api_key else 'N/A'}, "
f"原始模型={model}, 映射后={mapped_model or '无映射'}, URL模型={upstream_request.url_model}, "
f"代理={_proxy_label}"
)
from src.services.proxy_node.resolver import (
build_proxy_url_async,
resolve_delegate_config_async,
)
# 非流式请求使用 http_request_timeout 作为整体超时
# 优先使用 Provider 配置,否则使用全局配置
request_timeout = provider.request_timeout or config.http_request_timeout
delegate_cfg = await resolve_delegate_config_async(_effective_proxy)
is_tunnel_delegate = bool(delegate_cfg and delegate_cfg.get("tunnel"))
proxy_url: str | None = None
if _effective_proxy and not is_tunnel_delegate:
proxy_url = await build_proxy_url_async(_effective_proxy)
rust_plan = ExecutionPlan(
request_id=str(self.request_id or ""),
candidate_id=str(
getattr(candidate, "request_candidate_id", "")
or getattr(candidate, "id", "")
or ""
)
or None,
provider_name=str(provider.name),
provider_id=str(provider.id),
endpoint_id=str(endpoint.id),
key_id=str(key.id),
method="POST",
url=url,
headers=dict(provider_headers),
body=build_execution_plan_body(
provider_payload,
content_type=str(provider_headers.get("content-type") or "").strip() or None,
),
stream=upstream_is_stream,
provider_api_format=provider_api_format,
client_api_format=client_api_format,
model_name=str(model or ""),
content_type=str(provider_headers.get("content-type") or "").strip() or None,
content_encoding=effective_client_content_encoding,
proxy=ExecutionProxySnapshot.from_proxy_info(
sync_proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if is_tunnel_delegate else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None
if is_tunnel_delegate
else None
),
),
tls_profile=envelope_tls_profile,
timeouts=ExecutionPlanTimeouts(
connect_ms=int(config.http_connect_timeout * 1000),
read_ms=int(config.http_read_timeout * 1000),
write_ms=int(config.http_write_timeout * 1000),
pool_ms=int(config.http_pool_timeout * 1000),
total_ms=int(request_timeout * 1000),
),
)
if not is_remote_execution_runtime_contract_eligible(rust_plan):
raise ProviderNotAvailableException(
"CLI 请求暂不支持当前 Rust executor 契约",
provider_name=str(provider.name),
upstream_response="remote_contract_ineligible",
)
try:
rust_result = await ExecutionRuntimeClient().execute_sync_json(rust_plan)
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError) as exc:
logger.warning(
"[{}] CLI Rust executor unavailable: {}",
self.request_id,
exc,
)
raise ProviderNotAvailableException(
"执行器暂时不可用,请稍后重试",
provider_name=str(provider.name),
upstream_response=str(exc),
) from exc
status_code = rust_result.status_code
response_headers = dict(rust_result.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,
)
request = httpx.Request("POST", url, headers=provider_headers)
synthetic_content = rust_result.response_body_bytes
if synthetic_content is None:
synthetic_content = json.dumps(
rust_result.response_json or {},
ensure_ascii=False,
).encode("utf-8")
synthetic_response = httpx.Response(
status_code,
request=request,
headers=response_headers,
content=synthetic_content,
)
if status_code >= 400:
error = httpx.HTTPStatusError(
f"Upstream status error: {status_code}",
request=request,
response=synthetic_response,
)
error_body = ""
try:
if envelope and hasattr(envelope, "extract_error_text"):
error_body = await envelope.extract_error_text(synthetic_response)
else:
error_body = synthetic_response.text[:4000] if synthetic_response.text else ""
except Exception:
error_body = synthetic_response.text[:4000] if synthetic_response.text else ""
error.upstream_response = error_body[:4000] # type: ignore[attr-defined]
raise error
if upstream_is_stream:
if rust_result.response_body_bytes is None:
raise ExecutionRuntimeClientError(
"Rust executor stream sync result must contain body bytes"
)
response_json = await self._aggregate_upstream_stream_sync_response(
body_bytes=rust_result.response_body_bytes,
provider_api_format=provider_api_format,
client_api_format=client_api_format,
provider_name=str(provider.name),
provider_type=str(getattr(provider, "provider_type", "") or "").lower(),
model=str(model or ""),
request_id=str(self.request_id or ""),
envelope=envelope,
)
response_metadata_result = self._extract_response_metadata(response_json or {})
return response_json if isinstance(response_json, dict) else {}
response_json = rust_result.response_json or {}
if envelope:
response_json = envelope.unwrap_response(response_json)
envelope.postprocess_unwrapped_response(model=model, data=response_json)
response_metadata_result = self._extract_response_metadata(response_json)
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.core.context import TaskMode
exec_result = await TaskService(self.db, self.redis).execute(
task_type="cli",
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_state=request_state,
request_headers=original_headers,
request_body=original_request_body,
# 预创建失败时,回退到 TaskService 侧创建,避免丢失 pending 状态。
create_pending_usage=not pending_usage_created,
)
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 = int((time.time() - sync_start_time) * 1000)
# 确保 response_json 不为 None
if response_json is None:
response_json = {}
# 跨格式:响应转换回 client_format失败不触发 failover保守回退为原始响应
provider_response_json: dict[str, Any] | None = None
if (
needs_conversion
and provider_api_format
and api_format
and isinstance(response_json, dict)
):
try:
provider_response_json = response_json.copy()
registry = get_format_converter_registry()
response_json = registry.convert_response(
response_json,
provider_api_format,
api_format,
requested_model=model, # 使用用户请求的原始模型名
)
logger.debug(
"非流式响应格式转换完成: {} -> {}", provider_api_format, api_format
)
except Exception as conv_err:
logger.warning("非流式响应格式转换失败,使用原始响应: {}", conv_err)
provider_response_json = None
# 使用解析器提取 usage
usage = self.parser.extract_usage_from_response(response_json)
input_tokens = usage.get("input_tokens", 0)
output_tokens = usage.get("output_tokens", 0)
cached_tokens = usage.get("cache_read_tokens", 0)
cache_creation_tokens = usage.get("cache_creation_tokens", 0)
output_text = self.parser.extract_text_content(response_json)[:200]
# 非流式成功时,返回给客户端的是提供商响应头(透传)
client_response_headers = filter_proxy_response_headers(response_headers)
client_response_headers["content-type"] = "application/json"
client_response = build_json_response_for_client(
status_code=status_code,
content=response_json,
headers=client_response_headers,
client_accept_encoding=effective_client_accept_encoding,
)
actual_client_response_headers = dict(client_response.headers)
request_metadata = self._build_request_metadata() or {}
if sync_proxy_info:
request_metadata["proxy"] = sync_proxy_info
request_metadata = self._merge_scheduling_metadata(
request_metadata,
exec_result=exec_result,
selected_key_id=key_id,
)
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=original_request_body,
response_headers=response_headers,
client_response_headers=actual_client_response_headers,
response_body=provider_response_json or response_json,
client_response_body=response_json if provider_response_json else None,
provider_request_body=provider_request_body,
cache_creation_tokens=cache_creation_tokens,
cache_read_tokens=cached_tokens,
is_stream=False,
provider_request_headers=provider_request_headers,
api_format=api_format,
api_family=self.api_family,
endpoint_kind=self.endpoint_kind,
# 格式转换追踪
endpoint_api_format=provider_api_format or None,
has_format_conversion=is_format_converted(provider_api_format, str(api_format)),
# Provider 侧追踪信息(用于记录真实成本)
provider_id=provider_id,
provider_endpoint_id=endpoint_id,
provider_api_key_id=key_id,
# 模型映射信息
target_model=mapped_model_result,
# Provider 响应元数据(如 Gemini 的 modelVersion
response_metadata=response_metadata_result if response_metadata_result else None,
request_metadata=request_metadata,
)
logger.info("{} 非流式响应处理完成", self.FORMAT_ID)
# 透传提供商的响应头
return client_response
except ThinkingSignatureException as e:
# Thinking 签名错误TaskService 层已处理整流重试但仍失败
# 记录实际发送给 Provider 的请求体,便于排查问题根因
response_time_ms = int((time.time() - sync_start_time) * 1000)
request_metadata = self._build_request_metadata() or {}
if sync_proxy_info:
request_metadata["proxy"] = sync_proxy_info
request_metadata = self._merge_scheduling_metadata(
request_metadata,
selected_key_id=key_id,
pool_summary=getattr(exec_result, "pool_summary", None),
fallback_from_request=True,
)
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=original_request_body,
provider_request_body=provider_request_body,
error_message=str(e),
is_stream=False,
api_format=api_format,
api_family=self.api_family,
endpoint_kind=self.endpoint_kind,
request_metadata=request_metadata,
)
raise
except Exception as e:
response_time_ms = int((time.time() - sync_start_time) * 1000)
status_code = 503
if isinstance(e, ProviderAuthException):
status_code = 503
elif isinstance(e, ProviderRateLimitException):
status_code = 429
elif isinstance(e, ProviderTimeoutException):
status_code = 504
# 尝试从异常中提取响应头
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
request_metadata = self._merge_scheduling_metadata(
request_metadata,
selected_key_id=key_id,
pool_summary=getattr(exec_result, "pool_summary", None),
fallback_from_request=True,
)
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=original_request_body,
provider_request_body=provider_request_body,
is_stream=False,
api_format=api_format,
api_family=self.api_family,
endpoint_kind=self.endpoint_kind,
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 or None,
has_format_conversion=is_format_converted(provider_api_format, str(api_format)),
# 模型映射信息
target_model=mapped_model_result,
request_metadata=request_metadata,
)
raise
async def _extract_error_text(
self,
e: httpx.HTTPStatusError,
*,
envelope: Any = None,
) -> str:
"""从 HTTP 错误中提取错误文本"""
if envelope and hasattr(envelope, "extract_error_text"):
return await envelope.extract_error_text(e)
try:
if hasattr(e.response, "is_stream_consumed") and not e.response.is_stream_consumed:
error_bytes = await e.response.aread()
for encoding in ["utf-8", "gbk", "latin1"]:
try:
return error_bytes.decode(encoding)
except (UnicodeDecodeError, LookupError):
continue
return error_bytes.decode("utf-8", errors="replace")
else:
return (
e.response.text
if hasattr(e.response, "_content")
else "Unable to read response"
)
except Exception as decode_error:
return f"Unable to read error response: {decode_error}"

View File

@@ -0,0 +1,271 @@
"""
流式内容提取器 - 策略模式实现
为不同 API 格式OpenAI、Claude、Gemini提供内容提取和 chunk 构造的抽象。
StreamSmoother 使用这些提取器来处理不同格式的 SSE 事件。
"""
import copy
import json
from abc import ABC, abstractmethod
class ContentExtractor(ABC):
"""
流式内容提取器抽象基类
定义从 SSE 事件中提取文本内容和构造新 chunk 的接口。
每种 API 格式OpenAI、Claude、Gemini需要实现自己的提取器。
"""
@abstractmethod
def extract_content(self, data: dict) -> str | None:
"""
从 SSE 数据中提取可拆分的文本内容
Args:
data: 解析后的 JSON 数据
Returns:
提取的文本内容,如果无法提取则返回 None
"""
pass
@abstractmethod
def create_chunk(
self,
original_data: dict,
new_content: str,
event_type: str = "",
is_first: bool = False,
) -> bytes:
"""
使用新内容构造 SSE chunk
Args:
original_data: 原始 JSON 数据
new_content: 新的文本内容
event_type: SSE 事件类型(某些格式需要)
is_first: 是否是第一个 chunk用于保留 role 等字段)
Returns:
编码后的 SSE 字节数据
"""
pass
class OpenAIContentExtractor(ContentExtractor):
"""
OpenAI 格式内容提取器
处理 OpenAI Chat Completions API 的流式响应格式:
- 数据结构: choices[0].delta.content
- 只在 delta 仅包含 role/content 时允许拆分,避免破坏 tool_calls 等结构
"""
def extract_content(self, data: dict) -> str | None:
if not isinstance(data, dict):
return None
choices = data.get("choices")
if not isinstance(choices, list) or len(choices) != 1:
return None
first_choice = choices[0]
if not isinstance(first_choice, dict):
return None
delta = first_choice.get("delta")
if not isinstance(delta, dict):
return None
content = delta.get("content")
if not isinstance(content, str):
return None
# 只有 delta 仅包含 role/content 时才允许拆分
# 避免破坏 tool_calls、function_call 等复杂结构
allowed_keys = {"role", "content"}
if not all(key in allowed_keys for key in delta.keys()):
return None
return content
def create_chunk(
self,
original_data: dict,
new_content: str,
event_type: str = "",
is_first: bool = False,
) -> bytes:
new_data = original_data.copy()
if "choices" in new_data and new_data["choices"]:
new_choices = []
for choice in new_data["choices"]:
new_choice = choice.copy()
if "delta" in new_choice:
new_delta = {}
# 只有第一个 chunk 保留 role
if is_first and "role" in new_choice["delta"]:
new_delta["role"] = new_choice["delta"]["role"]
new_delta["content"] = new_content
new_choice["delta"] = new_delta
new_choices.append(new_choice)
new_data["choices"] = new_choices
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
class ClaudeContentExtractor(ContentExtractor):
"""
Claude 格式内容提取器
处理 Claude Messages API 的流式响应格式:
- 事件类型: content_block_delta
- 数据结构: delta.type=text_delta, delta.text
"""
def extract_content(self, data: dict) -> str | None:
if not isinstance(data, dict):
return None
# 检查事件类型
if data.get("type") != "content_block_delta":
return None
delta = data.get("delta", {})
if not isinstance(delta, dict):
return None
# 检查 delta 类型
if delta.get("type") != "text_delta":
return None
text = delta.get("text")
if not isinstance(text, str):
return None
return text
def create_chunk(
self,
original_data: dict,
new_content: str,
event_type: str = "",
is_first: bool = False,
) -> bytes:
new_data = original_data.copy()
if "delta" in new_data:
new_delta = new_data["delta"].copy()
new_delta["text"] = new_content
new_data["delta"] = new_delta
# Claude 格式需要 event: 前缀
event_name = event_type or "content_block_delta"
return f"event: {event_name}\ndata: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
class GeminiContentExtractor(ContentExtractor):
"""
Gemini 格式内容提取器
处理 Gemini API 的流式响应格式:
- 数据结构: candidates[0].content.parts[0].text
- 只有纯文本块才拆分
"""
def extract_content(self, data: dict) -> str | None:
if not isinstance(data, dict):
return None
candidates = data.get("candidates")
if not isinstance(candidates, list) or len(candidates) != 1:
return None
first_candidate = candidates[0]
if not isinstance(first_candidate, dict):
return None
content = first_candidate.get("content", {})
if not isinstance(content, dict):
return None
parts = content.get("parts", [])
if not isinstance(parts, list) or len(parts) != 1:
return None
first_part = parts[0]
if not isinstance(first_part, dict):
return None
text = first_part.get("text")
# 只有纯文本块(只有 text 字段)才拆分
if not isinstance(text, str) or len(first_part) != 1:
return None
return text
def create_chunk(
self,
original_data: dict,
new_content: str,
event_type: str = "",
is_first: bool = False,
) -> bytes:
new_data = copy.deepcopy(original_data)
if "candidates" in new_data and new_data["candidates"]:
first_candidate = new_data["candidates"][0]
if "content" in first_candidate:
content = first_candidate["content"]
if "parts" in content and content["parts"]:
content["parts"][0]["text"] = new_content
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
# 提取器注册表
_EXTRACTORS: dict[str, type[ContentExtractor]] = {
"openai": OpenAIContentExtractor,
"claude": ClaudeContentExtractor,
"gemini": GeminiContentExtractor,
}
def get_extractor(format_name: str) -> ContentExtractor | None:
"""
根据格式名获取对应的内容提取器实例
Args:
format_name: 格式名称openai, claude, gemini
Returns:
对应的提取器实例,如果格式不支持则返回 None
"""
extractor_class = _EXTRACTORS.get(format_name.lower())
if extractor_class:
return extractor_class()
return None
def register_extractor(format_name: str, extractor_class: type[ContentExtractor]) -> None:
"""
注册新的内容提取器
Args:
format_name: 格式名称
extractor_class: 提取器类
"""
_EXTRACTORS[format_name.lower()] = extractor_class
def get_extractor_formats() -> list[str]:
"""
获取所有已注册的格式名称列表
Returns:
格式名称列表
"""
return list(_EXTRACTORS.keys())

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,625 @@
"""
Handler Adapter 公共基类
从 ChatAdapterBase 和 CliAdapterBase 提取的共享逻辑:
- API 格式与头部处理
- 异常处理和错误响应
- 通过 `core.api_format` 注册表解析计费模板与抓模能力
- 端点测试辅助
- 路径参数合并
子类ChatAdapterBase / CliAdapterBase只需关注各自的 `handle()` 流程差异。
"""
from __future__ import annotations
import time
from typing import Any, ClassVar
import httpx
from fastapi import Request
from fastapi.responses import JSONResponse
from sqlalchemy.orm import Session
from src.api.base.adapter import ApiAdapter
from src.core.api_format import (
ApiFamily,
EndpointKind,
build_adapter_base_headers_for_endpoint,
build_adapter_headers_for_endpoint,
compute_total_input_context_for_api_format,
fetch_models_for_api_format,
get_adapter_protected_keys_for_endpoint,
get_auth_handler,
get_default_auth_method_for_endpoint,
resolve_billing_template_for_api_format,
resolve_header_name_case,
)
from src.core.exceptions import (
ProviderAuthException,
ProxyException,
UpstreamClientException,
)
from src.core.logger import logger
from src.services.billing import calculate_request_cost as _calculate_request_cost
from src.services.request.result import RequestResult
from src.services.usage.recorder import UsageRecorder
class HandlerAdapterBase(ApiAdapter):
"""
Chat/CLI Adapter 的公共基类
封装两者共享的逻辑:
- API 格式与头部处理
- 异常处理和错误响应
- 通过 `core.api_format` 注册表解析计费模板与模型抓取能力
- 端点测试辅助
子类ChatAdapterBase / CliAdapterBase只需实现 `handle()` 和格式特有的方法。
"""
# 子类必须覆盖
FORMAT_ID: str = "UNKNOWN"
# 结构化标识
API_FAMILY: ClassVar[ApiFamily | None] = None
ENDPOINT_KIND: ClassVar[EndpointKind] = EndpointKind.CHAT
# 兼容性回退:若 api_format 注册表未声明计费模板,则使用该默认值。
BILLING_TEMPLATE: str = "claude"
def __init__(self, allowed_api_formats: list[str] | None = None):
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
# =========================================================================
# API 格式与头部处理
# =========================================================================
def extract_api_key(self, request: Request) -> str | None:
"""从请求中提取 API 密钥"""
auth_method = get_default_auth_method_for_endpoint(self.FORMAT_ID)
handler = get_auth_handler(auth_method)
return handler.extract_credentials(request)
@classmethod
def build_base_headers(cls, api_key: str) -> dict[str, str]:
"""构建基础认证头"""
return build_adapter_base_headers_for_endpoint(cls.FORMAT_ID, api_key)
@classmethod
def build_headers_with_extra(
cls, api_key: str, extra_headers: dict[str, str] | None = None
) -> dict[str, str]:
"""构建带额外头部的完整请求头"""
return build_adapter_headers_for_endpoint(cls.FORMAT_ID, api_key, extra_headers)
@classmethod
def get_protected_header_keys(cls) -> tuple[str, ...]:
"""返回不应被 extra_headers 覆盖的头部 key"""
return get_adapter_protected_keys_for_endpoint(cls.FORMAT_ID)
# =========================================================================
# 路径参数合并
# =========================================================================
def _merge_path_params(
self, original_request_body: dict[str, Any], path_params: dict[str, Any]
) -> dict[str, Any]:
"""合并 URL 路径参数到请求体 - 子类可覆盖"""
merged = original_request_body.copy()
for key, value in path_params.items():
if key not in merged:
merged[key] = value
return merged
# =========================================================================
# 异常处理
# =========================================================================
async def _handle_provider_exception(
self,
e: Exception,
*,
db: Session,
user: Any,
api_key: Any,
model: str,
stream: bool,
start_time: float,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
client_ip: str,
request_id: str,
) -> JSONResponse:
"""处理 Provider 相关异常"""
logger.debug("Caught provider exception: {}", type(e).__name__)
response_time = int((time.time() - start_time) * 1000)
result = RequestResult.from_exception(
exception=e,
api_format=self.FORMAT_ID,
model=model,
response_time_ms=response_time,
is_stream=stream,
)
result.request_headers = original_headers
result.request_body = original_request_body
if isinstance(e, ProviderAuthException):
error_message = (
"上游服务认证失败" if result.metadata.provider != "unknown" else "服务暂时不可用"
)
result.error_message = error_message
if isinstance(e, UpstreamClientException):
result.status_code = e.status_code
result.error_message = e.message
recorder = UsageRecorder(
db=db,
user=user,
api_key=api_key,
client_ip=client_ip,
request_id=request_id,
)
await recorder.record_failure(result, original_headers, original_request_body)
if isinstance(e, UpstreamClientException):
error_type = "invalid_request_error"
elif result.status_code == 503:
error_type = "internal_server_error"
else:
error_type = "rate_limit_exceeded"
return self._error_response(
status_code=result.status_code,
error_type=error_type,
message=result.error_message or str(e),
)
async def _handle_unexpected_exception(
self,
e: Exception,
*,
db: Session,
user: Any,
api_key: Any,
model: str,
stream: bool,
start_time: float,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
client_ip: str,
request_id: str,
) -> JSONResponse:
"""处理未预期的异常"""
if isinstance(e, ProxyException):
logger.error("{} 请求处理业务异常: {}: {}", self.FORMAT_ID, type(e).__name__, e)
else:
logger.opt(exception=e).error(
"{} 请求处理意外异常: {}: {}", self.FORMAT_ID, type(e).__name__, e
)
response_time = int((time.time() - start_time) * 1000)
result = RequestResult.from_exception(
exception=e,
api_format=self.FORMAT_ID,
model=model,
response_time_ms=response_time,
is_stream=stream,
)
result.status_code = 500
result.error_type = "internal_error"
result.request_headers = original_headers
result.request_body = original_request_body
try:
recorder = UsageRecorder(
db=db,
user=user,
api_key=api_key,
client_ip=client_ip,
request_id=request_id,
)
await recorder.record_failure(result, original_headers, original_request_body)
except Exception as record_error:
logger.error("记录失败请求时出错: {}", record_error)
return self._error_response(
status_code=500, error_type="internal_server_error", message="处理请求时发生内部错误"
)
def _error_response(self, status_code: int, error_type: str, message: str) -> JSONResponse:
"""生成错误响应 - 子类可覆盖以自定义格式"""
return JSONResponse(
status_code=status_code,
content={
"error": {
"type": error_type,
"message": message,
}
},
)
# =========================================================================
# 计费能力委托
# =========================================================================
def compute_cost(
self,
input_tokens: int,
output_tokens: int,
cache_creation_input_tokens: int,
cache_read_input_tokens: int,
input_price_per_1m: float,
output_price_per_1m: float,
cache_creation_price_per_1m: float | None,
cache_read_price_per_1m: float | None,
price_per_request: float | None,
tiered_pricing: dict | None = None,
cache_ttl_minutes: int | None = None,
) -> dict[str, Any]:
"""计算请求成本"""
total_input_context = compute_total_input_context_for_api_format(
self.FORMAT_ID, input_tokens, cache_read_input_tokens, cache_creation_input_tokens
)
billing_template = (
resolve_billing_template_for_api_format(self.FORMAT_ID) or self.BILLING_TEMPLATE
)
return _calculate_request_cost(
input_tokens=input_tokens,
output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens,
input_price_per_1m=input_price_per_1m,
output_price_per_1m=output_price_per_1m,
cache_creation_price_per_1m=cache_creation_price_per_1m,
cache_read_price_per_1m=cache_read_price_per_1m,
price_per_request=price_per_request,
tiered_pricing=tiered_pricing,
cache_ttl_minutes=cache_ttl_minutes,
total_input_context=total_input_context,
billing_template=billing_template,
)
# =========================================================================
# 模型抓取委托与端点测试
# =========================================================================
@classmethod
async def fetch_models(
cls,
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询上游 API 支持的模型列表。"""
return await fetch_models_for_api_format(
client,
api_format=cls.FORMAT_ID,
base_url=base_url,
api_key=api_key,
extra_headers=extra_headers,
)
@classmethod
def build_request_body(
cls,
request_data: dict[str, Any] | None = None,
*,
base_url: str | None = None,
provider_type: str | None = None,
) -> dict[str, Any]:
"""构建测试请求体,使用转换器注册表自动处理格式转换"""
from src.api.handlers.base.request_builder import build_test_request_body
_ = base_url, provider_type
return build_test_request_body(cls.FORMAT_ID, request_data)
@staticmethod
def _validate_test_base_url(base_url: Any) -> str:
"""校验 test-model 场景传入的 base_url。"""
if not isinstance(base_url, str):
raise TypeError(f"base_url must be a non-empty string, got {type(base_url).__name__}")
normalized = base_url.strip()
if not normalized:
raise ValueError("base_url must be a non-empty string")
return normalized
@classmethod
async def check_endpoint(
cls,
client: httpx.AsyncClient,
base_url: str,
api_key: str,
request_data: dict[str, Any],
extra_headers: dict[str, str] | None = None,
# 端点规则参数
body_rules: list[dict[str, Any]] | None = None,
header_rules: list[dict[str, Any]] | None = None,
# 用量计算参数
db: Any | None = None,
user: Any | None = None,
provider_name: str | None = None,
provider_id: str | None = None,
api_key_id: str | None = None,
model_name: str | None = None,
# Provider 上下文(用于 OAuth 认证和特殊路由)
auth_type: str | None = None,
provider_type: str | None = None,
decrypted_auth_config: dict[str, Any] | None = None,
provider_endpoint: Any | None = None,
provider_api_key: Any | None = None,
# 代理配置
proxy_config: dict[str, Any] | None = None,
timeout_seconds: float | None = None,
) -> dict[str, Any]:
"""
测试模型连接性(非流式)
统一的 endpoint 测试方法,支持 OAuth/Antigravity/Kiro 等特殊路由。
"""
from src.api.handlers.base.endpoint_checker import run_endpoint_check
from src.api.handlers.base.request_builder import (
apply_body_rules,
evaluate_condition,
)
from src.core.api_format.headers import HeaderBuilder
from src.core.provider_types import ProviderType
from src.services.provider.adapters.vertex_ai.transport import is_vertex_ai_context
validated_base_url = cls._validate_test_base_url(base_url)
is_antigravity = provider_type == ProviderType.ANTIGRAVITY
is_gemini_cli = provider_type == ProviderType.GEMINI_CLI
is_vertex = is_vertex_ai_context(
base_url=validated_base_url,
provider_type=provider_type,
endpoint=provider_endpoint,
key=provider_api_key,
)
is_kiro = provider_type == ProviderType.KIRO
is_oauth = auth_type == "oauth"
vertex_auth_info: Any | None = None
kiro_cfg: Any | None = None
if is_kiro:
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
kiro_cfg = KiroAuthConfig.from_dict(decrypted_auth_config or {})
# ---- URL ----
if is_kiro:
from src.services.provider.adapters.kiro.request import (
build_kiro_generate_assistant_url,
)
assert kiro_cfg is not None
url = build_kiro_generate_assistant_url(validated_base_url, cfg=kiro_cfg)
elif is_antigravity:
from src.services.provider.adapters.antigravity.constants import (
V1INTERNAL_PATH_TEMPLATE,
get_v1internal_extra_headers,
)
from src.services.provider.adapters.antigravity.url_availability import (
url_availability,
)
ordered_urls = url_availability.get_ordered_urls(prefer_daily=True)
effective_base_url = ordered_urls[0] if ordered_urls else validated_base_url
path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent")
url = f"{str(effective_base_url).rstrip('/')}{path}"
elif is_gemini_cli:
from src.services.provider.adapters.gemini_cli.constants import V1INTERNAL_PATH_TEMPLATE
effective_base_url = validated_base_url
path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent")
url = f"{str(effective_base_url).rstrip('/')}{path}"
elif is_vertex and provider_endpoint is not None and provider_api_key is not None:
from src.services.provider.auth import get_provider_auth
from src.services.provider.transport import build_provider_url
vertex_auth_info = await get_provider_auth(provider_endpoint, provider_api_key)
effective_auth_config = (
vertex_auth_info.decrypted_auth_config
if vertex_auth_info
else decrypted_auth_config
)
if effective_auth_config:
decrypted_auth_config = effective_auth_config
effective_model_name = model_name or request_data.get("model", "")
path_params = {"model": effective_model_name} if effective_model_name else None
url = build_provider_url(
provider_endpoint,
path_params=path_params,
is_stream=bool(request_data.get("stream", False)),
key=provider_api_key,
decrypted_auth_config=effective_auth_config,
)
else:
url = cls.build_endpoint_url(
validated_base_url,
request_data,
model_name,
provider_type=provider_type,
)
# ---- Headers ----
cli_extra = cls.get_cli_extra_headers(
base_url=validated_base_url,
provider_type=provider_type,
)
merged_extra = dict(extra_headers) if extra_headers else {}
merged_extra.update(cli_extra)
if is_antigravity:
merged_extra.update(get_v1internal_extra_headers())
elif is_gemini_cli:
from src.services.provider.adapters.gemini_cli.constants import (
get_v1internal_extra_headers,
)
merged_extra.update(get_v1internal_extra_headers())
if is_kiro:
from src.services.provider.adapters.kiro.request import (
build_kiro_request_headers,
)
assert kiro_cfg is not None
kiro_headers = build_kiro_request_headers(
kiro_cfg,
access_token=api_key,
)
merged_extra.update(kiro_headers)
if is_vertex and provider_endpoint is not None and provider_api_key is not None:
headers = dict(merged_extra)
if (
vertex_auth_info
and getattr(vertex_auth_info, "auth_header", None)
and getattr(vertex_auth_info, "auth_value", None)
):
headers[str(vertex_auth_info.auth_header)] = str(vertex_auth_info.auth_value)
else:
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None)
if is_oauth:
from src.core.api_format import get_auth_config_for_endpoint
default_auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
if default_auth_header.lower() != "authorization":
headers.pop(default_auth_header, None)
auth_header_name = resolve_header_name_case(extra_headers, "Authorization")
headers[auth_header_name] = f"Bearer {api_key}"
# ---- Body ----
body = cls.build_request_body(
request_data,
base_url=validated_base_url,
provider_type=provider_type,
)
if body_rules:
body = apply_body_rules(
body,
body_rules,
original_body=body,
)
if is_antigravity:
from src.services.provider.adapters.antigravity.envelope import (
wrap_v1internal_request,
)
project_id = (decrypted_auth_config or {}).get("project_id", "")
effective_model = model_name or request_data.get("model", "")
body = wrap_v1internal_request(
body,
project_id=project_id,
model=effective_model,
)
elif is_gemini_cli:
from src.services.provider.adapters.gemini_cli.envelope import wrap_v1internal_request
project_id = (decrypted_auth_config or {}).get("project_id", "")
effective_model = model_name or request_data.get("model", "")
body = wrap_v1internal_request(
body,
project_id=project_id,
model=effective_model,
)
if is_kiro:
from src.services.provider.adapters.kiro.request import (
build_kiro_request_payload,
)
assert kiro_cfg is not None
effective_model = model_name or request_data.get("model", "")
body = build_kiro_request_payload(
body,
model=effective_model,
cfg=kiro_cfg,
)
# ---- Header Rules ----
if header_rules:
from src.core.api_format import get_auth_config_for_endpoint as _get_auth_cfg
if is_oauth:
protected_keys = {"authorization", "content-type"}
else:
ep_auth_header, _ = _get_auth_cfg(cls.FORMAT_ID)
protected_keys = {ep_auth_header.lower(), "content-type"}
if vertex_auth_info and getattr(vertex_auth_info, "auth_header", None):
protected_keys.add(str(vertex_auth_info.auth_header).lower())
header_builder = HeaderBuilder()
header_builder.add_many(headers)
header_builder.apply_rules(
header_rules,
protected_keys,
body=body,
original_body=body,
condition_evaluator=evaluate_condition,
)
headers = header_builder.build()
# ---- Execute ----
effective_model_name = model_name or request_data.get("model")
return await run_endpoint_check(
client=client,
url=url,
headers=headers,
json_body=body,
api_format=cls.FORMAT_ID,
is_stream=bool(request_data.get("stream", False)),
db=db,
user=user,
provider_name=provider_name,
provider_id=provider_id,
api_key_id=api_key_id,
model_name=effective_model_name,
proxy_config=proxy_config,
timeout=timeout_seconds,
)
# =========================================================================
# CLI Adapter 配置方法 - 子类可覆盖
# =========================================================================
@classmethod
def build_endpoint_url(
cls,
base_url: str,
request_data: dict[str, Any] | None = None,
model_name: str | None = None,
*,
provider_type: str | None = None,
) -> str:
"""构建 API 端点 URL - 子类应覆盖"""
return base_url
@classmethod
def get_cli_user_agent(cls) -> str | None:
"""获取 CLI User-Agent - 子类可覆盖"""
return None
@classmethod
def get_cli_extra_headers(
cls, *, base_url: str | None = None, provider_type: str | None = None
) -> dict[str, str]:
"""获取额外请求头 - 子类可覆盖"""
headers: dict[str, str] = {}
cli_user_agent = cls.get_cli_user_agent()
if cli_user_agent:
headers["User-Agent"] = cli_user_agent
return headers

View File

@@ -0,0 +1,717 @@
"""
响应解析器工厂
直接根据格式 ID 创建对应的 ResponseParser 实现,
不再经过 Protocol 抽象层。
"""
import re
from typing import Any
from src.api.handlers.base.response_parser import (
ParsedChunk,
ParsedResponse,
ResponseParser,
StreamStats,
)
# is_cli_format 权威定义在 core 层
from src.core.api_format import is_cli_format
from src.core.usage_tokens import extract_cache_creation_tokens, extract_cache_read_tokens
def _check_nested_error(response: dict[str, Any]) -> tuple[bool, dict[str, Any] | None]:
"""
检查响应中是否存在嵌套错误(某些代理服务返回 HTTP 200 但在响应体中包含错误)
检测格式:
1. 顶层 error: {"error": {...}}
2. 顶层 type=error: {"type": "error", ...}
3. chunks 内嵌套 error: {"chunks": [{"error": {...}}]}
Args:
response: 响应字典
Returns:
(is_error, error_dict): 是否为错误,以及提取的错误信息
"""
# 顶层 error
if "error" in response:
error = response["error"]
if isinstance(error, dict):
return True, error
return True, {"message": str(error)}
# 顶层 type=error
if response.get("type") == "error":
return True, response
# chunks 内嵌套 error (某些代理返回这种格式)
chunks = response.get("chunks", [])
if chunks and isinstance(chunks, list):
for chunk in chunks:
if isinstance(chunk, dict):
if "error" in chunk:
error = chunk["error"]
if isinstance(error, dict):
return True, error
return True, {"message": str(error)}
if chunk.get("type") == "error":
return True, chunk
return False, None
def _extract_embedded_status_code(error_info: dict[str, Any] | None) -> int | None:
"""
从错误信息中提取嵌套的状态码
支持多种格式:
1. 直接的 code 字段: {"code": 400}
2. status 字段: {"status": 400}
3. 从 message 中正则提取: "Request failed with status code 400"
4. 从 type 字段映射: "invalid_request_error" -> 400
Args:
error_info: 错误信息字典
Returns:
提取的状态码,如果无法提取则返回 None
"""
if not error_info:
return None
# 1. 直接的 code 字段Gemini 等)
code = error_info.get("code")
if isinstance(code, int) and 100 <= code < 600:
return code
if isinstance(code, str) and code.isdigit():
code_int = int(code)
if 100 <= code_int < 600:
return code_int
# 2. status 字段
status = error_info.get("status")
if isinstance(status, int) and 100 <= status < 600:
return status
if isinstance(status, str) and status.isdigit():
status_int = int(status)
if 100 <= status_int < 600:
return status_int
# 3. 从 message 中正则提取 (例如 "Request failed with status code 400")
message = error_info.get("message", "")
if message:
# 匹配 "status code XXX" 或 "status XXX" 或 "HTTP XXX"
match = re.search(r"(?:status\s*(?:code\s*)?|HTTP\s*)(\d{3})", message, re.IGNORECASE)
if match:
code_int = int(match.group(1))
if 100 <= code_int < 600:
return code_int
# 4. 从 type 字段映射常见的错误类型
error_type = error_info.get("type", "")
type_to_status = {
"invalid_request_error": 400,
"authentication_error": 401,
"permission_error": 403,
"not_found_error": 404,
"rate_limit_error": 429,
"overloaded_error": 503,
"api_error": 500,
"internal_error": 500,
}
if error_type and error_type.lower() in type_to_status:
return type_to_status[error_type.lower()]
return None
class OpenAIResponseParser(ResponseParser):
"""OpenAI 格式响应解析器"""
API_FORMAT = "openai:chat"
def __init__(self) -> None:
from src.api.handlers.openai.stream_parser import OpenAIStreamParser
self._parser = OpenAIStreamParser()
self.name = self.API_FORMAT
self.api_format = self.API_FORMAT
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
if not line or not line.strip():
return None
if line.startswith("data: "):
data_str = line[6:]
else:
data_str = line
parsed = self._parser.parse_line(data_str)
if parsed is None:
return None
chunk = ParsedChunk(
raw_line=line,
event_type=None,
data=parsed,
)
# 提取文本增量
text_delta = self._parser.extract_text_delta(parsed)
if text_delta:
chunk.text_delta = text_delta
stats.collected_text += text_delta
# 检查是否结束
if self._parser.is_done_chunk(parsed):
chunk.is_done = True
stats.has_completion = True
# 提取 usage 信息(某些 OpenAI 兼容 API 如豆包会在最后一个 chunk 中发送 usage
# 这个 chunk 通常 choices 为空数组,但包含完整的 usage 信息
# 使用取最大值策略确保正确统计
usage = parsed.get("usage")
if usage and isinstance(usage, dict):
chunk.input_tokens = usage.get("prompt_tokens", 0)
chunk.output_tokens = usage.get("completion_tokens", 0)
# 取最大值更新 stats
if chunk.input_tokens > stats.input_tokens:
stats.input_tokens = chunk.input_tokens
if chunk.output_tokens > stats.output_tokens:
stats.output_tokens = chunk.output_tokens
stats.chunk_count += 1
stats.data_count += 1
return chunk
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
result = ParsedResponse(
raw_response=response,
status_code=status_code,
)
# 提取文本内容
choices = response.get("choices", [])
if choices:
message = choices[0].get("message", {})
content = message.get("content")
if content:
result.text_content = content
result.response_id = response.get("id")
# 提取 usage
usage = response.get("usage") or {}
result.input_tokens = usage.get("prompt_tokens", 0)
result.output_tokens = usage.get("completion_tokens", 0)
result.cache_read_tokens = extract_cache_read_tokens(usage)
# 检查错误(支持嵌套错误格式)
is_error, error_info = _check_nested_error(response)
if is_error and error_info:
result.is_error = True
result.error_type = error_info.get("type")
result.error_message = error_info.get("message")
result.embedded_status_code = _extract_embedded_status_code(error_info)
return result
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
usage = response.get("usage") or {}
return {
"input_tokens": usage.get("prompt_tokens", 0),
"output_tokens": usage.get("completion_tokens", 0),
"cache_creation_tokens": 0,
"cache_read_tokens": extract_cache_read_tokens(usage),
}
def extract_text_content(self, response: dict[str, Any]) -> str:
choices = response.get("choices", [])
if choices:
message = choices[0].get("message", {})
content = message.get("content")
if isinstance(content, str):
return content
return ""
def is_error_response(self, response: dict[str, Any]) -> bool:
is_error, _ = _check_nested_error(response)
return is_error
class OpenAICliResponseParser(OpenAIResponseParser):
"""OpenAI CLI / Responses API 格式响应解析器
OpenAI Responses API 与 Chat Completions API 的关键差异:
- Usage 字段: input_tokens/output_tokens而非 prompt_tokens/completion_tokens
- 响应结构: output[].content[].text而非 choices[].message.content
- 流式事件: response.completed 事件中 usage 嵌套在 response 对象内
"""
API_FORMAT = "openai:cli"
def __init__(self) -> None:
super().__init__()
self.name = self.API_FORMAT
self.api_format = self.API_FORMAT
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
result = ParsedResponse(
raw_response=response,
status_code=status_code,
)
# Responses API: 文本在 output[].content[].text 中
result.text_content = self._extract_responses_api_text(response)
result.response_id = response.get("id")
# Responses API usage: input_tokens / output_tokens
usage = self._extract_responses_api_usage(response)
result.input_tokens = usage.get("input_tokens", 0)
result.output_tokens = usage.get("output_tokens", 0)
result.cache_creation_tokens = usage.get("cache_creation_tokens", 0)
result.cache_read_tokens = usage.get("cache_read_tokens", 0)
# 检查错误(支持嵌套错误格式)
is_error, error_info = _check_nested_error(response)
if is_error and error_info:
result.is_error = True
result.error_type = error_info.get("type")
result.error_message = error_info.get("message")
result.embedded_status_code = _extract_embedded_status_code(error_info)
return result
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
usage = self._extract_responses_api_usage(response)
return usage
def extract_text_content(self, response: dict[str, Any]) -> str:
return self._extract_responses_api_text(response)
@staticmethod
def _extract_responses_api_usage(response: dict[str, Any]) -> dict[str, int]:
"""从 Responses API 响应或流式事件中提取 usage
支持多种结构:
1. 顶层 usage非流式响应 / 部分转换后的响应)
2. response.usage流式 response.completed 事件)
3. 兼容 Chat Completions 字段名prompt_tokens/completion_tokens
"""
usage: dict[str, Any] = {}
# 优先从顶层 usage 提取
top_usage = response.get("usage")
if isinstance(top_usage, dict):
usage = top_usage
else:
# 流式事件: response.completed 中 usage 嵌套在 response 对象内
resp_obj = response.get("response")
if isinstance(resp_obj, dict):
nested_usage = resp_obj.get("usage")
if isinstance(nested_usage, dict):
usage = nested_usage
if not usage:
return {
"input_tokens": 0,
"output_tokens": 0,
"cache_creation_tokens": 0,
"cache_read_tokens": 0,
}
# Responses API 使用 input_tokens/output_tokens
# 兼容 Chat Completions 的 prompt_tokens/completion_tokens以防转换后的响应
input_tokens = usage.get("input_tokens") or usage.get("prompt_tokens") or 0
output_tokens = usage.get("output_tokens") or usage.get("completion_tokens") or 0
return {
"input_tokens": int(input_tokens),
"output_tokens": int(output_tokens),
"cache_creation_tokens": extract_cache_creation_tokens(usage),
"cache_read_tokens": extract_cache_read_tokens(usage),
}
@staticmethod
def _extract_responses_api_text(response: dict[str, Any]) -> str:
"""从 Responses API 响应中提取文本内容
支持结构: output[].content[].text 或 output[].text
"""
text_parts: list[str] = []
output = response.get("output")
if isinstance(output, list):
for item in output:
if not isinstance(item, dict):
continue
# message 类型: output[].content[].text
if item.get("type") == "message":
content = item.get("content")
if isinstance(content, list):
for part in content:
if isinstance(part, dict):
ptype = str(part.get("type") or "")
if ptype in ("output_text", "text") and isinstance(
part.get("text"), str
):
text_parts.append(part["text"])
# 直接文本类型: output[].text
elif item.get("type") in ("output_text", "text") and isinstance(
item.get("text"), str
):
text_parts.append(item["text"])
# 兼容: 部分实现可能直接给 output_text
if not text_parts and isinstance(response.get("output_text"), str):
text_parts.append(response["output_text"])
# 兼容: 如果是 Chat Completions 格式(可能来自转换后的响应),回退到 choices 结构
if not text_parts:
choices = response.get("choices", [])
if isinstance(choices, list) and choices:
message = choices[0].get("message", {}) if isinstance(choices[0], dict) else {}
content = message.get("content") if isinstance(message, dict) else None
if isinstance(content, str):
text_parts.append(content)
return "".join(text_parts)
class ClaudeResponseParser(ResponseParser):
"""Claude 格式响应解析器"""
API_FORMAT = "claude:chat"
def __init__(self) -> None:
from src.api.handlers.claude.stream_parser import ClaudeStreamParser
self._parser = ClaudeStreamParser()
self.name = self.API_FORMAT
self.api_format = self.API_FORMAT
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
if not line or not line.strip():
return None
if line.startswith("data: "):
data_str = line[6:]
else:
data_str = line
parsed = self._parser.parse_line(data_str)
if parsed is None:
return None
chunk = ParsedChunk(
raw_line=line,
event_type=self._parser.get_event_type(parsed),
data=parsed,
)
# 提取文本增量
text_delta = self._parser.extract_text_delta(parsed)
if text_delta:
chunk.text_delta = text_delta
stats.collected_text += text_delta
# 检查是否结束
if self._parser.is_done_event(parsed):
chunk.is_done = True
stats.has_completion = True
# 提取 usage
# Claude 流式响应的 usage 可能在首个 chunkmessage_start或最后一个 chunkmessage_delta
# 首个 chunk 通常包含 input_tokens最后一个 chunk 包含 output_tokens
# 使用取最大值策略确保正确统计
usage = self._parser.extract_usage(parsed)
if usage:
chunk.input_tokens = usage.get("input_tokens", 0)
chunk.output_tokens = usage.get("output_tokens", 0)
chunk.cache_creation_tokens = usage.get("cache_creation_tokens", 0)
chunk.cache_read_tokens = usage.get("cache_read_tokens", 0)
# 取最大值更新 stats
if chunk.input_tokens > stats.input_tokens:
stats.input_tokens = chunk.input_tokens
if chunk.output_tokens > stats.output_tokens:
stats.output_tokens = chunk.output_tokens
if chunk.cache_creation_tokens > stats.cache_creation_tokens:
stats.cache_creation_tokens = chunk.cache_creation_tokens
if chunk.cache_read_tokens > stats.cache_read_tokens:
stats.cache_read_tokens = chunk.cache_read_tokens
# 检查错误
if self._parser.is_error_event(parsed):
chunk.is_error = True
error = parsed.get("error", {})
if isinstance(error, dict):
chunk.error_message = error.get("message", str(error))
else:
chunk.error_message = str(error)
stats.chunk_count += 1
stats.data_count += 1
return chunk
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
result = ParsedResponse(
raw_response=response,
status_code=status_code,
)
# 提取文本内容
content = response.get("content", [])
if isinstance(content, list):
text_parts = []
for block in content:
if isinstance(block, dict) and block.get("type") == "text":
text_parts.append(block.get("text", ""))
result.text_content = "".join(text_parts)
result.response_id = response.get("id")
# 提取 usage
usage = response.get("usage") or {}
result.input_tokens = usage.get("input_tokens", 0)
result.output_tokens = usage.get("output_tokens", 0)
result.cache_creation_tokens = extract_cache_creation_tokens(usage)
result.cache_read_tokens = usage.get("cache_read_input_tokens", 0)
# 检查错误(支持嵌套错误格式)
is_error, error_info = _check_nested_error(response)
if is_error and error_info:
result.is_error = True
result.error_type = error_info.get("type")
result.error_message = error_info.get("message")
result.embedded_status_code = _extract_embedded_status_code(error_info)
return result
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
# 对于 message_start 事件usage 在 message.usage 路径下
# 对于其他响应usage 在顶层
usage = response.get("usage") or {}
if not usage and "message" in response:
usage = (response.get("message") or {}).get("usage") or {}
return {
"input_tokens": usage.get("input_tokens", 0),
"output_tokens": usage.get("output_tokens", 0),
"cache_creation_tokens": extract_cache_creation_tokens(usage),
"cache_read_tokens": usage.get("cache_read_input_tokens", 0),
}
def extract_text_content(self, response: dict[str, Any]) -> str:
content = response.get("content", [])
if isinstance(content, list):
text_parts = []
for block in content:
if isinstance(block, dict) and block.get("type") == "text":
text_parts.append(block.get("text", ""))
return "".join(text_parts)
return ""
def is_error_response(self, response: dict[str, Any]) -> bool:
is_error, _ = _check_nested_error(response)
return is_error
class GeminiResponseParser(ResponseParser):
"""Gemini 格式响应解析器"""
API_FORMAT = "gemini:chat"
def __init__(self) -> None:
from src.api.handlers.gemini.stream_parser import GeminiStreamParser
self._parser = GeminiStreamParser()
self.name = self.API_FORMAT
self.api_format = self.API_FORMAT
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
"""
解析 Gemini SSE 行
Gemini 的流式响应使用 SSE 格式 (data: {...})
"""
if not line or not line.strip():
return None
# Gemini SSE 格式: data: {...}
if line.startswith("data: "):
data_str = line[6:]
else:
data_str = line
parsed = self._parser.parse_line(data_str)
if parsed is None:
return None
chunk = ParsedChunk(
raw_line=line,
event_type="content",
data=parsed,
)
# 提取文本增量
text_delta = self._parser.extract_text_delta(parsed)
if text_delta:
chunk.text_delta = text_delta
stats.collected_text += text_delta
# 检查是否结束
if self._parser.is_done_event(parsed):
chunk.is_done = True
stats.has_completion = True
# 提取 usage
# Gemini 流式响应的 usage 可能出现在多个 chunk 中
# 使用取最大值策略确保正确统计
usage = self._parser.extract_usage(parsed)
if usage:
chunk.input_tokens = usage.get("input_tokens", 0)
chunk.output_tokens = usage.get("output_tokens", 0)
chunk.cache_read_tokens = usage.get("cached_tokens", 0)
# 取最大值更新 stats
if chunk.input_tokens > stats.input_tokens:
stats.input_tokens = chunk.input_tokens
if chunk.output_tokens > stats.output_tokens:
stats.output_tokens = chunk.output_tokens
if chunk.cache_read_tokens > stats.cache_read_tokens:
stats.cache_read_tokens = chunk.cache_read_tokens
# 检查错误
if self._parser.is_error_event(parsed):
chunk.is_error = True
error = parsed.get("error", {})
if isinstance(error, dict):
chunk.error_message = error.get("message", str(error))
else:
chunk.error_message = str(error)
stats.chunk_count += 1
stats.data_count += 1
return chunk
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
result = ParsedResponse(
raw_response=response,
status_code=status_code,
)
# 提取文本内容
candidates = response.get("candidates", [])
if candidates:
content = candidates[0].get("content", {})
parts = content.get("parts", [])
text_parts = []
for part in parts:
if "text" in part:
text_parts.append(part["text"])
result.text_content = "".join(text_parts)
result.response_id = response.get("modelVersion")
# 提取 usage调用 GeminiStreamParser.extract_usage 作为单一实现源)
usage = self._parser.extract_usage(response)
if usage:
result.input_tokens = usage.get("input_tokens", 0)
result.output_tokens = usage.get("output_tokens", 0)
result.cache_read_tokens = usage.get("cached_tokens", 0)
# 检查错误(使用增强的错误检测)
error_info = self._parser.extract_error_info(response)
if error_info:
result.is_error = True
result.error_type = error_info.get("status")
result.error_message = error_info.get("message")
result.embedded_status_code = _extract_embedded_status_code(error_info)
return result
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
"""
从 Gemini 响应中提取 token 使用量
调用 GeminiStreamParser.extract_usage 作为单一实现源
"""
usage = self._parser.extract_usage(response)
if not usage:
return {
"input_tokens": 0,
"output_tokens": 0,
"cache_creation_tokens": 0,
"cache_read_tokens": 0,
}
return {
"input_tokens": usage.get("input_tokens", 0),
"output_tokens": usage.get("output_tokens", 0),
"cache_creation_tokens": 0,
"cache_read_tokens": usage.get("cached_tokens", 0),
}
def extract_text_content(self, response: dict[str, Any]) -> str:
candidates = response.get("candidates", [])
if candidates:
content = candidates[0].get("content", {})
parts = content.get("parts", [])
text_parts = []
for part in parts:
if "text" in part:
text_parts.append(part["text"])
return "".join(text_parts)
return ""
def is_error_response(self, response: dict[str, Any]) -> bool:
"""
判断响应是否为错误响应
使用增强的错误检测逻辑,支持嵌套在 chunks 中的错误
"""
return bool(self._parser.is_error_event(response))
# 注册解析器到 core 层注册表(供 services 层通过 format_id 获取)
from src.core.stream_types import get_parser_for_format, register_parser
def register_default_parsers() -> None:
"""自动发现所有 ResponseParser 子类并注册
通过 __subclasses__() 递归收集所有 ResponseParser 子类,
使用类级别 API_FORMAT 属性获取格式 ID无需实例化。
"""
def _collect_subclasses(base: type) -> list[type]:
subs = base.__subclasses__()
return subs + [s for c in subs for s in _collect_subclasses(c)]
for cls in _collect_subclasses(ResponseParser):
api_format = getattr(cls, "API_FORMAT", None)
if api_format:
register_parser(api_format, cls)
# 模块加载时自动注册(保证 import parsers 即可用,测试也不需要手动初始化)
# main.py lifespan 中的显式调用是冗余但无害的安全保障dict 覆盖幂等)
register_default_parsers()
__all__ = [
"OpenAIResponseParser",
"OpenAICliResponseParser",
"ClaudeResponseParser",
"GeminiResponseParser",
"register_default_parsers",
"get_parser_for_format",
"is_cli_format",
]

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,19 @@
"""
响应解析器基类 - re-export from src.core.stream_types
实际定义已下沉到 src/core/stream_types.py此文件保留向后兼容的 re-export。
"""
from src.core.stream_types import (
ParsedChunk,
ParsedResponse,
ResponseParser,
StreamStats,
)
__all__ = [
"ParsedChunk",
"ParsedResponse",
"ResponseParser",
"StreamStats",
]

View File

@@ -0,0 +1,473 @@
"""
流式处理上下文 - 类型安全的数据类替代 dict
提供流式请求处理过程中的状态跟踪,包括:
- Provider/Endpoint/Key 信息
- Token 统计
- 响应状态
- 请求/响应数据
"""
from __future__ import annotations
import json
import time
from contextlib import contextmanager
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Iterator
if TYPE_CHECKING:
from src.core.api_format.conversion.stream_state import StreamState
def extract_proxy_timing(proxy_info: dict[str, Any] | None, headers: dict[str, str]) -> None:
"""从响应头中提取代理分阶段耗时X-Proxy-Timing写入 proxy_info"""
if proxy_info is None:
return
timing_raw = headers.get("x-proxy-timing")
if not timing_raw:
return
try:
timing = json.loads(timing_raw)
if isinstance(timing, dict):
proxy_info["timing"] = timing
except (json.JSONDecodeError, TypeError):
pass
def is_format_converted(
provider_api_format: str | None,
client_api_format: str | None,
) -> bool:
"""client 与 provider 的 api_format 是否真正不同(用于 usage 展示层)"""
return bool(
provider_api_format
and client_api_format
and provider_api_format.strip().lower() != client_api_format.strip().lower()
)
_MAX_COLLECTED_TEXT_CHARS = 16 * 1024
@dataclass
class RecordedStreamBodies:
"""统一封装 telemetry/usage 使用的流式响应体引用。"""
response_body: dict[str, Any] | None
client_response_body: dict[str, Any] | None
def ensure_populated(self, ctx: StreamContext, response_time_ms: int) -> None:
"""在 fallback 到需要 body 的路径时按需补建响应体。"""
if self.response_body is None:
self.response_body = ctx.build_response_body(response_time_ms)
if self.client_response_body is None:
self.client_response_body = ctx.build_client_response_body(response_time_ms)
@dataclass
class StreamContext:
"""
流式处理上下文
用于在流式请求处理过程中跟踪状态,替代原有的 ctx dict。
所有字段都有类型注解,提供更好的 IDE 支持和运行时类型安全。
"""
# 请求基本信息
model: str
api_format: str
api_family: str | None = None # 协议族(从 Adapter 层透传)
endpoint_kind: str | None = None # 端点类型(从 Adapter 层透传)
# 请求标识信息CLI handler 需要)
request_id: str = ""
user_id: int = 0
api_key_id: int = 0
# Provider 信息(在请求执行时填充)
provider_name: str | None = None
provider_id: str | None = None
provider_type: str | None = None # Provider 类型(如 codex用于元数据采集
# Transport 层选中的 base_url用于 URL 可用性更新/故障转移等场景)
selected_base_url: str | None = None
endpoint_id: str | None = None
key_id: str | None = None
attempt_id: str | None = None
attempt_synced: bool = False
provider_api_format: str | None = None # Provider 的响应格式
# 模型映射
mapped_model: str | None = None
# Token 统计
input_tokens: int = 0
output_tokens: int = 0
cached_tokens: int = 0
cache_creation_tokens: int = 0
cache_creation_tokens_5m: int = 0 # 5min TTL 缓存创建
cache_creation_tokens_1h: int = 0 # 1h TTL 缓存创建
# 响应内容
_collected_text_parts: list[str] = field(default_factory=list, repr=False)
_collected_text_chars: int = field(default=0, repr=False)
_stored_collected_text_chars: int = field(default=0, repr=False)
response_id: str | None = None
final_usage: dict[str, Any] | None = None
final_response: dict[str, Any] | None = None
# 时间指标
first_byte_time_ms: int | None = None # 首字时间 (TTFB - Time To First Byte)
start_time: float = field(default_factory=time.time)
# 响应状态
status_code: int = 200
error_message: str | None = None # 客户端友好的错误消息
upstream_response: str | None = None # 原始 Provider 响应(用于请求链路追踪)
has_completion: bool = False
# 请求/响应数据
response_headers: dict[str, str] = field(default_factory=dict) # 提供商响应头
client_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
# 格式转换信息CLI handler 需要)
client_api_format: str = ""
needs_conversion: bool = False # 是否需要跨格式转换(由 handler 层设置)
# Provider 响应元数据CLI handler 需要)
response_metadata: dict[str, Any] = field(default_factory=dict)
# 整流标记Thinking Rectifier
rectified: bool = False # 请求是否经过整流(移除 thinking 块后重试)
# 流式处理统计
data_count: int = 0
chunk_count: int = 0
parsed_chunks: list[dict[str, Any]] = field(default_factory=list)
# 格式转换时保留提供商原始 chunks转换前的数据
provider_parsed_chunks: list[dict[str, Any]] = field(default_factory=list)
# 是否记录 parsed_chunks可用于降低高并发/长流式响应的内存占用)
record_parsed_chunks: bool = True
# 性能采集(可选)
perf_sampled: bool = False
perf_metrics: dict[str, Any] = field(default_factory=dict)
# 代理信息(用于 usage 记录和日志,含 ttfb_ms
proxy_info: dict[str, Any] | None = None
# 号池调度摘要(来自 ExecutionResult.pool_summary
pool_summary: dict[str, Any] | None = None
# 候选轨迹(来自 ExecutionResult.candidate_keys写入 usage metadata
candidate_keys: list[dict[str, Any]] = field(default_factory=list)
# 内部调度审计摘要(重试/故障转移/账号使用轨迹)
scheduling_audit: dict[str, Any] | None = None
# 流式格式转换状态(跨 chunk 追踪)
stream_conversion_state: StreamState | None = None
stream_conversion_event_count: int = 0 # 流式转换成功的 event 计数
def reset_for_retry(self) -> None:
"""
重试时重置状态
在故障转移重试时调用,清除之前的数据避免累积。
保留 model 和 api_format重置其他所有状态。
"""
self.release_recorded_chunks()
self.chunk_count = 0
self.data_count = 0
self.has_completion = False
self._collected_text_parts = []
self._collected_text_chars = 0
self._stored_collected_text_chars = 0
self.input_tokens = 0
self.output_tokens = 0
self.cached_tokens = 0
self.cache_creation_tokens = 0
self.cache_creation_tokens_5m = 0
self.cache_creation_tokens_1h = 0
self.error_message = None
self.upstream_response = None
self.status_code = 200
self.first_byte_time_ms = None
self.response_headers = {}
self.client_response_headers = {}
self.provider_request_headers = {}
self.provider_request_body = None
self.response_id = None
self.final_usage = None
self.final_response = None
self.proxy_info = None
self.pool_summary = None
self.candidate_keys = []
self.scheduling_audit = None
self.stream_conversion_state = None
self.stream_conversion_event_count = 0
self.needs_conversion = False
self.selected_base_url = None
def release_recorded_chunks(self) -> None:
"""释放 telemetry/usage 已消费完的 chunk 列表,避免后台任务继续持有大对象。"""
self.parsed_chunks = []
self.provider_parsed_chunks = []
@contextmanager
def managed_recorded_bodies(
self,
response_time_ms: int,
*,
include_bodies: bool = True,
) -> Iterator[RecordedStreamBodies]:
"""统一管理响应体构建与 chunk 释放,避免 telemetry 路径重复写 finally。"""
recorded_bodies = RecordedStreamBodies(
response_body=self.build_response_body(response_time_ms) if include_bodies else None,
client_response_body=(
self.build_client_response_body(response_time_ms) if include_bodies else None
),
)
try:
yield recorded_bodies
finally:
self.release_recorded_chunks()
recorded_bodies.response_body = None
recorded_bodies.client_response_body = None
@property
def collected_text(self) -> str:
"""已收集的文本内容(按需拼接,避免在流式过程中频繁做字符串拷贝)"""
return "".join(self._collected_text_parts)
@property
def collected_text_length(self) -> int:
"""已收集文本的总字符数(包含未保留到内存的截断部分)"""
return self._collected_text_chars
def append_text(self, text: str) -> None:
"""追加文本内容(仅保留有限前缀,避免长流导致内存增长)"""
if not text:
return
text_len = len(text)
self._collected_text_chars += text_len
remaining = _MAX_COLLECTED_TEXT_CHARS - self._stored_collected_text_chars
if remaining <= 0:
return
if text_len <= remaining:
self._collected_text_parts.append(text)
self._stored_collected_text_chars += text_len
return
self._collected_text_parts.append(text[:remaining])
self._stored_collected_text_chars += remaining
def update_provider_info(
self,
provider_name: str,
provider_id: str,
endpoint_id: str,
key_id: str,
provider_api_format: str | None = None,
) -> None:
"""更新 Provider 信息"""
self.provider_name = provider_name
self.provider_id = provider_id
self.endpoint_id = endpoint_id
self.key_id = key_id
self.provider_api_format = provider_api_format
def update_usage(
self,
input_tokens: int | None = None,
output_tokens: int | None = None,
cached_tokens: int | None = None,
cache_creation_tokens: int | None = None,
) -> None:
"""
更新 Token 使用统计
采用防御性更新策略:只有当新值 > 0 或当前值为 0 时才更新,避免用 0 覆盖已有的正确值。
设计原理:
- 在流式响应中,某些事件可能不包含完整的 usage 信息(字段为 0 或不存在)
- 后续事件可能会提供完整的统计数据
- 通过这种策略,确保一旦获得非零值就保留它,不会被后续的 0 值覆盖
示例场景:
- message_start 事件input_tokens=100, output_tokens=0
- message_delta 事件input_tokens=0, output_tokens=50
- 最终结果input_tokens=100, output_tokens=50
注意事项:
- 此策略假设初始值为 0 是正确的默认状态
- 如果需要将已有值重置为 0请直接修改实例属性不使用此方法
Args:
input_tokens: 输入 tokens 数量
output_tokens: 输出 tokens 数量
cached_tokens: 缓存命中 tokens 数量
cache_creation_tokens: 缓存创建 tokens 数量
"""
if input_tokens is not None and (input_tokens > 0 or self.input_tokens == 0):
self.input_tokens = input_tokens
if output_tokens is not None and (output_tokens > 0 or self.output_tokens == 0):
self.output_tokens = output_tokens
if cached_tokens is not None and (cached_tokens > 0 or self.cached_tokens == 0):
self.cached_tokens = cached_tokens
if cache_creation_tokens is not None and (
cache_creation_tokens > 0 or self.cache_creation_tokens == 0
):
self.cache_creation_tokens = cache_creation_tokens
def mark_failed(
self,
status_code: int,
error_message: str,
upstream_response: str | None = None,
) -> None:
"""
标记请求失败
Args:
status_code: HTTP 状态码
error_message: 客户端友好的错误消息
upstream_response: 原始 Provider 响应(用于请求链路追踪)
"""
self.status_code = status_code
self.error_message = error_message
if upstream_response:
self.upstream_response = upstream_response
def record_first_byte_time(self, start_time: float) -> None:
"""
记录首字时间 (TTFB - Time To First Byte)
应在第一次向客户端发送数据时调用。
如果已记录过,则不会覆盖(避免重试时重复记录)。
Args:
start_time: 请求开始时间 (time.time())
"""
if self.first_byte_time_ms is None:
self.first_byte_time_ms = int((time.time() - start_time) * 1000)
@property
def has_format_conversion(self) -> bool:
"""是否发生了真正的格式转换client 和 provider 的 api_format 不同)
区别于 needs_conversion后者包含 envelope rewrite如 Antigravity v1internal
不代表客户端与上游的数据格式真正不同。此属性用于 usage 展示层。
"""
return is_format_converted(self.provider_api_format, self.client_api_format)
def is_success(self) -> bool:
"""检查请求是否成功"""
return self.status_code < 400
def is_client_disconnected(self) -> bool:
"""检查是否因客户端断开连接而结束"""
return self.status_code == 499
def has_partial_response(self) -> bool:
"""是否已收到部分流式响应数据。"""
return self.data_count > 0 or self.chunk_count > 0 or self.collected_text_length > 0
def ensure_estimated_output_tokens(self) -> bool:
"""在缺少 usage 时,基于已收集文本补充输出 tokens。"""
if self.output_tokens > 0 or self.collected_text_length <= 0:
return False
self.output_tokens = max(1, self.collected_text_length // 4)
return True
def should_estimate_incomplete_tokens(self) -> bool:
"""流异常结束且尚无 usage 时,是否应做兜底 token 估算。
使用 or 而非 andCancelledError 路径中 ensure_estimated_output_tokens
可能已补了 output_tokens但 input_tokens 仍为 0此时仍需估算。
"""
return (
not self.has_completion
and (self.input_tokens == 0 or self.output_tokens == 0)
and self.has_partial_response()
)
def set_ttfb_ms(self, ms: int) -> None:
"""将首字节响应耗时TTFB注入到 proxy_info 中"""
if self.proxy_info is not None:
self.proxy_info["ttfb_ms"] = ms
def set_proxy_timing(self, headers: dict[str, str]) -> None:
"""从代理响应头中提取分阶段耗时信息X-Proxy-Timing"""
extract_proxy_timing(self.proxy_info, headers)
def build_response_body(self, response_time_ms: int) -> dict[str, Any]:
"""
构建响应体元数据
用于记录到 Usage 表的 response_body 字段。
当有格式转换时,返回提供商原始 chunks否则返回 parsed_chunks。
"""
chunks = self.provider_parsed_chunks if self.provider_parsed_chunks else self.parsed_chunks
return {
"chunks": chunks,
"metadata": {
"stream": True,
"total_chunks": len(chunks),
"data_count": self.data_count,
"has_completion": self.has_completion,
"response_time_ms": response_time_ms,
},
}
def build_client_response_body(self, response_time_ms: int) -> dict[str, Any] | None:
"""
构建客户端侧响应体元数据
仅当有格式转换时返回parsed_chunks 是转换后的客户端格式);
无格式转换时返回 None此时 parsed_chunks 已在 response_body 中)。
"""
if not self.provider_parsed_chunks:
return None
return {
"chunks": self.parsed_chunks,
"metadata": {
"stream": True,
"total_chunks": len(self.parsed_chunks),
"data_count": self.data_count,
"has_completion": self.has_completion,
"response_time_ms": response_time_ms,
},
}
def get_log_summary(self, request_id: str, response_time_ms: int) -> str:
"""
获取日志摘要
用于请求完成/失败时的日志输出。
包含首字时间 (TTFB) 和总响应时间,分两行显示。
"""
if self.is_success():
status = "OK"
elif self.is_client_disconnected():
status = "CANCEL"
else:
status = "FAIL"
# 第一行:基本信息 + 首字时间
line1 = (
f"[{status}] {request_id[:8]} | {self.model} | " f"{self.provider_name or 'unknown'}"
)
if self.first_byte_time_ms is not None:
line1 += f" | TTFB: {self.first_byte_time_ms}ms"
# 第二行:总响应时间 + tokens
line2 = (
f" Total: {response_time_ms}ms | "
f"in:{self.input_tokens} out:{self.output_tokens}"
)
return f"{line1}\n{line2}"

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,674 @@
"""
流式遥测记录器 - 从 ChatHandlerBase 提取的统计记录逻辑
职责:
1. 记录流式请求的成功/失败统计
2. 更新 Usage 状态
3. 更新候选记录状态
"""
import asyncio
import json
import time
from typing import Any
from sqlalchemy.orm import Session
from src.api.handlers.base.base_handler import MessageTelemetry
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.utils import filter_proxy_response_headers
from src.config.settings import config
from src.core.logger import logger
from src.database import get_db
from src.models.database import ApiKey, User
from src.services.system.config import SystemConfigService
from src.services.usage.telemetry_writer import (
DbTelemetryWriter,
QueueTelemetryWriter,
TelemetryWriter,
)
class StreamTelemetryRecorder:
"""
流式遥测记录器
负责在流式请求完成后记录统计信息。
从 ChatHandlerBase 中提取的 _record_stream_stats 逻辑。
"""
def __init__(
self,
request_id: str,
user_id: str,
api_key_id: str,
client_ip: str,
format_id: str,
):
"""
初始化遥测记录器
Args:
request_id: 请求 ID
user_id: 用户 ID
api_key_id: API Key ID
client_ip: 客户端 IP
format_id: API 格式标识
"""
self.request_id = request_id
self.user_id = user_id
self.api_key_id = api_key_id
self.client_ip = client_ip
self.format_id = format_id
async def record_stream_stats(
self,
ctx: StreamContext,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
start_time: float,
) -> None:
"""
记录流式统计信息
Args:
ctx: 流式上下文
original_headers: 原始请求头
original_request_body: 原始请求体
start_time: 请求开始时间 (time.time())
"""
bg_db = None
try:
# 在流结束后计算响应时间,与首字时间使用相同的时间基准
# 注意不要把统计延迟stream_stats_delay算进响应时间里
response_time_ms = int((time.time() - start_time) * 1000)
await asyncio.sleep(config.stream_stats_delay) # 等待流完全关闭
if not ctx.provider_name:
await self._update_usage_status_on_error(
response_time_ms=response_time_ms,
error_message="Provider name not available",
)
return
db_gen = get_db()
bg_db = next(db_gen)
try:
writer = await self._get_telemetry_writer(bg_db, ctx, response_time_ms)
if writer is None:
ctx.release_recorded_chunks()
return
# 兜底估算:流未正常完成且 token 均为 0 时,从请求体粗略估算。
# 覆盖成功但缺少 completion以及已传出部分数据后被中断的场景。
if ctx.should_estimate_incomplete_tokens():
# 用实际发给 Provider 的请求体估算 token格式转换时与客户端请求体不同
self._estimate_tokens_for_incomplete_stream(
ctx, ctx.provider_request_body or original_request_body
)
should_log_body = SystemConfigService.should_log_body(bg_db)
include_bodies = (
writer.include_bodies
if isinstance(writer, QueueTelemetryWriter)
else should_log_body
)
with ctx.managed_recorded_bodies(
response_time_ms, include_bodies=include_bodies
) as recorded_bodies:
try:
await self._dispatch_record(
bg_db,
writer,
ctx,
original_headers,
original_request_body,
recorded_bodies.response_body,
response_time_ms,
client_response_body=recorded_bodies.client_response_body,
)
except Exception as writer_error:
if not isinstance(writer, QueueTelemetryWriter):
raise
logger.warning(
f"[{self.request_id}] Queue writer failed, falling back to DB: {writer_error}"
)
db_writer = self._build_db_writer(bg_db)
if db_writer is None:
await self._update_usage_status_directly(
bg_db,
status=self._get_status_from_ctx(ctx),
response_time_ms=response_time_ms,
status_code=ctx.status_code,
)
return
if should_log_body:
recorded_bodies.ensure_populated(ctx, response_time_ms)
await self._dispatch_record(
bg_db,
db_writer,
ctx,
original_headers,
original_request_body,
recorded_bodies.response_body,
response_time_ms,
client_response_body=recorded_bodies.client_response_body,
)
# 更新候选记录状态
await self._update_candidate_status(bg_db, ctx, response_time_ms, start_time)
finally:
if bg_db:
bg_db.close()
except Exception as e:
logger.exception("记录流式统计信息时出错")
await self._update_usage_status_on_error(
response_time_ms=response_time_ms,
error_message=f"记录统计信息失败: {str(e)[:200]}",
)
finally:
# 遥测写入后主动释放大列表,避免长流式请求对象滞留在 worker 堆中。
ctx.release_recorded_chunks()
async def _record_success(
self,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
response_body: dict[str, Any] | None,
response_time_ms: int,
client_response_body: dict[str, Any] | None = None,
) -> None:
"""记录成功的请求"""
# 流式成功时,返回给客户端的是提供商响应头 + SSE 必需头
client_response_headers = filter_proxy_response_headers(ctx.response_headers)
client_response_headers.update(
{
"Cache-Control": "no-cache, no-transform",
"X-Accel-Buffering": "no",
"content-type": "text/event-stream",
}
)
metadata: dict[str, Any] = {"stream": True, "content_length": ctx.data_count}
if ctx.perf_metrics:
metadata["perf"] = ctx.perf_metrics
if ctx.proxy_info:
metadata["proxy"] = ctx.proxy_info
if ctx.pool_summary:
metadata["pool_summary"] = ctx.pool_summary
if ctx.candidate_keys:
metadata["candidate_keys"] = ctx.candidate_keys
if ctx.scheduling_audit:
metadata["scheduling_audit"] = ctx.scheduling_audit
await writer.record_success(
provider=ctx.provider_name or "unknown",
model=ctx.model,
input_tokens=ctx.input_tokens,
output_tokens=ctx.output_tokens,
response_time_ms=response_time_ms,
first_byte_time_ms=ctx.first_byte_time_ms, # 传递首字时间
status_code=ctx.status_code,
request_headers=original_headers,
request_body=original_request_body,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
response_body=response_body,
client_response_body=client_response_body,
provider_request_body=ctx.provider_request_body,
cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens,
cache_creation_tokens_5m=ctx.cache_creation_tokens_5m,
cache_creation_tokens_1h=ctx.cache_creation_tokens_1h,
is_stream=True,
provider_request_headers=ctx.provider_request_headers,
api_format=ctx.api_format,
api_family=ctx.api_family,
endpoint_kind=ctx.endpoint_kind,
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
target_model=ctx.mapped_model,
request_type="chat",
metadata=metadata,
endpoint_api_format=ctx.provider_api_format,
has_format_conversion=ctx.has_format_conversion,
)
logger.debug(f"{self.format_id} 流式响应完成")
logger.info(ctx.get_log_summary(self.request_id, response_time_ms))
async def _record_failure(
self,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
response_body: dict[str, Any] | None,
response_time_ms: int,
client_response_body: dict[str, Any] | None = None,
) -> None:
"""记录失败的请求"""
# 失败时返回给客户端的是 JSON 错误响应,如果没有设置则使用默认值
client_response_headers = ctx.client_response_headers or {
"content-type": "application/json"
}
metadata: dict[str, Any] = {"stream": True, "content_length": ctx.data_count}
if ctx.perf_metrics:
metadata["perf"] = ctx.perf_metrics
if ctx.proxy_info:
metadata["proxy"] = ctx.proxy_info
if ctx.pool_summary:
metadata["pool_summary"] = ctx.pool_summary
if ctx.candidate_keys:
metadata["candidate_keys"] = ctx.candidate_keys
if ctx.scheduling_audit:
metadata["scheduling_audit"] = ctx.scheduling_audit
await writer.record_failure(
provider=ctx.provider_name or "unknown",
model=ctx.model,
response_time_ms=response_time_ms,
status_code=ctx.status_code,
error_message=ctx.error_message or f"HTTP {ctx.status_code}",
request_headers=original_headers,
request_body=original_request_body,
is_stream=True,
api_format=ctx.api_format,
api_family=ctx.api_family,
endpoint_kind=ctx.endpoint_kind,
provider_request_headers=ctx.provider_request_headers,
provider_request_body=ctx.provider_request_body,
input_tokens=ctx.input_tokens,
output_tokens=ctx.output_tokens,
cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens,
cache_creation_tokens_5m=ctx.cache_creation_tokens_5m,
cache_creation_tokens_1h=ctx.cache_creation_tokens_1h,
response_body=response_body,
client_response_body=client_response_body,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
target_model=ctx.mapped_model,
request_type="chat",
metadata=metadata,
endpoint_api_format=ctx.provider_api_format,
has_format_conversion=ctx.has_format_conversion,
)
logger.debug(f"{self.format_id} 流式响应中断")
log_summary = ctx.get_log_summary(self.request_id, response_time_ms)
# 对于失败日志,添加缓存信息
logger.info(f"{log_summary} cache:{ctx.cached_tokens}")
async def _record_cancelled(
self,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
response_body: dict[str, Any] | None,
response_time_ms: int,
client_response_body: dict[str, Any] | None = None,
) -> None:
"""记录客户端取消的请求"""
client_response_headers = ctx.client_response_headers or {
"content-type": "application/json"
}
metadata: dict[str, Any] = {"stream": True, "content_length": ctx.data_count}
if ctx.perf_metrics:
metadata["perf"] = ctx.perf_metrics
if ctx.proxy_info:
metadata["proxy"] = ctx.proxy_info
if ctx.pool_summary:
metadata["pool_summary"] = ctx.pool_summary
if ctx.candidate_keys:
metadata["candidate_keys"] = ctx.candidate_keys
if ctx.scheduling_audit:
metadata["scheduling_audit"] = ctx.scheduling_audit
await writer.record_cancelled(
provider=ctx.provider_name or "unknown",
model=ctx.model,
response_time_ms=response_time_ms,
first_byte_time_ms=ctx.first_byte_time_ms,
status_code=ctx.status_code,
request_headers=original_headers,
request_body=original_request_body,
is_stream=True,
api_format=ctx.api_format,
api_family=ctx.api_family,
endpoint_kind=ctx.endpoint_kind,
provider_request_headers=ctx.provider_request_headers,
provider_request_body=ctx.provider_request_body,
input_tokens=ctx.input_tokens,
output_tokens=ctx.output_tokens,
cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens,
cache_creation_tokens_5m=ctx.cache_creation_tokens_5m,
cache_creation_tokens_1h=ctx.cache_creation_tokens_1h,
response_body=response_body,
client_response_body=client_response_body,
response_headers=ctx.response_headers,
client_response_headers=client_response_headers,
provider_id=ctx.provider_id,
provider_endpoint_id=ctx.endpoint_id,
provider_api_key_id=ctx.key_id,
target_model=ctx.mapped_model,
request_type="chat",
metadata=metadata,
endpoint_api_format=ctx.provider_api_format,
has_format_conversion=ctx.has_format_conversion,
)
logger.debug(f"{self.format_id} 流式响应被客户端取消")
logger.info(ctx.get_log_summary(self.request_id, response_time_ms))
async def _update_candidate_status(
self,
db: Session,
ctx: StreamContext,
response_time_ms: int,
request_start_time: float,
) -> None:
"""更新候选记录状态"""
if not ctx.attempt_id:
return
# Capture all needed ctx attrs before handing off to thread
attempt_id = ctx.attempt_id
is_success = ctx.is_success()
is_disconnected = ctx.is_client_disconnected()
status_code = ctx.status_code
data_count = ctx.data_count
rectified = ctx.rectified
proxy_info = ctx.proxy_info
first_byte_time_ms_val = ctx.first_byte_time_ms
upstream_response = ctx.upstream_response
error_message = ctx.error_message
def _sync() -> None:
from src.services.request.candidate import RequestCandidateService
extra_data: dict[str, Any] = {
"stream_completed": is_success,
"data_count": data_count,
}
if rectified:
extra_data["rectified"] = True
if proxy_info:
extra_data["proxy"] = proxy_info
if first_byte_time_ms_val is not None:
candidate_ttfb = RequestCandidateService.calculate_candidate_ttfb(
db=db,
candidate_id=attempt_id,
request_start_time=request_start_time,
global_first_byte_time_ms=first_byte_time_ms_val,
)
extra_data["first_byte_time_ms"] = candidate_ttfb
if is_success:
RequestCandidateService.mark_candidate_success(
db=db,
candidate_id=attempt_id,
status_code=status_code,
latency_ms=response_time_ms,
extra_data=extra_data,
)
elif is_disconnected:
RequestCandidateService.mark_candidate_cancelled(
db=db,
candidate_id=attempt_id,
status_code=status_code,
latency_ms=response_time_ms,
extra_data=extra_data,
)
else:
trace_error_message = upstream_response or error_message or f"HTTP {status_code}"
RequestCandidateService.mark_candidate_failed(
db=db,
candidate_id=attempt_id,
error_type="stream_error",
error_message=trace_error_message,
status_code=status_code,
latency_ms=response_time_ms,
extra_data=extra_data,
)
await asyncio.to_thread(_sync)
async def _update_usage_status_on_error(
self,
response_time_ms: int,
error_message: str,
) -> None:
"""在记录失败时更新 Usage 状态"""
try:
db_gen = get_db()
error_db = next(db_gen)
try:
await self._update_usage_status_directly(
error_db,
status="failed",
response_time_ms=response_time_ms,
status_code=500,
error_message=error_message,
)
finally:
error_db.close()
except Exception as inner_e:
logger.error(f"[{self.request_id}] 更新 Usage 状态失败: {inner_e}")
async def _update_usage_status_directly(
self,
db: Session,
status: str,
response_time_ms: int,
status_code: int = 200,
error_message: str | None = None,
) -> None:
"""直接更新 Usage 表的状态字段"""
request_id = self.request_id
def _sync() -> None:
from src.models.database import Usage
usage = db.query(Usage).filter(Usage.request_id == request_id).first()
if usage:
if getattr(usage, "billing_status", None) in {"settled", "void"}:
logger.debug("[{}] Usage 已终态,跳过快速状态更新: {}", self.request_id, status)
return
setattr(usage, "status", status)
setattr(usage, "status_code", status_code)
setattr(usage, "response_time_ms", response_time_ms)
if error_message:
setattr(usage, "error_message", error_message)
db.commit()
logger.debug(f"[{request_id}] Usage 状态已更新: {status}")
try:
await asyncio.to_thread(_sync)
except Exception as e:
logger.error(f"[{self.request_id}] 直接更新 Usage 状态失败: {e}")
async def _get_telemetry_writer(
self, bg_db: Session, ctx: StreamContext, response_time_ms: int
) -> TelemetryWriter | None:
if config.usage_queue_enabled and self.user_id and self.api_key_id:
# Queue payload detail follows system config request_record_level.
log_level = SystemConfigService.get_request_record_level(bg_db).value
sensitive_headers = SystemConfigService.get_sensitive_headers(bg_db) or []
max_request_body_size = int(
SystemConfigService.get_config(bg_db, "max_request_body_size", 5242880) or 0
)
max_response_body_size = int(
SystemConfigService.get_config(bg_db, "max_response_body_size", 5242880) or 0
)
return QueueTelemetryWriter(
request_id=self.request_id,
user_id=self.user_id,
api_key_id=self.api_key_id,
log_level=log_level,
sensitive_headers=sensitive_headers,
max_request_body_size=max_request_body_size,
max_response_body_size=max_response_body_size,
)
db_writer = self._build_db_writer(bg_db)
if db_writer is None:
await self._update_usage_status_directly(
bg_db,
status=self._get_status_from_ctx(ctx),
response_time_ms=response_time_ms,
status_code=ctx.status_code,
)
return None
return db_writer
async def _dispatch_record(
self,
db: Session,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
response_body: dict[str, Any] | None,
response_time_ms: int,
client_response_body: dict[str, Any] | None = None,
) -> None:
"""根据上下文状态分发到对应的记录方法"""
if ctx.is_success():
await self._record_success(
writer,
ctx,
original_headers,
original_request_body,
response_body,
response_time_ms,
client_response_body=client_response_body,
)
# Queue writer 异步落库可能造成 UI 延迟,先直接更新 Usage 状态
if isinstance(writer, QueueTelemetryWriter):
await self._update_usage_status_directly(
db=db,
status=self._get_status_from_ctx(ctx),
response_time_ms=response_time_ms,
status_code=ctx.status_code,
)
elif ctx.is_client_disconnected():
await self._record_cancelled(
writer,
ctx,
original_headers,
original_request_body,
response_body,
response_time_ms,
client_response_body=client_response_body,
)
# Queue writer 异步落库可能造成 UI 延迟,先直接更新 Usage 状态
if isinstance(writer, QueueTelemetryWriter):
await self._update_usage_status_directly(
db=db,
status="cancelled",
response_time_ms=response_time_ms,
status_code=ctx.status_code,
)
else:
await self._record_failure(
writer,
ctx,
original_headers,
original_request_body,
response_body,
response_time_ms,
client_response_body=client_response_body,
)
# Queue writer 异步落库可能造成 UI 延迟,先直接更新 Usage 状态
if isinstance(writer, QueueTelemetryWriter):
await self._update_usage_status_directly(
db=db,
status=self._get_status_from_ctx(ctx),
response_time_ms=response_time_ms,
status_code=ctx.status_code,
error_message=ctx.error_message or f"HTTP {ctx.status_code}",
)
def _get_status_from_ctx(self, ctx: StreamContext) -> str:
"""根据上下文获取状态字符串"""
if ctx.is_success():
return "completed"
if ctx.is_client_disconnected():
return "cancelled"
return "failed"
@staticmethod
def _estimate_tokens_for_incomplete_stream(
ctx: StreamContext,
request_body: dict[str, Any],
) -> None:
"""
流未正常完成(无 response.completed且 token 均为 0 时的兜底估算。
从已收集的输出文本和请求体粗略估算 token 数,确保 usage 记录不为 0。
估算采用 ~4 字符/token 的保守比例。
"""
# 输出 tokens从已收集的文本估算
if ctx.collected_text_length > 0:
ctx.output_tokens = max(1, ctx.collected_text_length // 4)
# 输入 tokens从请求体文本内容估算
try:
total_input_len = 0
instructions = request_body.get("instructions")
if isinstance(instructions, str):
total_input_len += len(instructions)
# OpenAI Responses API 使用 input 字段Claude 使用 messages
input_items = request_body.get("input") or request_body.get("messages") or []
if isinstance(input_items, list):
for item in input_items:
if isinstance(item, str):
total_input_len += len(item)
elif isinstance(item, dict):
content = item.get("content", "")
if isinstance(content, str):
total_input_len += len(content)
elif isinstance(content, list):
for block in content:
if isinstance(block, dict):
text = block.get("text", "")
if isinstance(text, str):
total_input_len += len(text)
if total_input_len > 0:
ctx.input_tokens = max(1, total_input_len // 4)
else:
# fallback: 整个请求体 JSON 大小
body_str = json.dumps(request_body, ensure_ascii=False)
ctx.input_tokens = max(1, len(body_str) // 4)
except Exception:
pass
if ctx.input_tokens > 0 or ctx.output_tokens > 0:
logger.warning(
f"[{ctx.request_id}] 流未正常完成 (has_completion=False, data_count={ctx.data_count}), "
f"使用估算 tokens: in={ctx.input_tokens}, out={ctx.output_tokens}"
)
def _build_db_writer(self, bg_db: Session) -> DbTelemetryWriter | None:
user = bg_db.query(User).filter(User.id == self.user_id).first()
api_key_obj = bg_db.query(ApiKey).filter(ApiKey.id == self.api_key_id).first()
if not user or not api_key_obj:
logger.warning(
f"[{self.request_id}] User or ApiKey not found, updating status directly"
)
return None
bg_telemetry = MessageTelemetry(bg_db, user, api_key_obj, self.request_id, self.client_ip)
return DbTelemetryWriter(bg_telemetry)

View File

@@ -0,0 +1,246 @@
"""Upstream stream bridging helpers (handler layer).
This module provides small utilities used when handler-layer policies force an
upstream request to be streaming (SSE) even when the client asked for sync.
It intentionally stays lightweight and works with:
- standard SSE `data: {...}` lines (OpenAI/Claude/etc.)
- Gemini CLI JSON-array lines (best-effort)
"""
from __future__ import annotations
import codecs
import json
from collections import Counter
from collections.abc import AsyncIterator
from typing import Any
from src.api.handlers.base.response_parser import ResponseParser
from src.api.handlers.base.utils import (
ensure_stream_buffer_limit,
get_format_converter_registry,
)
from src.core.api_format.conversion.internal import InternalResponse
from src.core.api_format.conversion.stream_bridge import InternalStreamAggregator
from src.core.api_format.conversion.stream_state import StreamState
from src.core.exceptions import EmbeddedErrorException
from src.core.logger import logger
from src.services.provider.envelope import ProviderEnvelope
def _parse_sse_data_line(line: str) -> tuple[Any | None, str]:
"""Parse `data: {...}` as JSON."""
payload = line[5:].strip()
if not payload:
return None, "empty"
try:
return json.loads(payload), "ok"
except json.JSONDecodeError:
return None, "invalid"
def _parse_sse_event_data_line(line: str) -> tuple[Any | None, str]:
"""Parse `event: xxx data: {...}` as JSON (best-effort)."""
# Split only on the first " data:" occurrence.
try:
_, data_part = line.split(" data:", 1)
except ValueError:
return None, "invalid"
payload = data_part.strip()
if not payload:
return None, "empty"
try:
return json.loads(payload), "ok"
except json.JSONDecodeError:
return None, "invalid"
def _parse_gemini_json_array_line(line: str) -> tuple[Any | None, str]:
"""Parse Gemini CLI JSON-array streaming line (best-effort).
Gemini CLI may stream objects in a JSON array form like:
- "[{...},"
- " {...},"
- " {...}]"
"""
stripped = (line or "").strip()
if not stripped:
return None, "empty"
# Quick filter: must contain a JSON object boundary.
if "{" not in stripped:
return None, "skip"
candidate = stripped.lstrip(",").rstrip(",").strip()
# Drop array brackets on edges.
if candidate.startswith("["):
candidate = candidate[1:].strip()
if candidate.endswith("]"):
candidate = candidate[:-1].strip()
candidate = candidate.lstrip(",").rstrip(",").strip()
if not candidate:
return None, "empty"
try:
return json.loads(candidate), "ok"
except json.JSONDecodeError:
logger.debug(f"Gemini JSON-array line skip: {stripped[:50]}")
return None, "invalid"
def parse_provider_stream_line_to_json(
line: str,
provider_format: str,
) -> tuple[Any | None, str]:
"""Best-effort parse for upstream streaming lines (SSE or Gemini JSON-array)."""
if not line:
return None, "skip"
normalized = line.rstrip("\r").strip("\n")
if not normalized or normalized.strip() == "":
return None, "skip"
# Standard SSE data line.
if normalized.startswith("data:"):
# `data: [DONE]` is a sentinel.
if normalized[5:].strip() == "[DONE]":
return None, "skip"
return _parse_sse_data_line(normalized)
# event + data on same line.
if normalized.startswith("event:") and " data:" in normalized:
return _parse_sse_event_data_line(normalized)
# Other control lines.
if normalized.startswith(("event:", "id:", "retry:")):
return None, "skip"
# Gemini JSON-array/chunked streaming (no SSE prefix).
if str(provider_format or "").strip().lower().startswith("gemini"):
return _parse_gemini_json_array_line(normalized)
return None, "skip"
async def aggregate_upstream_stream_to_internal_response(
byte_iter: AsyncIterator[bytes],
*,
provider_api_format: str,
provider_name: str,
model: str,
request_id: str,
envelope: ProviderEnvelope | None = None,
provider_parser: ResponseParser | None = None,
) -> InternalResponse:
"""Aggregate upstream SSE/streaming bytes into an InternalResponse (best-effort)."""
registry = get_format_converter_registry()
src_norm = registry.get_normalizer(provider_api_format) if provider_api_format else None
if src_norm is None:
raise RuntimeError(f"未注册 Normalizer: {provider_api_format}")
if not getattr(src_norm, "capabilities", None) or not src_norm.capabilities.supports_stream:
raise RuntimeError(f"上游格式不支持流式: {provider_api_format}")
state = StreamState(model=str(model or ""), message_id=str(request_id or ""))
aggregator = InternalStreamAggregator(
fallback_id=str(request_id or "resp"),
fallback_model=str(model or ""),
)
buffer = b""
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
_event_type_counts: Counter[str] = Counter()
_total_events = 0
def _feed_line(normalized_line: str) -> None:
nonlocal _total_events
data_obj, st = parse_provider_stream_line_to_json(normalized_line, provider_api_format)
if st != "ok" or data_obj is None:
return
if not isinstance(data_obj, dict):
return
if envelope:
unwrapped = envelope.unwrap_response(data_obj)
if not isinstance(unwrapped, dict):
return
data_obj = unwrapped
envelope.postprocess_unwrapped_response(model=model, data=data_obj)
if provider_parser and provider_parser.is_error_response(data_obj):
parsed = provider_parser.parse_response(data_obj, 200)
raise EmbeddedErrorException(
provider_name=str(provider_name),
error_code=parsed.embedded_status_code,
error_message=parsed.error_message,
error_status=parsed.error_type,
)
etype = str(data_obj.get("type") or "")
_event_type_counts[etype] += 1
_total_events += 1
internal_events = src_norm.stream_chunk_to_internal(data_obj, state)
aggregator.feed(internal_events)
async for chunk in byte_iter:
buffer += chunk
ensure_stream_buffer_limit(
buffer,
request_id=str(request_id or ""),
provider_name=str(provider_name or "unknown"),
)
while b"\n" in buffer:
line_bytes, buffer = buffer.split(b"\n", 1)
line = decoder.decode(line_bytes + b"\n", False).rstrip("\n")
normalized_line = line.rstrip("\r")
_feed_line(normalized_line)
# Flush remaining buffered bytes (in case upstream doesn't end with newline).
if buffer:
try:
tail = decoder.decode(buffer, True)
except Exception:
tail = ""
normalized_tail = (tail or "").rstrip("\r\n")
if normalized_tail:
_feed_line(normalized_tail)
# 诊断日志:在 build() 之前记录聚合器状态
open_count = aggregator.open_count
final_count = aggregator.final_count
if not final_count and not open_count:
logger.warning(
"[{}] aggregate_upstream_stream: 聚合器无内容, "
"open={}, final={}, usage={}, stop_reason={}, "
"event_types={}, total_events={}",
request_id,
open_count,
final_count,
aggregator.usage,
aggregator.stop_reason,
dict(_event_type_counts),
_total_events,
)
elif not final_count and open_count:
logger.warning(
"[{}] aggregate_upstream_stream: final 为空但 open 有内容, "
"open={}, final={}, event_types={}, total_events={}",
request_id,
open_count,
final_count,
dict(_event_type_counts),
_total_events,
)
result = aggregator.build()
return result
__all__ = [
"aggregate_upstream_stream_to_internal_response",
"parse_provider_stream_line_to_json",
]

View File

@@ -0,0 +1,255 @@
"""
Handler 基础工具函数
"""
from __future__ import annotations
import gzip
import json
from typing import TYPE_CHECKING, Any
from fastapi.responses import JSONResponse, Response
from src.config.constants import StreamDefaults
from src.core.api_format import filter_response_headers
from src.core.api_format.headers import get_header_value
from src.core.exceptions import EmbeddedErrorException, ProviderNotAvailableException
from src.core.http_compression import accepts_gzip, normalize_content_encoding
from src.core.logger import logger
if TYPE_CHECKING:
from src.core.api_format.conversion.registry import FormatConversionRegistry
def get_format_converter_registry() -> FormatConversionRegistry:
"""
获取格式转换注册表(线程安全)
该函数确保 normalizers 已注册后再返回全局注册表实例。
register_default_normalizers 内部已有双重检查锁,可安全多次调用。
"""
from src.core.api_format.conversion.registry import (
format_conversion_registry,
register_default_normalizers,
)
register_default_normalizers()
return format_conversion_registry
def build_sse_headers(extra_headers: dict[str, str] | None = None) -> dict[str, str]:
"""
构建 SSEtext/event-stream推荐响应头用于减少代理缓冲带来的卡顿/成段输出。
说明:
- Cache-Control: no-transform 可避免部分代理对流做压缩/改写导致缓冲
- X-Accel-Buffering: no 可显式提示 Nginx 关闭缓冲(即使全局已关闭也无害)
"""
headers: dict[str, str] = {
"Cache-Control": "no-cache, no-transform",
"X-Accel-Buffering": "no",
}
if extra_headers:
headers.update(extra_headers)
return headers
def filter_proxy_response_headers(headers: dict[str, str] | None) -> dict[str, str]:
"""
过滤上游响应头中不应透传给客户端的字段。
主要用于“解析/转换后再返回”的场景:
- 非流式:我们会 `resp.json()` 后再由 `JSONResponse` 重新序列化
- 流式:我们会解析/重组 SSE 行再输出
如果透传上游的 `content-length/content-encoding/...`,会导致客户端解码失败或等待更多字节。
"""
return filter_response_headers(headers)
def resolve_client_content_encoding(
original_headers: dict[str, str],
hinted_content_encoding: str | None = None,
) -> str | None:
"""解析客户端请求体编码(优先使用上层透传值)。"""
if hinted_content_encoding is not None:
return normalize_content_encoding(hinted_content_encoding)
return normalize_content_encoding(get_header_value(original_headers, "content-encoding"))
def resolve_client_accept_encoding(
original_headers: dict[str, str],
hinted_accept_encoding: str | None = None,
) -> str | None:
"""解析客户端 Accept-Encoding优先使用上层透传值"""
if isinstance(hinted_accept_encoding, str):
normalized_hint = hinted_accept_encoding.strip()
if normalized_hint:
return normalized_hint
header_value = get_header_value(original_headers, "accept-encoding")
normalized_header = header_value.strip()
return normalized_header or None
def build_json_response_for_client(
*,
status_code: int,
content: Any,
headers: dict[str, str] | None,
client_accept_encoding: str | None,
) -> Response:
"""根据客户端 Accept-Encoding 返回普通或 gzip 压缩 JSON 响应。"""
response_headers = dict(headers or {})
response_headers.setdefault("content-type", "application/json")
if not accepts_gzip(client_accept_encoding):
return JSONResponse(status_code=status_code, content=content, headers=response_headers)
json_bytes = json.dumps(content, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
compressed_bytes = gzip.compress(json_bytes, compresslevel=6)
cleaned_headers = {
key: value
for key, value in response_headers.items()
if key.lower() not in {"content-length", "content-encoding", "vary"}
}
cleaned_headers["Content-Encoding"] = "gzip"
existing_vary = next((v for k, v in response_headers.items() if k.lower() == "vary"), "")
vary_values = [part.strip() for part in str(existing_vary).split(",") if part.strip()]
if not any(part.lower() == "accept-encoding" for part in vary_values):
vary_values.append("Accept-Encoding")
if vary_values:
cleaned_headers["Vary"] = ", ".join(vary_values)
return Response(
status_code=status_code,
content=compressed_bytes,
headers=cleaned_headers,
media_type="application/json",
)
def check_html_response(line: str) -> bool:
"""
检查行是否为 HTML 响应base_url 配置错误的常见症状)
Args:
line: 要检查的行内容
Returns:
True 如果检测到 HTML 响应
"""
lower_line = line.lstrip().lower()
return lower_line.startswith("<!doctype") or lower_line.startswith("<html")
def ensure_stream_buffer_limit(
buffer: bytes,
*,
request_id: str,
provider_name: str | None = None,
) -> None:
"""防止上游单行流式数据过大导致内存失控。"""
total_limit = StreamDefaults.MAX_STREAM_BUFFER_TOTAL_BYTES
if len(buffer) > total_limit:
raise ProviderNotAvailableException(
"上游流式响应异常:总缓冲区超过安全上限",
provider_name=provider_name,
upstream_status=502,
upstream_response=(
f"stream buffer total overflow: {len(buffer)} bytes > {total_limit}, "
f"request_id={request_id}"
),
)
limit = StreamDefaults.MAX_STREAM_BUFFER_BYTES
if len(buffer) <= limit:
return
# 允许 chunk 中存在大量完整行;只限制“最后一行未闭合缓冲”的体积。
trailing_line = buffer.rsplit(b"\n", 1)[-1]
if len(trailing_line) <= limit:
return
raise ProviderNotAvailableException(
"上游流式响应异常:单行数据超过安全上限",
provider_name=provider_name,
upstream_status=502,
upstream_response=f"stream buffer overflow: {len(buffer)} bytes > {limit}, request_id={request_id}",
)
def check_prefetched_response_error(
prefetched_chunks: list,
parser: Any,
request_id: str,
provider_name: str,
endpoint_id: str | None,
base_url: str | None,
) -> None:
"""
检查预读的响应是否为非 SSE 格式的错误响应HTML 或纯 JSON 错误)
某些代理可能返回:
1. HTML 页面base_url 配置错误)
2. 纯 JSON 错误(无换行或多行 JSON
Args:
prefetched_chunks: 预读的字节块列表
parser: 响应解析器(需要有 is_error_response 和 parse_response 方法)
request_id: 请求 ID用于日志
provider_name: Provider 名称
endpoint_id: Endpoint ID
base_url: Endpoint 的 base_url
Raises:
ProviderNotAvailableException: 如果检测到 HTML 响应
EmbeddedErrorException: 如果检测到 JSON 错误响应
"""
if not prefetched_chunks:
return
try:
prefetched_bytes = b"".join(prefetched_chunks)
stripped = prefetched_bytes.lstrip()
# 去除 BOM
if stripped.startswith(b"\xef\xbb\xbf"):
stripped = stripped[3:]
# HTML 响应(通常是 base_url 配置错误导致返回网页)
lower_prefix = stripped[:32].lower()
if lower_prefix.startswith(b"<!doctype") or lower_prefix.startswith(b"<html"):
endpoint_short = endpoint_id[:8] + "..." if endpoint_id else "N/A"
logger.error(
f" [{request_id}] 检测到 HTML 响应,可能是 base_url 配置错误: "
f"Provider={provider_name}, Endpoint={endpoint_short}, "
f"base_url={base_url}"
)
raise ProviderNotAvailableException(
"上游服务返回了非预期的响应格式",
provider_name=provider_name,
upstream_status=200,
upstream_response=stripped.decode("utf-8", errors="replace")[:500],
)
# 纯 JSON可能无换行/多行 JSON
if stripped.startswith(b"{") or stripped.startswith(b"["):
payload_str = stripped.decode("utf-8", errors="replace").strip()
data = json.loads(payload_str)
if isinstance(data, dict) and parser.is_error_response(data):
parsed = parser.parse_response(data, 200)
logger.warning(
f" [{request_id}] 检测到 JSON 错误响应: "
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=provider_name,
error_code=parsed.embedded_status_code,
error_message=parsed.error_message,
error_status=parsed.error_type,
)
except json.JSONDecodeError:
pass

View File

@@ -0,0 +1,160 @@
"""
Video Adapter 通用基类
负责请求分发、Handler 创建、认证头提取等通用逻辑。
"""
from __future__ import annotations
from typing import Any, ClassVar
from fastapi import HTTPException, Request
from fastapi.responses import Response
from src.api.base.adapter import ApiAdapter, ApiMode
from src.api.base.context import ApiRequestContext
from src.api.handlers.base.video_handler_base import VideoHandlerBase
from src.core.api_format import (
ApiFamily,
EndpointKind,
get_auth_handler,
get_default_auth_method_for_endpoint,
)
from src.core.logger import logger
class VideoAdapterBase(ApiAdapter):
"""视频生成适配器基类"""
FORMAT_ID: str = "UNKNOWN"
HANDLER_CLASS: type[VideoHandlerBase]
# 新架构:结构化标识(逐步替代直接依赖 FORMAT_ID 的语义)
API_FAMILY: ClassVar[ApiFamily | None] = None
ENDPOINT_KIND: ClassVar[EndpointKind] = EndpointKind.VIDEO
name: str = "video.base"
mode = ApiMode.STANDARD
eager_request_body = False
def __init__(self, allowed_api_formats: list[str] | None = None):
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
def extract_api_key(self, request: Request) -> str | None:
auth_method = get_default_auth_method_for_endpoint(self.FORMAT_ID)
handler = get_auth_handler(auth_method)
return handler.extract_credentials(request)
async def handle(self, context: ApiRequestContext) -> Response:
http_request = context.request
path_params = context.path_params or {}
if context.api_key is None or context.user is None:
raise HTTPException(status_code=401, detail="Unauthorized")
handler = self._create_handler(context)
method = http_request.method.upper()
path = http_request.url.path.lower()
task_id = path_params.get("task_id")
# Note: not every POST endpoint requires a body (e.g. cancel).
original_request_body: dict[str, Any] = {}
logger.debug(
"[VideoAdapter] dispatch method={} path={} task_id={}",
method,
path,
task_id,
)
# Download content
if method == "GET" and path.endswith("/content") and task_id:
return await handler.handle_download_content(
task_id=task_id,
http_request=http_request,
original_headers=context.original_headers,
query_params=context.query_params,
path_params=path_params,
)
# Cancel task (POST /videos/{id}/cancel or explicit action=cancel)
if (method == "POST" and path.endswith("/cancel")) or path_params.get("action") == "cancel":
if not task_id:
raise HTTPException(
status_code=400, detail="Task ID is required for cancel operation"
)
return await handler.handle_cancel_task(
task_id=task_id,
http_request=http_request,
original_headers=context.original_headers,
query_params=context.query_params,
path_params=path_params,
)
# Delete task (DELETE /videos/{id})
if method == "DELETE" and task_id:
return await handler.handle_delete_task(
task_id=task_id,
http_request=http_request,
original_headers=context.original_headers,
query_params=context.query_params,
path_params=path_params,
)
# Remix task
if method == "POST" and path.endswith("/remix") and task_id:
original_request_body = await context.ensure_json_body_async()
return await handler.handle_remix_task(
task_id=task_id,
http_request=http_request,
original_headers=context.original_headers,
original_request_body=original_request_body,
query_params=context.query_params,
path_params=path_params,
)
# Get task
if method == "GET" and task_id:
return await handler.handle_get_task(
task_id=task_id,
http_request=http_request,
original_headers=context.original_headers,
query_params=context.query_params,
path_params=path_params,
)
# List tasks
if method == "GET" and not task_id:
return await handler.handle_list_tasks(
http_request=http_request,
original_headers=context.original_headers,
query_params=context.query_params,
path_params=path_params,
)
# Create task (default)
if method in {"POST", "PUT", "PATCH"}:
original_request_body = await context.ensure_json_body_async()
return await handler.handle_create_task(
http_request=http_request,
original_headers=context.original_headers,
original_request_body=original_request_body,
query_params=context.query_params,
path_params=path_params,
)
def _create_handler(self, context: ApiRequestContext) -> VideoHandlerBase:
return self.HANDLER_CLASS(
db=context.db,
user=context.user,
api_key=context.api_key,
request_id=context.request_id,
client_ip=context.client_ip,
user_agent=context.user_agent,
start_time=context.start_time,
allowed_api_formats=self.allowed_api_formats,
)
__all__ = ["VideoAdapterBase"]

View File

@@ -0,0 +1,495 @@
"""
Video Handler 基类
定义视频生成相关操作的统一接口。
"""
from __future__ import annotations
import json
from abc import ABC, abstractmethod
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any
import httpx
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse, Response, StreamingResponse
from sqlalchemy.orm import Session
from src.config.settings import config
from src.core.api_format.conversion.internal_video import InternalVideoTask, VideoStatus
from src.core.exceptions import ProviderNotAvailableException
from src.core.logger import logger
from src.core.video_utils import (
extract_short_id_from_operation,
normalize_gemini_operation_id,
sanitize_error_message,
)
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult
from src.services.scheduling.aware_scheduler import ProviderCandidate
if TYPE_CHECKING:
from src.services.candidate.submit import SubmitOutcome
class VideoHandlerBase(ABC):
"""视频处理器基类"""
FORMAT_ID: str = ""
def __init__(
self,
db: Session,
user: User,
api_key: ApiKey,
request_id: str,
client_ip: str,
user_agent: str,
start_time: float,
allowed_api_formats: list[str] | None = None,
):
self.db = db
self.user = user
self.api_key = api_key
self.request_id = request_id
self.client_ip = client_ip
self.user_agent = user_agent
self.start_time = start_time
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
@abstractmethod
async def handle_create_task(
self,
*,
http_request: Request,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
"""创建视频任务"""
@abstractmethod
async def handle_get_task(
self,
*,
task_id: str,
http_request: Request,
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
"""获取视频任务状态"""
@abstractmethod
async def handle_list_tasks(
self,
*,
http_request: Request,
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
"""列出任务"""
@abstractmethod
async def handle_cancel_task(
self,
*,
task_id: str,
http_request: Request,
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
"""取消任务"""
async def handle_delete_task(
self,
*,
task_id: str,
http_request: Request,
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
"""删除已完成或失败的视频任务 - 可选实现"""
raise HTTPException(status_code=501, detail="Delete not supported for this provider")
async def handle_remix_task(
self,
*,
task_id: str,
http_request: Request,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> JSONResponse:
"""Remix 任务(基于已完成视频创建新视频)- 可选实现"""
raise HTTPException(status_code=501, detail="Remix not supported for this provider")
@abstractmethod
async def handle_download_content(
self,
*,
task_id: str,
http_request: Request,
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
path_params: dict[str, Any] | None = None,
) -> Response | StreamingResponse:
"""下载视频内容"""
def _build_error_response(self, response: "httpx.Response") -> JSONResponse:
"""
构建脱敏后的错误响应
子类可重写 _format_error_payload 自定义错误格式。
"""
content_type = response.headers.get("content-type", "")
if "application/json" in content_type:
try:
error_data = response.json()
if isinstance(error_data, dict) and "error" in error_data:
payload = self._format_error_payload(error_data["error"], response.status_code)
return JSONResponse(
status_code=response.status_code,
content={"error": payload},
)
except (ValueError, KeyError, TypeError):
pass
message = sanitize_error_message(response.text or "Upstream error")
fallback_payload = self._format_error_payload({"message": message}, response.status_code)
return JSONResponse(
status_code=response.status_code,
content={"error": fallback_payload},
)
async def _try_rust_sync_http_response(
self,
*,
method: str,
url: str,
headers: dict[str, str],
body: Any = None,
provider_name: str | None = None,
provider_id: str | None = None,
endpoint_id: str | None = None,
key_id: str | None = None,
provider_api_format: str | None = None,
client_api_format: str | None = None,
model_name: str | None = None,
content_type: str | None = None,
content_encoding: str | None = None,
proxy: Any = None,
tls_profile: str | None = None,
request_timeout_ms: int = 300_000,
connect_timeout_ms: int = 30_000,
pool_timeout_ms: int = 30_000,
log_label: str = "VideoRequest",
) -> httpx.Response:
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanTimeouts,
build_execution_plan_body,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
resolved_provider_name = str(provider_name or "").strip() or None
if config.execution_runtime_backend != "rust":
raise ProviderNotAvailableException(
"Video 请求仅支持 Rust executor",
provider_name=resolved_provider_name,
upstream_response=f"executor_backend={config.execution_runtime_backend}",
)
request_headers = dict(headers)
if (
body is not None
and content_type
and not any(str(key).lower() == "content-type" for key in request_headers)
):
request_headers["content-type"] = content_type
try:
plan = ExecutionPlan(
request_id=str(self.request_id or ""),
candidate_id=None,
provider_name=str(provider_name or ""),
provider_id=str(provider_id or ""),
endpoint_id=str(endpoint_id or ""),
key_id=str(key_id or ""),
method=str(method or "POST").upper(),
url=url,
headers=request_headers,
body=build_execution_plan_body(body, content_type=content_type),
stream=False,
provider_api_format=str(provider_api_format or self.FORMAT_ID),
client_api_format=str(client_api_format or self.FORMAT_ID),
model_name=str(model_name or ""),
content_type=content_type,
content_encoding=content_encoding,
proxy=proxy,
tls_profile=tls_profile,
timeouts=ExecutionPlanTimeouts(
connect_ms=connect_timeout_ms,
read_ms=request_timeout_ms,
write_ms=request_timeout_ms,
pool_ms=pool_timeout_ms,
total_ms=request_timeout_ms,
),
)
except Exception as exc:
logger.warning(
"[{}] Rust execution plan build failed request_id={} method={} url={}: {}",
log_label,
self.request_id,
method,
url,
sanitize_error_message(str(exc)),
)
raise ProviderNotAvailableException(
"Rust executor 请求计划构建失败",
provider_name=resolved_provider_name,
upstream_response=sanitize_error_message(str(exc)),
) from exc
try:
rust_result = await ExecutionRuntimeClient().execute_sync_json(plan)
except (ExecutionRuntimeClientError, httpx.HTTPError, json.JSONDecodeError) as exc:
logger.warning(
"[{}] Rust executor unavailable request_id={} method={} url={}: {}",
log_label,
self.request_id,
method,
url,
sanitize_error_message(str(exc)),
)
raise ProviderNotAvailableException(
"执行器暂时不可用,请稍后重试",
provider_name=resolved_provider_name,
upstream_response=sanitize_error_message(str(exc)),
) from exc
response_headers = dict(rust_result.headers)
if rust_result.response_json is not None:
response_headers.setdefault("content-type", "application/json")
response_body = json.dumps(rust_result.response_json, ensure_ascii=False).encode(
"utf-8"
)
elif rust_result.response_body_bytes is not None:
response_body = rust_result.response_body_bytes
else:
response_body = b""
return httpx.Response(
status_code=rust_result.status_code,
request=httpx.Request(str(method or "POST").upper(), url, headers=request_headers),
headers=response_headers,
content=response_body,
)
def _format_error_payload(self, error: dict[str, Any], status_code: int) -> dict[str, Any]:
"""
格式化错误负载,子类可重写以匹配特定 API 格式
默认返回 OpenAI 风格格式。
"""
return {
"type": error.get("type", "upstream_error"),
"message": sanitize_error_message(error.get("message", "Request failed")),
}
def _get_task(self, task_id: str) -> VideoTask:
"""通过 UUID 查找任务OpenAI Sora 风格)"""
task = (
self.db.query(VideoTask)
.filter(VideoTask.id == task_id, VideoTask.user_id == self.user.id)
.first()
)
if not task:
raise HTTPException(status_code=404, detail="Video task not found")
return task
def _get_endpoint_and_key(self, task: VideoTask) -> tuple[ProviderEndpoint, ProviderAPIKey]:
endpoint = (
self.db.query(ProviderEndpoint).filter(ProviderEndpoint.id == task.endpoint_id).first()
)
key = self.db.query(ProviderAPIKey).filter(ProviderAPIKey.id == task.key_id).first()
if not endpoint or not key:
raise HTTPException(status_code=500, detail="Provider endpoint or key not found")
return endpoint, key
def _task_to_internal(self, task: VideoTask) -> InternalVideoTask:
try:
status = VideoStatus(task.status)
except ValueError:
status = VideoStatus.PENDING
return InternalVideoTask(
id=task.id, # OpenAI Sora 使用 UUID
external_id=task.external_task_id,
status=status,
progress_percent=task.progress_percent or 0,
progress_message=task.progress_message,
video_url=task.video_url,
video_urls=task.video_urls or [],
created_at=task.created_at,
completed_at=task.completed_at,
error_code=task.error_code,
error_message=task.error_message,
extra={"model": task.model},
)
def _finalize_usage_on_submit_failure(
self,
candidate_keys: list[dict[str, Any]],
status_code: int | None,
) -> None:
"""
提交失败时结算 pending usage避免遗留 pending 状态)。
从 candidate_keys 中提取最后尝试的 provider 信息,更新 Usage 记录。
"""
from src.services.usage.service import UsageService
# 提取 provider 信息:优先取最后一个有 attempt 的候选
provider_name = "unknown"
provider_id = None
endpoint_id = None
key_id = None
for ck in reversed(candidate_keys):
if ck.get("attempt_status") or ck.get("selected"):
provider_name = ck.get("provider_name") or "unknown"
provider_id = ck.get("provider_id")
endpoint_id = ck.get("endpoint_id")
key_id = ck.get("key_id")
break
try:
# 更新 usage 状态并设置 provider 信息
UsageService.update_usage_status(
self.db,
request_id=self.request_id,
status="failed",
error_message=f"submit_failed (status_code={status_code or 'unknown'})",
provider=provider_name,
provider_id=provider_id,
provider_endpoint_id=endpoint_id,
provider_api_key_id=key_id,
status_code=status_code,
)
except Exception as exc:
logger.warning(
"Failed to finalize usage on submit failure: request_id={}, error={}",
self.request_id,
sanitize_error_message(str(exc)),
)
def _build_billing_rule_snapshot(
self, rule_lookup: BillingRuleLookupResult | None
) -> dict[str, Any]:
"""
构建 billing_rule 快照,用于冻结到视频任务的 request_metadata 中。
快照确保异步任务完成时使用创建时刻的计费规则,避免规则变更导致成本计算不一致。
"""
if not rule_lookup:
return {"status": "no_rule"}
rule = rule_lookup.rule
return {
"status": "ok",
"scope": rule_lookup.scope,
"effective_task_type": rule_lookup.effective_task_type,
"rule_id": rule.id,
"rule_name": rule.name,
"expression": rule.expression,
"variables": rule.variables,
"dimension_mappings": rule.dimension_mappings,
}
async def _submit_with_failover(
self,
*,
api_format: str,
model_name: str,
task_type: str,
submit_func: Callable[[ProviderCandidate], Awaitable["httpx.Response"]],
extract_external_task_id: Callable[[dict[str, Any]], str | None],
supported_auth_types: set[str] | None,
allow_format_conversion: bool = False,
capability_requirements: dict[str, bool] | None = None,
max_candidates: int = 10,
) -> "SubmitOutcome | JSONResponse":
"""
提交阶段故障转移(只负责拿到 external_task_id
返回:
- 成功SubmitOutcome
- 命中上游终止规则:直接返回脱敏后的 JSONResponse保留 API 格式差异)
失败时:
- 无可用候选 / 全部失败:抛 HTTPException(503)
"""
# 延迟导入,避免 handler 基类层引入过多依赖导致循环
from src.services.candidate.submit import (
AllCandidatesFailedError,
SubmitOutcome,
UpstreamClientRequestError,
)
# 统一入口:总是通过 TaskService内部可继续委托 CandidateService便于逐步内核统一
from src.services.task import TaskService
submitter: Any = TaskService(self.db)
submit_call = submitter.submit_with_failover
try:
return await submit_call(
api_format=api_format,
model_name=model_name,
affinity_key=str(self.api_key.id),
user_api_key=self.api_key,
request_id=self.request_id,
task_type=task_type,
submit_func=submit_func,
extract_external_task_id=extract_external_task_id,
supported_auth_types=supported_auth_types,
allow_format_conversion=allow_format_conversion,
capability_requirements=capability_requirements,
max_candidates=max_candidates,
)
except UpstreamClientRequestError as exc:
# 将 pending usage 结算为 failed并记录 provider 信息
self._finalize_usage_on_submit_failure(exc.candidate_keys, exc.response.status_code)
return self._build_error_response(exc.response)
except AllCandidatesFailedError as exc:
# 将 pending usage 结算为 failed
self._finalize_usage_on_submit_failure(exc.candidate_keys, exc.last_status_code)
detail = "No available provider for video generation"
if config.billing_require_rule:
detail = "No available provider with billing rule for video generation"
# 记录候选信息到日志
logger.warning(
"[VideoHandler] All candidates failed: reason={}, candidate_keys={}",
exc.reason,
exc.candidate_keys,
)
# 创建带有 candidate_keys 的 HTTPException
http_exc = HTTPException(status_code=503, detail=detail)
http_exc.candidate_keys = exc.candidate_keys # type: ignore[attr-defined]
raise http_exc
except ProviderNotAvailableException:
detail = "No available provider for video generation"
if config.billing_require_rule:
detail = "No available provider with billing rule for video generation"
raise HTTPException(status_code=503, detail=detail)
__all__ = ["VideoHandlerBase", "normalize_gemini_operation_id", "sanitize_error_message"]