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,99 @@
"""
API Handlers - 请求处理器
按 API 格式组织的 Adapter 和 Handler
- Adapter: 请求验证、格式转换、错误处理
- Handler: 业务逻辑、调用 Provider、记录用量
支持的格式:
- claude: Claude Chat API (/v1/messages)
- claude_cli: Claude CLI 透传模式
- openai: OpenAI Chat API (/v1/chat/completions)
- openai_cli: OpenAI CLI 透传模式
注意Handler 基类和具体 Handler 使用延迟导入以避免循环依赖。
"""
# Adapter 基类(不会引起循环导入,可以直接导入)
from src.api.handlers.base import (
ChatAdapterBase,
CliAdapterBase,
)
__all__ = [
# Adapter 基类
"ChatAdapterBase",
"CliAdapterBase",
# Handler 基类(延迟导入)
"ChatHandlerBase",
"CliMessageHandlerBase",
"BaseMessageHandler",
"MessageHandlerProtocol",
"MessageTelemetry",
"StreamContext",
# Claude
"ClaudeChatAdapter",
"ClaudeTokenCountAdapter",
"build_claude_adapter",
"ClaudeChatHandler",
# Claude CLI
"ClaudeCliAdapter",
"ClaudeCliMessageHandler",
# OpenAI
"OpenAIChatAdapter",
"OpenAIChatHandler",
# OpenAI CLI
"OpenAICliAdapter",
"OpenAICliMessageHandler",
]
# 延迟导入映射表
_LAZY_IMPORTS = {
# Handler 基类
"ChatHandlerBase": ("src.api.handlers.base.chat_handler_base", "ChatHandlerBase"),
"CliMessageHandlerBase": (
"src.api.handlers.base.cli_handler_base",
"CliMessageHandlerBase",
),
"StreamContext": ("src.api.handlers.base.cli_handler_base", "StreamContext"),
"BaseMessageHandler": ("src.api.handlers.base.base_handler", "BaseMessageHandler"),
"MessageHandlerProtocol": (
"src.api.handlers.base.base_handler",
"MessageHandlerProtocol",
),
"MessageTelemetry": ("src.api.handlers.base.base_handler", "MessageTelemetry"),
# Claude
"ClaudeChatAdapter": ("src.api.handlers.claude.adapter", "ClaudeChatAdapter"),
"ClaudeTokenCountAdapter": (
"src.api.handlers.claude.adapter",
"ClaudeTokenCountAdapter",
),
"build_claude_adapter": ("src.api.handlers.claude.adapter", "build_claude_adapter"),
"ClaudeChatHandler": ("src.api.handlers.claude.handler", "ClaudeChatHandler"),
# Claude CLI
"ClaudeCliAdapter": ("src.api.handlers.claude_cli.adapter", "ClaudeCliAdapter"),
"ClaudeCliMessageHandler": (
"src.api.handlers.claude_cli.handler",
"ClaudeCliMessageHandler",
),
# OpenAI
"OpenAIChatAdapter": ("src.api.handlers.openai.adapter", "OpenAIChatAdapter"),
"OpenAIChatHandler": ("src.api.handlers.openai.handler", "OpenAIChatHandler"),
# OpenAI CLI
"OpenAICliAdapter": ("src.api.handlers.openai_cli.adapter", "OpenAICliAdapter"),
"OpenAICliMessageHandler": (
"src.api.handlers.openai_cli.handler",
"OpenAICliMessageHandler",
),
}
def __getattr__(name: str) -> None:
"""延迟导入以避免循环依赖"""
if name in _LAZY_IMPORTS:
module_path, attr_name = _LAZY_IMPORTS[name]
import importlib
module = importlib.import_module(module_path)
return getattr(module, attr_name)
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")

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"]

View File

@@ -0,0 +1,26 @@
"""Claude handler package (lazy exports)."""
from __future__ import annotations
from importlib import import_module
from typing import Any
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
"ClaudeChatAdapter": (".adapter", "ClaudeChatAdapter"),
"ClaudeTokenCountAdapter": (".adapter", "ClaudeTokenCountAdapter"),
"build_claude_adapter": (".adapter", "build_claude_adapter"),
"ClaudeChatHandler": (".handler", "ClaudeChatHandler"),
}
__all__ = list(_LAZY_EXPORTS.keys())
def __getattr__(name: str) -> Any:
if name not in _LAZY_EXPORTS:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
module_name, attr_name = _LAZY_EXPORTS[name]
module = import_module(module_name, __name__)
value = getattr(module, attr_name)
globals()[name] = value
return value

View File

@@ -0,0 +1,334 @@
"""
Claude Chat Adapter - 基于 ChatAdapterBase 的 Claude Chat API 适配器
处理 /v1/messages 端点的 Claude Chat 格式请求。
"""
from __future__ import annotations
from typing import Any
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse
from src.api.base.adapter import ApiAdapter, ApiMode
from src.api.base.context import ApiRequestContext
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase, register_adapter
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.core.api_format import ApiFamily, get_header_value
from src.core.logger import logger
from src.models.claude import ClaudeMessagesRequest, ClaudeTokenCountRequest
class ClaudeCapabilityDetector:
"""Claude API 能力检测器"""
@staticmethod
def detect_from_headers(
headers: dict[str, str],
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""
从 Claude 请求头和请求体检测能力需求
检测规则:
- anthropic-beta: context-1m-xxx -> context_1m: True
- 请求体中 cache_control.ttl = "1h" -> cache_1h: True
Args:
headers: 请求头字典
request_body: 请求体(用于检测 cache_control.ttl
"""
requirements: dict[str, bool] = {}
# 使用统一的大小写不敏感获取
beta_header = get_header_value(headers, "anthropic-beta")
if beta_header and "context-1m" in beta_header.lower():
requirements["context_1m"] = True
# 从请求体检测 cache_1h
if request_body and _detect_cache_1h_in_body(request_body):
requirements["cache_1h"] = True
return requirements
def _has_cache_1h_ttl(block: dict[str, Any]) -> bool:
"""检查单个内容块是否包含 cache_control.ttl = '1h'"""
cache_control = block.get("cache_control")
if isinstance(cache_control, dict):
return cache_control.get("ttl") == "1h"
return False
def _detect_cache_1h_in_body(body: dict[str, Any]) -> bool:
"""
扫描 Claude 请求体,检测是否包含 cache_control.ttl = "1h"
检查位置:
- system[].cache_control.ttl
- messages[].content[].cache_control.ttl
- tools[].cache_control.ttl
"""
# 检查 system数组格式
system = body.get("system")
if isinstance(system, list):
for block in system:
if isinstance(block, dict) and _has_cache_1h_ttl(block):
return True
# 检查 messages
messages = body.get("messages")
if isinstance(messages, list):
for msg in messages:
if not isinstance(msg, dict):
continue
content = msg.get("content")
if isinstance(content, list):
for block in content:
if isinstance(block, dict) and _has_cache_1h_ttl(block):
return True
# 检查 tools
tools = body.get("tools")
if isinstance(tools, list):
for tool in tools:
if isinstance(tool, dict) and _has_cache_1h_ttl(tool):
return True
return False
_TOKEN_COUNTER_PLUGIN: Any = None
def _get_token_counter() -> Any:
global _TOKEN_COUNTER_PLUGIN # noqa: PLW0603
if _TOKEN_COUNTER_PLUGIN is None:
from src.plugins.token.tiktoken_counter import TiktokenCounterPlugin
_TOKEN_COUNTER_PLUGIN = TiktokenCounterPlugin(name="tiktoken")
return _TOKEN_COUNTER_PLUGIN
async def _count_text_tokens_with_fallback(text: str, model: str) -> int:
"""使用 tiktoken 插件计数,失败时回退到轻量估算。"""
if not text:
return 0
try:
plugin = _get_token_counter()
if plugin.enabled:
return await plugin.count_tokens(text, model)
except Exception as exc:
logger.debug("tiktoken token 计数失败,使用估算回退: {}", exc)
# 与旧实现保持一致:按字符估算
return max(1, len(text) // 4)
async def _count_messages_tokens_with_fallback(messages: list[dict[str, Any]], model: str) -> int:
"""按历史逻辑统计 messages token每条消息固定开销 + 内容 token"""
total = 0
for message in messages:
if not isinstance(message, dict):
continue
total += 4 # 角色与分隔符开销
content = message.get("content", "")
if isinstance(content, str):
total += await _count_text_tokens_with_fallback(content, model)
elif isinstance(content, list):
for item in content:
if isinstance(item, dict):
text = item.get("text")
if isinstance(text, str):
total += await _count_text_tokens_with_fallback(text, model)
return total
@register_adapter
class ClaudeChatAdapter(ChatAdapterBase):
"""
Claude Chat API 适配器
处理 Claude Chat 格式的请求(/v1/messages 端点,进行格式验证)。
"""
FORMAT_ID = "claude:chat"
API_FAMILY = ApiFamily.CLAUDE
name = "claude.chat"
@property
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.claude.handler import ClaudeChatHandler
return ClaudeChatHandler
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats)
logger.info(f"[{self.name}] 初始化Chat模式适配器 | API格式: {self.allowed_api_formats}")
def detect_capability_requirements(
self,
headers: dict[str, str],
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""检测 Claude 请求中隐含的能力需求"""
return ClaudeCapabilityDetector.detect_from_headers(headers, request_body)
def _validate_request_body(
self, original_request_body: dict, path_params: dict | None = None
) -> None:
"""验证请求体"""
try:
if not isinstance(original_request_body, dict):
raise ValueError("Request body must be a JSON object")
required_fields = ["model", "messages", "max_tokens"]
missing_fields = [f for f in required_fields if f not in original_request_body]
if missing_fields:
raise ValueError(f"Missing required fields: {', '.join(missing_fields)}")
request = ClaudeMessagesRequest.model_validate(
original_request_body,
strict=False,
)
except ValueError as e:
logger.error(f"请求体基本验证失败: {str(e)}")
raise HTTPException(status_code=400, detail=str(e))
except Exception as e:
logger.warning(f"Pydantic验证警告(将继续处理): {str(e)}")
request = ClaudeMessagesRequest.model_construct(
model=original_request_body.get("model"),
max_tokens=original_request_body.get("max_tokens"),
messages=original_request_body.get("messages", []),
stream=original_request_body.get("stream", False),
)
return request
def _build_audit_metadata(self, _payload: dict[str, Any], request_obj: Any) -> dict[str, Any]:
"""构建 Claude Chat 特定的审计元数据"""
role_counts: dict[str, int] = {}
for message in request_obj.messages:
role_counts[message.role] = role_counts.get(message.role, 0) + 1
return {
"action": "claude_messages",
"model": request_obj.model,
"stream": bool(request_obj.stream),
"max_tokens": request_obj.max_tokens,
"temperature": getattr(request_obj, "temperature", None),
"top_p": getattr(request_obj, "top_p", None),
"top_k": getattr(request_obj, "top_k", None),
"messages_count": len(request_obj.messages),
"message_roles": role_counts,
"stop_sequences": len(request_obj.stop_sequences or []),
"tools_count": len(request_obj.tools or []),
"system_present": bool(request_obj.system),
"metadata_present": bool(request_obj.metadata),
"thinking_enabled": bool(request_obj.thinking),
}
@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:
"""构建Claude API端点URL"""
base_url = base_url.rstrip("/")
if base_url.endswith("/v1"):
return f"{base_url}/messages"
else:
return f"{base_url}/v1/messages"
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE
def build_claude_adapter(request: Request) -> Any:
"""根据认证头构造 Chat 或 Claude Code 适配器。
- Authorization: Bearer (且无 x-api-key) -> CLI 模式
- x-api-key -> Chat 模式
"""
auth_header = request.headers.get("authorization", "")
has_bearer = auth_header.lower().startswith("bearer ")
has_api_key = bool(request.headers.get("x-api-key"))
if has_bearer and not has_api_key:
from src.api.handlers.claude_cli.adapter import ClaudeCliAdapter
return ClaudeCliAdapter()
return ClaudeChatAdapter()
class ClaudeTokenCountAdapter(ApiAdapter):
"""计算 Claude 请求 Token 数的轻量适配器。"""
name = "claude.token_count"
mode = ApiMode.STANDARD
eager_request_body = False
def extract_api_key(self, request: Request) -> str | None:
"""从请求中提取 API 密钥 (x-api-key 或 Authorization: Bearer)"""
from src.core.api_format import get_auth_handler
from src.core.api_format.enums import AuthMethod
handler = get_auth_handler(AuthMethod.API_KEY)
api_key = handler.extract_credentials(request)
if api_key:
return api_key
bearer_handler = get_auth_handler(AuthMethod.BEARER)
return bearer_handler.extract_credentials(request)
async def handle(self, context: ApiRequestContext) -> Any:
payload = await context.ensure_json_body_async()
try:
request = ClaudeTokenCountRequest.model_validate(payload, strict=False)
except Exception as e:
logger.error(f"Token count payload invalid: {e}")
raise HTTPException(status_code=400, detail="Invalid token count payload") from e
total_tokens = 0
if request.system:
if isinstance(request.system, str):
total_tokens += await _count_text_tokens_with_fallback(
request.system, request.model
)
elif isinstance(request.system, list):
for block in request.system:
if hasattr(block, "text"):
total_tokens += await _count_text_tokens_with_fallback(
block.text, request.model
)
messages_dict = [
msg.model_dump() if hasattr(msg, "model_dump") else msg for msg in request.messages
]
total_tokens += await _count_messages_tokens_with_fallback(messages_dict, request.model)
context.add_audit_metadata(
action="claude_token_count",
model=request.model,
messages_count=len(request.messages),
system_present=bool(request.system),
tools_count=len(request.tools or []),
thinking_enabled=bool(request.thinking),
input_tokens=total_tokens,
)
return JSONResponse({"input_tokens": total_tokens})
__all__ = [
"ClaudeChatAdapter",
"ClaudeTokenCountAdapter",
"build_claude_adapter",
]

View File

@@ -0,0 +1,128 @@
"""
Claude Chat Handler - 基于通用 Chat Handler 基类的简化实现
继承 ChatHandlerBase只需覆盖格式特定的方法。
代码量从原来的 ~1470 行减少到 ~120 行。
"""
from typing import Any
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.core.api_format import ApiFamily, EndpointKind
from src.core.usage_tokens import extract_cache_creation_tokens_detail
class ClaudeChatHandler(ChatHandlerBase):
"""
Claude Chat Handler - 处理 Claude Chat/CLI API 格式的请求
格式特点:
- 使用 input_tokens/output_tokens
- 支持 cache_creation_input_tokens/cache_read_input_tokens
- 请求格式ClaudeMessagesRequest
"""
FORMAT_ID = "claude:chat"
API_FAMILY = ApiFamily.CLAUDE
ENDPOINT_KIND = EndpointKind.CHAT
def extract_model_from_request(
self,
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> str:
"""
从请求中提取模型名 - Claude 格式实现
Claude API 的 model 在请求体顶级字段。
Args:
request_body: 请求体
path_params: URL 路径参数Claude 不使用)
Returns:
模型名
"""
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,
) -> dict[str, Any]:
"""
将映射后的模型名应用到请求体
Claude API 的 model 在请求体顶级字段。
Args:
request_body: 原始请求体
mapped_model: 映射后的模型名
Returns:
更新了 model 字段的请求体
"""
result = dict(request_body)
result["model"] = mapped_model
return result
async def _convert_request(self, request: Any) -> Any:
"""
将请求转换为 Claude 格式的 Pydantic 对象
注意此方法只做类型转换dict → Pydantic不做跨格式转换。
跨格式转换由调度/执行层TaskService + RequestDispatcher在选中候选后、发送请求前执行
并受全局开关和端点配置控制。
Args:
request: 原始请求对象(应已是 Claude 格式)
Returns:
ClaudeMessagesRequest 对象
"""
from src.models.claude import ClaudeMessagesRequest
# 如果已经是 Claude 格式 Pydantic 对象,直接返回
if isinstance(request, ClaudeMessagesRequest):
return request
# 如果是字典,转换为 Pydantic 对象(假设已是 Claude 格式)
if isinstance(request, dict):
return ClaudeMessagesRequest(**request)
return request
def _extract_usage(self, response: dict) -> dict[str, int]:
"""
从 Claude 响应中提取 token 使用情况
Claude 格式使用:
- input_tokens / output_tokens
- cache_creation_input_tokens / cache_read_input_tokens
- 新格式claude_cache_creation_5_m_tokens / claude_cache_creation_1_h_tokens
"""
usage = response.get("usage", {})
total, t5m, t1h = extract_cache_creation_tokens_detail(usage)
return {
"input_tokens": usage.get("input_tokens", 0),
"output_tokens": usage.get("output_tokens", 0),
"cache_creation_input_tokens": total,
"cache_read_input_tokens": usage.get("cache_read_input_tokens", 0),
"cache_creation_input_tokens_5m": t5m,
"cache_creation_input_tokens_1h": t1h,
}
def _normalize_response(self, response: dict[str, Any]) -> dict[str, Any]:
"""
规范化 Claude 响应
Args:
response: 原始响应
Returns:
规范化后的响应
"""
# 作为中转站,直接透传响应,不做标准化处理
return response

View File

@@ -0,0 +1,248 @@
"""
Claude SSE 流解析器
解析 Claude Messages API 的 Server-Sent Events 流。
"""
import json
from typing import Any
from src.core.usage_tokens import extract_cache_creation_tokens
class ClaudeStreamParser:
"""
Claude SSE 流解析器
解析 Claude Messages API 的 SSE 事件流。
事件类型:
- message_start: 消息开始,包含初始 message 对象
- content_block_start: 内容块开始
- content_block_delta: 内容块增量(文本、工具输入等)
- content_block_stop: 内容块结束
- message_delta: 消息增量,包含 stop_reason 和最终 usage
- message_stop: 消息结束
- ping: 心跳事件
- error: 错误事件
"""
# Claude SSE 事件类型
EVENT_MESSAGE_START = "message_start"
EVENT_MESSAGE_STOP = "message_stop"
EVENT_MESSAGE_DELTA = "message_delta"
EVENT_CONTENT_BLOCK_START = "content_block_start"
EVENT_CONTENT_BLOCK_STOP = "content_block_stop"
EVENT_CONTENT_BLOCK_DELTA = "content_block_delta"
EVENT_PING = "ping"
EVENT_ERROR = "error"
# Delta 类型
DELTA_TEXT = "text_delta"
DELTA_INPUT_JSON = "input_json_delta"
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
"""
解析 SSE 数据块
Args:
chunk: 原始 SSE 数据bytes 或 str
Returns:
解析后的事件列表
"""
if isinstance(chunk, bytes):
text = chunk.decode("utf-8")
else:
text = chunk
events: list[dict[str, Any]] = []
lines = text.strip().split("\n")
current_event_type: str | None = None
for line in lines:
line = line.strip()
if not line:
continue
# 解析事件类型行
if line.startswith("event: "):
current_event_type = line[7:]
continue
# 解析数据行
if line.startswith("data: "):
data_str = line[6:]
# 处理 [DONE] 标记
if data_str == "[DONE]":
events.append({"type": "__done__", "raw": "[DONE]"})
continue
try:
data = json.loads(data_str)
# 如果数据中没有 type使用事件行的类型
if "type" not in data and current_event_type:
data["type"] = current_event_type
events.append(data)
except json.JSONDecodeError:
# 无法解析的数据,跳过
pass
current_event_type = None
return events
def parse_line(self, line: str) -> dict[str, Any] | None:
"""
解析单行 SSE 数据
Args:
line: SSE 数据行(已去除 "data: " 前缀)
Returns:
解析后的事件字典,如果无法解析返回 None
"""
if not line or line == "[DONE]":
return None
try:
result = json.loads(line)
if isinstance(result, dict):
return result
return None
except json.JSONDecodeError:
return None
def is_done_event(self, event: dict[str, Any]) -> bool:
"""
判断是否为结束事件
Args:
event: 事件字典
Returns:
True 如果是结束事件
"""
event_type = event.get("type")
return event_type in (self.EVENT_MESSAGE_STOP, "__done__")
def is_error_event(self, event: dict[str, Any]) -> bool:
"""
判断是否为错误事件
Args:
event: 事件字典
Returns:
True 如果是错误事件
"""
return event.get("type") == self.EVENT_ERROR
def get_event_type(self, event: dict[str, Any]) -> str | None:
"""
获取事件类型
Args:
event: 事件字典
Returns:
事件类型字符串
"""
event_type = event.get("type")
return str(event_type) if event_type is not None else None
def extract_text_delta(self, event: dict[str, Any]) -> str | None:
"""
从 content_block_delta 事件中提取文本增量
Args:
event: 事件字典
Returns:
文本增量,如果不是文本 delta 返回 None
"""
if event.get("type") != self.EVENT_CONTENT_BLOCK_DELTA:
return None
delta = event.get("delta", {})
if delta.get("type") == self.DELTA_TEXT:
text = delta.get("text")
return str(text) if text is not None else None
return None
def extract_usage(self, event: dict[str, Any]) -> dict[str, int] | None:
"""
从事件中提取 token 使用量
Args:
event: 事件字典
Returns:
使用量字典,如果没有使用量信息返回 None
"""
event_type = event.get("type")
# message_start 事件包含初始 usage
if event_type == self.EVENT_MESSAGE_START:
message = event.get("message", {})
usage = message.get("usage", {})
if usage:
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),
}
# message_delta 事件包含最终 usage
if event_type == self.EVENT_MESSAGE_DELTA:
usage = event.get("usage", {})
if usage:
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),
}
return None
def extract_message_id(self, event: dict[str, Any]) -> str | None:
"""
从 message_start 事件中提取消息 ID
Args:
event: 事件字典
Returns:
消息 ID如果不是 message_start 返回 None
"""
if event.get("type") != self.EVENT_MESSAGE_START:
return None
message = event.get("message", {})
msg_id = message.get("id")
return str(msg_id) if msg_id is not None else None
def extract_stop_reason(self, event: dict[str, Any]) -> str | None:
"""
从 message_delta 事件中提取停止原因
Args:
event: 事件字典
Returns:
停止原因,如果没有返回 None
"""
if event.get("type") != self.EVENT_MESSAGE_DELTA:
return None
delta = event.get("delta", {})
reason = delta.get("stop_reason")
return str(reason) if reason is not None else None
__all__ = ["ClaudeStreamParser"]

View File

@@ -0,0 +1,11 @@
"""
Claude CLI 透传处理器
"""
from src.api.handlers.claude_cli.adapter import ClaudeCliAdapter
from src.api.handlers.claude_cli.handler import ClaudeCliMessageHandler
__all__ = [
"ClaudeCliAdapter",
"ClaudeCliMessageHandler",
]

View File

@@ -0,0 +1,105 @@
"""
Claude CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
继承 CliAdapterBase只需配置 FORMAT_ID 和 HANDLER_CLASS。
"""
from __future__ import annotations
from typing import Any
from src.api.handlers.base.cli_adapter_base import CliAdapterBase, register_cli_adapter
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
from src.api.handlers.claude.adapter import ClaudeCapabilityDetector
from src.config.settings import config
from src.core.api_format import ApiFamily
@register_cli_adapter
class ClaudeCliAdapter(CliAdapterBase):
"""
Claude CLI API 适配器
处理 Claude CLI 格式的请求(/v1/messages 端点,使用 Bearer 认证)。
"""
FORMAT_ID = "claude:cli"
API_FAMILY = ApiFamily.CLAUDE
name = "claude.cli"
@property
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.claude_cli.handler import ClaudeCliMessageHandler
return ClaudeCliMessageHandler
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats)
def detect_capability_requirements(
self,
headers: dict[str, str],
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""检测 Claude CLI 请求中隐含的能力需求"""
return ClaudeCapabilityDetector.detect_from_headers(headers, request_body)
def _extract_message_count(self, payload: dict[str, Any]) -> int:
"""Claude CLI 使用 messages 字段"""
messages = payload.get("messages", [])
return len(messages) if isinstance(messages, list) else 0
def _build_audit_metadata(
self,
payload: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> dict[str, Any]:
"""Claude CLI 特定的审计元数据"""
model = payload.get("model", "unknown")
stream = payload.get("stream", False)
messages = payload.get("messages", [])
role_counts = {}
for msg in messages:
role = msg.get("role", "unknown")
role_counts[role] = role_counts.get(role, 0) + 1
return {
"action": "claude_cli_request",
"model": model,
"stream": bool(stream),
"max_tokens": payload.get("max_tokens"),
"messages_count": len(messages),
"message_roles": role_counts,
"temperature": payload.get("temperature"),
"top_p": payload.get("top_p"),
"tool_count": len(payload.get("tools") or []),
"system_present": bool(payload.get("system")),
}
@classmethod
def build_endpoint_url(
cls,
base_url: str,
request_data: dict[str, Any],
model_name: str | None = None,
*,
provider_type: str | None = None,
) -> str:
"""构建Claude CLI API端点URL"""
base_url = base_url.rstrip("/")
if base_url.endswith("/v1"):
return f"{base_url}/messages"
else:
return f"{base_url}/v1/messages"
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE_CLI
@classmethod
def get_cli_user_agent(cls) -> str | None:
"""获取Claude CLI User-Agent"""
return config.internal_user_agent_claude_cli
__all__ = ["ClaudeCliAdapter"]

View File

@@ -0,0 +1,217 @@
"""
Claude CLI Message Handler - 基于通用 CLI Handler 基类的简化实现
继承 CliMessageHandlerBase只需覆盖格式特定的配置和事件处理逻辑。
"""
from typing import Any
from src.api.handlers.base.cli_handler_base import (
CliMessageHandlerBase,
StreamContext,
)
from src.core.api_format import ApiFamily, EndpointKind
from src.core.usage_tokens import extract_cache_creation_tokens_detail
class ClaudeCliMessageHandler(CliMessageHandlerBase):
"""
Claude CLI Message Handler - 处理 Claude CLI API 格式
使用新三层架构 (Provider -> ProviderEndpoint -> ProviderAPIKey)
通过 TaskService/FailoverEngine 实现自动故障转移、健康监控和并发控制
响应格式特点:
- 使用 content[] 数组
- 使用 text 类型
- 流式事件message_start, content_block_delta, message_delta, message_stop
- 支持 cache_creation_input_tokens 和 cache_read_input_tokens
模型字段:请求体顶级 model 字段
"""
FORMAT_ID = "claude:cli"
API_FAMILY = ApiFamily.CLAUDE
ENDPOINT_KIND = EndpointKind.CLI
def extract_model_from_request(
self,
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> str:
"""
从请求中提取模型名 - Claude 格式实现
Claude API 的 model 在请求体顶级字段。
Args:
request_body: 请求体
path_params: URL 路径参数Claude 不使用)
Returns:
模型名
"""
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,
) -> dict[str, Any]:
"""
Claude API 的 model 在请求体顶级
Args:
request_body: 原始请求体
mapped_model: 映射后的模型名
Returns:
更新了 model 字段的请求体
"""
result = dict(request_body)
result["model"] = mapped_model
return result
def _process_event_data(
self,
ctx: StreamContext,
event_type: str,
data: dict[str, Any],
) -> None:
"""
处理 Claude CLI 格式的 SSE 事件
事件类型:
- message_start: 消息开始,包含初始 usage含缓存 tokens
- content_block_delta: 文本增量
- message_delta: 消息增量,包含最终 usage
- message_stop: 消息结束
跨格式转换时(如 provider=openai:cli原始事件数据是 Provider 格式而非 Claude 格式。
此时委托基类方法通过 Provider 格式解析器提取 usage。
"""
# 跨格式转换时:原始事件是 Provider 格式,
# 基类 _process_event_data 会自动选择正确的 Provider 解析器提取 usage/text
if ctx.provider_api_format and ctx.provider_api_format != ctx.client_api_format:
super()._process_event_data(ctx, event_type, data)
return
# 以下是同格式claude:cli / claude:chat的处理逻辑
# 处理 message_start 事件
if event_type == "message_start":
message = data.get("message", {})
if message.get("id"):
ctx.response_id = message["id"]
# 提取初始 usage包含缓存 tokens
usage = message.get("usage", {})
if usage:
ctx.input_tokens = usage.get("input_tokens", 0)
cache_read = usage.get("cache_read_input_tokens", 0)
if cache_read:
ctx.cached_tokens = cache_read
total, t5m, t1h = extract_cache_creation_tokens_detail(usage)
if total:
ctx.cache_creation_tokens = total
ctx.cache_creation_tokens_5m = t5m
ctx.cache_creation_tokens_1h = t1h
# 处理文本增量
elif event_type == "content_block_delta":
delta = data.get("delta", {})
if delta.get("type") == "text_delta":
text = delta.get("text", "")
if text:
ctx.append_text(text)
# 处理消息增量(包含最终 usage
elif event_type == "message_delta":
usage = data.get("usage", {})
if usage:
if "input_tokens" in usage:
ctx.input_tokens = usage["input_tokens"]
if "output_tokens" in usage:
ctx.output_tokens = usage["output_tokens"]
# 更新缓存读取 tokens
if "cache_read_input_tokens" in usage:
ctx.cached_tokens = usage["cache_read_input_tokens"]
# 更新缓存创建 tokens
total, t5m, t1h = extract_cache_creation_tokens_detail(usage)
if total > 0:
ctx.cache_creation_tokens = total
ctx.cache_creation_tokens_5m = t5m
ctx.cache_creation_tokens_1h = t1h
# 检查是否结束
delta = data.get("delta", {})
if delta.get("stop_reason"):
ctx.has_completion = True
ctx.final_response = data
# 处理消息结束
elif event_type == "message_stop":
ctx.has_completion = True
def _extract_response_metadata(
self,
response: dict[str, Any],
) -> dict[str, Any]:
"""
从 Claude 响应中提取元数据
提取 model、stop_reason 等字段作为元数据。
Args:
response: Claude API 响应
Returns:
提取的元数据字典
"""
metadata: dict[str, Any] = {}
# 提取模型名称(实际使用的模型)
if "model" in response:
metadata["model"] = response["model"]
# 提取停止原因
if "stop_reason" in response:
metadata["stop_reason"] = response["stop_reason"]
# 提取消息 ID
if "id" in response:
metadata["message_id"] = response["id"]
# 提取消息类型
if "type" in response:
metadata["type"] = response["type"]
return metadata
def _finalize_stream_metadata(self, ctx: StreamContext) -> None:
"""
从流上下文中提取最终元数据
在流传输完成后调用,从收集的事件中提取元数据。
Args:
ctx: 流上下文
"""
# 从 response_id 提取消息 ID
if ctx.response_id:
ctx.response_metadata["message_id"] = ctx.response_id
# 从 final_response 提取停止原因message_delta 事件中的 delta.stop_reason
if ctx.final_response:
delta = ctx.final_response.get("delta", {})
if "stop_reason" in delta:
ctx.response_metadata["stop_reason"] = delta["stop_reason"]
# 记录模型名称
if ctx.model:
ctx.response_metadata["model"] = ctx.model

View File

@@ -0,0 +1,28 @@
"""Gemini handler package (lazy exports)."""
from __future__ import annotations
from importlib import import_module
from typing import Any
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
"GeminiChatAdapter": (".adapter", "GeminiChatAdapter"),
"build_gemini_adapter": (".adapter", "build_gemini_adapter"),
"GeminiChatHandler": (".handler", "GeminiChatHandler"),
"GeminiStreamParser": (".stream_parser", "GeminiStreamParser"),
"GeminiVeoAdapter": (".video_adapter", "GeminiVeoAdapter"),
"GeminiVeoHandler": (".video_handler", "GeminiVeoHandler"),
}
__all__ = list(_LAZY_EXPORTS.keys())
def __getattr__(name: str) -> Any:
if name not in _LAZY_EXPORTS:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
module_name, attr_name = _LAZY_EXPORTS[name]
module = import_module(module_name, __name__)
value = getattr(module, attr_name)
globals()[name] = value
return value

View File

@@ -0,0 +1,438 @@
"""
Gemini Chat Adapter
处理 Gemini API 格式的请求适配
"""
from __future__ import annotations
from typing import Any
import httpx
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase, register_adapter
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.core.api_format import ApiFamily, get_auth_handler, resolve_header_name_case
from src.core.api_format.enums import AuthMethod
from src.core.logger import logger
from src.models.gemini import GeminiRequest
from src.services.gemini_files_mapping import extract_file_names_from_request
class GeminiCapabilityDetector:
"""Gemini API 能力检测器"""
@staticmethod
def detect_from_request(
headers: dict[str, str], # noqa: ARG004 - 预留
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""
从请求体检测 Gemini 能力需求
检测规则:
- fileData.fileUri -> gemini_files: True
"""
requirements: dict[str, bool] = {}
if request_body and extract_file_names_from_request(request_body):
requirements["gemini_files"] = True
return requirements
@register_adapter
class GeminiChatAdapter(ChatAdapterBase):
"""
Gemini Chat API 适配器
处理 Gemini Chat 格式的请求
端点: /v1beta/models/{model}:generateContent
"""
FORMAT_ID = "gemini:chat"
API_FAMILY = ApiFamily.GEMINI
name = "gemini.chat"
@property
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.gemini.handler import GeminiChatHandler
return GeminiChatHandler
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats)
logger.info(
f"[{self.name}] 初始化 Gemini Chat 适配器 | API格式: {self.allowed_api_formats}"
)
def extract_api_key(self, request: Request) -> str | None:
"""
从请求中提取 API 密钥 - Gemini 支持 header 和 query 两种方式
优先级(与 Google SDK 行为一致):
1. URL 参数 ?key=
2. x-goog-api-key 请求头
"""
handler = get_auth_handler(AuthMethod.GOOG_API_KEY)
return handler.extract_credentials(request)
def detect_capability_requirements(
self,
headers: dict[str, str],
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""从请求体检测 Gemini 能力需求fileData.fileUri -> gemini_files"""
return GeminiCapabilityDetector.detect_from_request(headers, request_body)
def _merge_path_params(
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
) -> dict[str, Any]:
"""
合并 URL 路径参数到请求体 - Gemini 特化版本
Gemini API 特点:
- model 不合并到请求体(通过 extract_model_from_request 从 path_params 获取)
- stream 不合并到请求体Gemini API 通过 URL 端点区分流式/非流式)
Handler 层的 extract_model_from_request 会从 path_params 获取 model
prepare_provider_request_body 会确保发送给 Gemini API 的请求体不含 model。
Args:
original_request_body: 原始请求体字典
path_params: URL 路径参数字典(不使用)
Returns:
原始请求体(不合并任何 path_params
"""
return original_request_body.copy()
def _validate_request_body(
self, original_request_body: dict, path_params: dict | None = None
) -> None:
"""验证请求体"""
path_params = path_params or {}
is_stream = path_params.get("stream", False)
model = path_params.get("model", "unknown")
try:
if not isinstance(original_request_body, dict):
raise ValueError("Request body must be a JSON object")
# Gemini 必需字段: contents
if "contents" not in original_request_body:
raise ValueError("Missing required field: contents")
request = GeminiRequest.model_validate(
original_request_body,
strict=False,
)
except ValueError as e:
logger.error(f"请求体基本验证失败: {str(e)}")
raise HTTPException(status_code=400, detail=str(e))
except Exception as e:
logger.warning(f"Pydantic验证警告(将继续处理): {str(e)}")
request = GeminiRequest.model_construct(
contents=original_request_body.get("contents", []),
)
# 设置 model从 path_params 获取,用于日志和审计)
request.model = model
# 设置 stream 属性(用于 ChatAdapterBase 判断流式模式)
request.stream = is_stream
return request
def _extract_message_count(self, payload: dict[str, Any], request_obj: Any) -> int:
"""提取消息数量"""
contents = payload.get("contents", [])
if hasattr(request_obj, "contents"):
contents = request_obj.contents
return len(contents) if isinstance(contents, list) else 0
def _build_audit_metadata(self, payload: dict[str, Any], request_obj: Any) -> dict[str, Any]:
"""构建 Gemini Chat 特定的审计元数据"""
role_counts: dict[str, int] = {}
contents = getattr(request_obj, "contents", []) or []
for content in contents:
if isinstance(content, dict):
role = content.get("role", "unknown")
else:
role = getattr(content, "role", None) or "unknown"
role_counts[role] = role_counts.get(role, 0) + 1
generation_config = getattr(request_obj, "generation_config", None) or {}
if hasattr(generation_config, "dict"):
generation_config = generation_config.dict()
elif not isinstance(generation_config, dict):
generation_config = {}
# 判断流式模式
stream = getattr(request_obj, "stream", False)
return {
"action": "gemini_generate_content",
"model": getattr(request_obj, "model", payload.get("model", "unknown")),
"stream": bool(stream),
"max_output_tokens": generation_config.get("max_output_tokens"),
"temperature": generation_config.get("temperature"),
"top_p": generation_config.get("top_p"),
"top_k": generation_config.get("top_k"),
"contents_count": len(contents),
"content_roles": role_counts,
"tools_count": len(getattr(request_obj, "tools", None) or []),
"system_instruction_present": bool(getattr(request_obj, "system_instruction", None)),
"safety_settings_count": len(getattr(request_obj, "safety_settings", None) or []),
}
def _error_response(self, status_code: int, error_type: str, message: str) -> JSONResponse:
"""生成 Gemini 格式的错误响应"""
# Gemini 错误响应格式
return JSONResponse(
status_code=status_code,
content={
"error": {
"code": status_code,
"message": message,
"status": error_type.upper(),
}
},
)
@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:
"""构建Gemini API端点URL"""
base_url = base_url.rstrip("/")
if base_url.endswith("/v1beta"):
return base_url # 子类需要处理model参数
else:
return f"{base_url}/v1beta"
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> GEMINI
@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 上下文
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]:
"""测试 Gemini API 模型连接性(非流式)"""
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.services.provider.adapters.vertex_ai.transport import is_vertex_ai_context
# Gemini需要从request_data或model_name参数获取model名称
effective_model_name = model_name or request_data.get("model", "")
if not effective_model_name:
return {
"error": "Model name is required for Gemini API",
"status_code": 400,
}
is_antigravity = provider_type and provider_type.lower() == "antigravity"
is_gemini_cli = provider_type and provider_type.lower() == "gemini_cli"
is_vertex = is_vertex_ai_context(
base_url=base_url,
provider_type=provider_type,
endpoint=provider_endpoint,
key=provider_api_key,
)
is_oauth = auth_type == "oauth"
vertex_auth_info: Any | None = None
# Antigravity provider 使用 v1internal 路径,而非标准 Gemini API 路径
if 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)
ag_base = ordered_urls[0] if ordered_urls else base_url
path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent")
url = f"{str(ag_base).rstrip('/')}{path}"
elif is_gemini_cli:
from src.services.provider.adapters.gemini_cli.constants import V1INTERNAL_PATH_TEMPLATE
path = V1INTERNAL_PATH_TEMPLATE.format(action="generateContent")
url = f"{str(base_url).rstrip('/')}{path}"
elif is_vertex and provider_endpoint is not None and provider_api_key is not None:
# Vertex AI: test-model 必须走统一 provider transport/auth
# 否则会错误命中普通 Gemini URL导致 404
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
url = build_provider_url(
provider_endpoint,
path_params={"model": effective_model_name},
is_stream=bool(request_data.get("stream", False)),
key=provider_api_key,
decrypted_auth_config=effective_auth_config,
)
else:
# 使用基类配置方法但重写URL构建逻辑
base_url_resolved = cls.build_endpoint_url(base_url)
url = f"{base_url_resolved}/models/{effective_model_name}:generateContent"
# 构建请求组件
# Antigravity 需要特定的 User-Agent
merged_extra = dict(extra_headers) if extra_headers else {}
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_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)
# OAuth 统一处理替换端点默认认证头x-goog-api-key为 Authorization: Bearer
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 = cls.build_request_body(request_data)
# 应用请求体规则(在格式转换后应用,确保规则效果不被覆盖)
if body_rules:
body = apply_body_rules(
body,
body_rules,
original_body=body,
)
# Antigravity 需要将请求体包装为 v1internal 信封格式
if is_antigravity:
from src.services.provider.adapters.antigravity.envelope import wrap_v1internal_request
project_id = (decrypted_auth_config or {}).get("project_id", "")
body = wrap_v1internal_request(
body,
project_id=project_id,
model=effective_model_name,
request_type="endpoint_test",
)
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", "")
body = wrap_v1internal_request(
body,
project_id=project_id,
model=effective_model_name,
)
# 应用请求头规则(在请求头构建后应用)
if header_rules:
# 获取认证头名称,防止被规则覆盖
from src.core.api_format import get_auth_config_for_endpoint
auth_header, _ = get_auth_config_for_endpoint(cls.FORMAT_ID)
protected_keys = {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()
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,
)
def build_gemini_adapter(x_app_header: str = "") -> GeminiChatAdapter: # noqa: ARG001
"""
根据请求头构建适当的 Gemini 适配器
Args:
x_app_header: X-App 请求头值
Returns:
GeminiChatAdapter 实例
"""
# 目前只有一种 Gemini 适配器
# 未来可以根据 x_app_header 返回不同的适配器(如 CLI 模式)
return GeminiChatAdapter()
__all__ = ["GeminiChatAdapter", "build_gemini_adapter"]

View File

@@ -0,0 +1,198 @@
"""
Gemini Chat Handler
处理 Gemini API 格式的请求
"""
from __future__ import annotations
from typing import Any
from starlette.requests import Request
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.core.api_format import ApiFamily, EndpointKind
class GeminiChatHandler(ChatHandlerBase):
"""
Gemini Chat Handler - 处理 Google Gemini API 格式的请求
格式特点:
- 使用 promptTokenCount / candidatesTokenCount
- 支持 cachedContentTokenCount
- 请求格式: GeminiRequest
- 响应格式: JSON 数组流(非 SSE
"""
FORMAT_ID = "gemini:chat"
API_FAMILY = ApiFamily.GEMINI
ENDPOINT_KIND = EndpointKind.CHAT
async def _resolve_preferred_key_ids(
self,
model_name: str, # noqa: ARG002 - 仅做文件绑定
request_body: dict[str, Any] | None = None,
) -> list[str] | None:
"""
从 files/xxx 绑定关系中解析优先 Key ID 列表。
Gemini 文件与上传它的 API Key 绑定,必须使用同一 Key 访问。
此方法从缓存中查找文件→Key 映射,优先使用正确的 Key。
当同一源文件被上传到多个 Key 时,会返回所有可用的 Key ID
让系统能够选择任意可用的 Key。
注意事项:
- 如果映射缺失(缓存过期/重启),会记录警告,请求可能失败
- 优先返回所有支持该文件的 Key让调度器选择可用的
"""
from src.core.logger import logger
from src.services.gemini_files_mapping import (
extract_file_names_from_request,
get_all_key_ids_for_file,
)
file_names = extract_file_names_from_request(request_body or {})
if not file_names:
return None
all_key_ids: set[str] = set()
unmapped_files: list[str] = []
for file_name in file_names:
# 获取所有支持该文件的 Key包括通过 source_hash 关联的)
key_ids = await get_all_key_ids_for_file(file_name)
if key_ids:
all_key_ids.update(key_ids)
else:
unmapped_files.append(file_name)
# 警告:映射缺失
if unmapped_files:
logger.warning(
f"[{self.request_id}] Gemini 文件→Key 映射缺失: {unmapped_files}"
"请求可能失败(文件属于其他 Key 或映射已过期)"
)
if all_key_ids:
logger.debug(f"[{self.request_id}] 文件引用可用的 Key: {list(all_key_ids)}")
return list(all_key_ids) if all_key_ids else None
def extract_model_from_request(
self,
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None,
) -> str:
"""
从请求中提取模型名 - Gemini Chat 格式实现
Gemini Chat 模式下model 在请求体中(经过转换后的 GeminiRequest
与 Gemini CLI 不同CLI 模式的 model 在 URL 路径中。
Args:
request_body: 请求体
path_params: URL 路径参数Chat 模式通常不使用)
Returns:
模型名
"""
# 优先从请求体获取,其次从 path_params
model = request_body.get("model")
if model:
return str(model)
if path_params and "model" in path_params:
return str(path_params["model"])
return "unknown"
async def _convert_request(self, request: Request) -> None:
"""
将请求转换为 Gemini 格式的 Pydantic 对象
注意此方法只做类型转换dict → Pydantic不做跨格式转换。
跨格式转换由调度/执行层TaskService + RequestDispatcher在选中候选后、发送请求前执行
并受全局开关和端点配置控制。
Args:
request: 原始请求对象(应已是 Gemini 格式)
Returns:
GeminiRequest 对象
"""
from src.models.gemini import GeminiRequest
# 如果已经是 Gemini 格式 Pydantic 对象,直接返回
if isinstance(request, GeminiRequest):
return request
# 如果是字典,转换为 Pydantic 对象(假设已是 Gemini 格式)
if isinstance(request, dict):
return GeminiRequest(**request)
return request
def _extract_usage(self, response: dict) -> dict[str, int]:
"""
从 Gemini 响应中提取 token 使用情况
调用 GeminiStreamParser.extract_usage 作为单一实现源
"""
from src.api.handlers.gemini.stream_parser import GeminiStreamParser
usage = GeminiStreamParser().extract_usage(response)
if not usage:
return {
"input_tokens": 0,
"output_tokens": 0,
"cache_creation_input_tokens": 0,
"cache_read_input_tokens": 0,
}
return {
"input_tokens": usage.get("input_tokens", 0),
"output_tokens": usage.get("output_tokens", 0),
"cache_creation_input_tokens": 0, # Gemini 不区分缓存创建
"cache_read_input_tokens": usage.get("cached_tokens", 0),
}
def finalize_provider_request(
self,
request_body: dict[str, Any],
*,
mapped_model: str | None,
provider_api_format: str | None, # noqa: ARG002
) -> dict[str, Any]:
from src.api.handlers.gemini.image_gen import (
adapt_request_for_image_gen,
is_image_gen_model,
)
# Sanitize Gemini contents: strip parts without a valid data-oneof
# field and merge consecutive same-role entries. This catches cases
# missed by the normalizer (passthrough) or the antigravity envelope.
from src.core.api_format.conversion.normalizers.gemini import (
compact_gemini_contents,
)
contents = request_body.get("contents")
if isinstance(contents, list):
request_body["contents"] = compact_gemini_contents(contents)
if not is_image_gen_model(mapped_model):
return request_body
return adapt_request_for_image_gen(request_body)
def _normalize_response(self, response: dict) -> dict:
"""
规范化 Gemini 响应
Args:
response: 原始响应
Returns:
规范化后的响应
"""
# 作为中转站,直接透传响应,不做标准化处理
return response

View File

@@ -0,0 +1,41 @@
"""
Gemini 图像生成模型请求适配
- 图像生成模型不支持 tools / system_instruction需要移除
- responseModalities / responseMimeType 与 imageConfig 冲突,需要移除
"""
from typing import Any
from src.core.video_utils import is_image_gen_model
__all__ = ["is_image_gen_model", "adapt_request_for_image_gen"]
def adapt_request_for_image_gen(body: dict[str, Any]) -> dict[str, Any]:
"""为图像生成模型清理不兼容字段"""
# 移除图像生成不支持的顶层字段
for key in ("tools", "tool_config", "toolConfig", "system_instruction", "systemInstruction"):
if key in body:
body.pop(key)
# 处理 generationConfig
gc_key = "generationConfig" if "generationConfig" in body else "generation_config"
gc = body.get(gc_key)
if not isinstance(gc, dict):
gc = {}
body[gc_key] = gc
# 移除与图像生成冲突的字段
for key in (
"responseMimeType",
"response_mime_type",
"responseModalities",
"response_modalities",
):
gc.pop(key, None)
# 设置输出模态
gc["responseModalities"] = ["TEXT", "IMAGE"]
return body

View File

@@ -0,0 +1,316 @@
"""
Gemini 流解析器SSE + JSON-array 兼容)
Gemini streamGenerateContent 常见两种返回(与上游/代理实现有关):
1) `?alt=sse`SSE`data: {GenerateContentResponse}`
2) 默认JSON-array / JSON-chunks`[{...},{...},...]`,可能跨 chunk/跨行)
本解析器提供:
- parse_line(): 适用于 SSE data 行或逐行 JSON 对象
- parse_chunk(): 适用于 JSON-array/chunks可跨 chunk 拼接)
参考:
- https://ai.google.dev/gemini-api/docs/text-generation?lang=python#generate-a-text-stream
- https://generativelanguage.googleapis.com/$discovery/rest?version=v1beta
"""
import json
from typing import Any
class GeminiStreamParser:
"""
Gemini 流解析器
解析 Gemini streamGenerateContent API 的响应流。
Gemini 流式响应特点:
- 每个事件块本质上都是一个 GenerateContentResponse JSON 对象(包含 candidates、usageMetadata 等)
- 结束判定以 `candidates[].finishReason` 为准(存在且不为 FINISH_REASON_UNSPECIFIED
"""
# finishReason官方枚举值很多见 discovery这里仅保留一个明确的“未结束”哨兵
FINISH_REASON_UNSPECIFIED = "FINISH_REASON_UNSPECIFIED"
def __init__(self) -> None:
self._buffer = ""
self._in_array = False
self._brace_depth = 0
def reset(self) -> None:
"""重置解析器状态"""
self._buffer = ""
self._in_array = False
self._brace_depth = 0
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
"""
解析流式数据块
Args:
chunk: 原始数据bytes 或 str
Returns:
解析后的事件列表
"""
if isinstance(chunk, bytes):
text = chunk.decode("utf-8")
else:
text = chunk
events: list[dict[str, Any]] = []
for char in text:
if char == "[" and not self._in_array:
self._in_array = True
continue
if char == "]" and self._in_array and self._brace_depth == 0:
# 数组结束
self._in_array = False
if self._buffer.strip():
try:
obj = json.loads(self._buffer.strip().rstrip(","))
events.append(obj)
except json.JSONDecodeError:
pass
self._buffer = ""
continue
if self._in_array:
if char == "{":
self._brace_depth += 1
elif char == "}":
self._brace_depth -= 1
self._buffer += char
# 当 brace_depth 回到 0 时,说明一个完整的 JSON 对象结束
if self._brace_depth == 0 and self._buffer.strip():
try:
obj = json.loads(self._buffer.strip().rstrip(","))
events.append(obj)
self._buffer = ""
except json.JSONDecodeError:
# 可能还不完整,继续累积
pass
return events
def parse_line(self, line: str) -> dict[str, Any] | None:
"""
解析单行 JSON 数据
Args:
line: JSON 数据行
Returns:
解析后的事件字典,如果无法解析返回 None
"""
if not line or line.strip() in ["[", "]", ","]:
return None
try:
result = json.loads(line.strip().rstrip(","))
if isinstance(result, dict):
return result
return None
except json.JSONDecodeError:
return None
def is_done_event(self, event: dict[str, Any]) -> bool:
"""
判断是否为结束事件
Args:
event: 事件字典
Returns:
True 如果是结束事件
"""
candidates = event.get("candidates", [])
if not candidates:
return False
for candidate in candidates:
finish_reason = candidate.get("finishReason")
if not finish_reason:
continue
# 只要出现非 UNSPECIFIED 的 finishReason通常表示该 candidate 已结束。
# 例如STOP/MAX_TOKENS/SAFETY/RECITATION/MALFORMED_FUNCTION_CALL/...(枚举持续演进)
if str(finish_reason) != self.FINISH_REASON_UNSPECIFIED:
return True
return False
def is_error_event(self, event: dict[str, Any]) -> bool:
"""
判断是否为错误事件
检测多种 Gemini 错误格式:
1. 顶层 error: {"error": {...}}
2. chunks 内嵌套 error: {"chunks": [{"error": {...}}]}
3. candidates 内的错误状态
Args:
event: 事件字典
Returns:
True 如果是错误事件
"""
# 顶层 error
if "error" in event:
return True
# chunks 内嵌套 error (某些 Gemini 响应格式)
chunks = event.get("chunks", [])
if chunks and isinstance(chunks, list):
for chunk in chunks:
if isinstance(chunk, dict) and "error" in chunk:
return True
return False
def extract_error_info(self, event: dict[str, Any]) -> dict[str, Any] | None:
"""
从事件中提取错误信息
Args:
event: 事件字典
Returns:
错误信息字典 {"code": int, "message": str, "status": str},无错误返回 None
"""
# 顶层 error
if "error" in event:
error = event["error"]
if isinstance(error, dict):
return {
"code": error.get("code"),
"message": error.get("message", str(error)),
"status": error.get("status"),
}
return {"code": None, "message": str(error), "status": None}
# chunks 内嵌套 error
chunks = event.get("chunks", [])
if chunks and isinstance(chunks, list):
for chunk in chunks:
if isinstance(chunk, dict) and "error" in chunk:
error = chunk["error"]
if isinstance(error, dict):
return {
"code": error.get("code"),
"message": error.get("message", str(error)),
"status": error.get("status"),
}
return {"code": None, "message": str(error), "status": None}
return None
def get_finish_reason(self, event: dict[str, Any]) -> str | None:
"""
获取结束原因
Args:
event: 事件字典
Returns:
结束原因字符串
"""
candidates = event.get("candidates", [])
if candidates:
reason = candidates[0].get("finishReason")
return str(reason) if reason is not None else None
return None
def extract_text_delta(self, event: dict[str, Any]) -> str | None:
"""
从响应中提取文本内容
Args:
event: 事件字典
Returns:
文本内容,如果没有文本返回 None
"""
candidates = event.get("candidates", [])
if not candidates:
return None
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) if text_parts else None
def extract_usage(self, event: dict[str, Any]) -> dict[str, int] | None:
"""
从事件中提取 token 使用量
这是 Gemini token 提取的单一实现源,其他地方都应该调用此方法。
Args:
event: 事件字典(包含 usageMetadata
Returns:
使用量字典,如果没有完整的使用量信息返回 None
注意:
- 只有当 totalTokenCount 存在时才提取(确保是完整的 usage 数据)
- 输出 token = thoughtsTokenCount + candidatesTokenCount
"""
usage_metadata = event.get("usageMetadata", {})
if not usage_metadata or "totalTokenCount" not in usage_metadata:
return None
# 输出 token = thoughtsTokenCount + candidatesTokenCount
thoughts_tokens = usage_metadata.get("thoughtsTokenCount", 0)
candidates_tokens = usage_metadata.get("candidatesTokenCount", 0)
output_tokens = thoughts_tokens + candidates_tokens
return {
"input_tokens": usage_metadata.get("promptTokenCount", 0),
"output_tokens": output_tokens,
"total_tokens": usage_metadata.get("totalTokenCount", 0),
"cached_tokens": usage_metadata.get("cachedContentTokenCount", 0),
}
def extract_model_version(self, event: dict[str, Any]) -> str | None:
"""
从响应中提取模型版本
Args:
event: 事件字典
Returns:
模型版本,如果没有返回 None
"""
version = event.get("modelVersion")
return str(version) if version is not None else None
def extract_safety_ratings(self, event: dict[str, Any]) -> list[dict[str, Any]] | None:
"""
从响应中提取安全评级
Args:
event: 事件字典
Returns:
安全评级列表,如果没有返回 None
"""
candidates = event.get("candidates", [])
if not candidates:
return None
ratings = candidates[0].get("safetyRatings")
if isinstance(ratings, list):
return ratings
return None
__all__ = ["GeminiStreamParser"]

View File

@@ -0,0 +1,24 @@
"""
Gemini Video Adapter - 基于 VideoAdapterBase 的 Veo 适配器
"""
from __future__ import annotations
from src.api.handlers.base.video_adapter_base import VideoAdapterBase
from src.api.handlers.base.video_handler_base import VideoHandlerBase
from src.core.api_format import ApiFamily
class GeminiVeoAdapter(VideoAdapterBase):
FORMAT_ID = "gemini:video"
API_FAMILY = ApiFamily.GEMINI
name = "gemini.video"
@property
def HANDLER_CLASS(self) -> type[VideoHandlerBase]:
from src.api.handlers.gemini.video_handler import GeminiVeoHandler
return GeminiVeoHandler
__all__ = ["GeminiVeoAdapter"]

View File

@@ -0,0 +1,870 @@
"""
Gemini Video Handler - Veo 视频生成实现
"""
from __future__ import annotations
import time
from datetime import datetime, timedelta, timezone
from typing import Any, AsyncIterator
from uuid import uuid4
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse, Response, StreamingResponse
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from src.api.handlers.base.request_builder import (
apply_body_rules,
evaluate_condition,
get_provider_auth,
)
from src.api.handlers.base.video_handler_base import (
VideoHandlerBase,
normalize_gemini_operation_id,
sanitize_error_message,
)
from src.config.settings import config
from src.core.api_format import (
ApiFamily,
EndpointKind,
build_upstream_headers_for_endpoint,
get_extra_headers_from_endpoint,
make_signature_key,
)
from src.core.api_format.conversion.internal_video import (
InternalVideoRequest,
InternalVideoTask,
VideoStatus,
)
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
from src.core.api_format.conversion.registry import format_conversion_registry
from src.core.api_format.headers import HOP_BY_HOP_HEADERS
from src.core.crypto import crypto_service
from src.core.exceptions import ProviderNotAvailableException
from src.core.logger import logger
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint, User, VideoTask
from src.services.billing.rule_service import BillingRuleLookupResult, BillingRuleService
from src.services.provider.provider_context import resolve_provider_proxy
from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.usage.service import UsageService
class GeminiVeoHandler(VideoHandlerBase):
FORMAT_ID = "gemini:video"
API_FAMILY = ApiFamily.GEMINI
ENDPOINT_KIND = EndpointKind.VIDEO
DEFAULT_BASE_URL = "https://generativelanguage.googleapis.com"
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,
):
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_api_formats,
)
self._normalizer = GeminiNormalizer()
@staticmethod
def _get_request_base_url(http_request: Request) -> str:
"""从 HTTP 请求中获取基础 URL协议 + 主机)"""
# 优先使用 X-Forwarded-Proto 和 X-Forwarded-Host代理场景
proto = http_request.headers.get("x-forwarded-proto") or http_request.url.scheme
host = http_request.headers.get("x-forwarded-host") or http_request.headers.get("host")
if host:
return f"{proto}://{host}"
# 回退到 request.url
return f"{http_request.url.scheme}://{http_request.url.netloc}"
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:
# 将路径中的 model 合并到请求体再解析
model = path_params.get("model") if path_params else None
request_with_model = {**original_request_body}
if model:
request_with_model["model"] = str(model)
try:
internal_request = self._normalizer.video_request_to_internal(request_with_model)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
# 异步任务:提前创建 pending usage便于前端看到“处理中”
try:
UsageService.create_pending_usage(
db=self.db,
request_id=self.request_id,
user=self.user,
api_key=self.api_key,
model=internal_request.model,
is_stream=False,
request_type="video",
api_format=self.FORMAT_ID,
request_headers=original_headers,
request_body=original_request_body,
)
except Exception as exc:
logger.warning(
"Failed to create pending usage for video request_id={}: {}",
self.request_id,
sanitize_error_message(str(exc)),
)
# 用于跟踪是否发生了格式转换
format_conversion_info: dict[str, Any] = {
"converted": False,
"provider_format": None,
}
async def _submit(candidate: ProviderCandidate) -> Any:
upstream_key, endpoint, _key, auth_info = await self._resolve_upstream_key(candidate)
# 检测目标格式
provider_format = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
needs_conversion = provider_format.upper() != self.FORMAT_ID.upper()
format_conversion_info["provider_format"] = provider_format
format_conversion_info["converted"] = needs_conversion
# 应用端点的请求体规则
endpoint_body_rules = getattr(endpoint, "body_rules", None)
if needs_conversion and provider_format.upper().startswith("OPENAI:"):
# Gemini -> OpenAI 格式转换
converted_body = format_conversion_registry.convert_video_request(
original_request_body,
self.FORMAT_ID,
provider_format,
)
# 确保 seconds 字段为字符串类型(上游 Go 服务要求 string
if "seconds" in converted_body and converted_body["seconds"] is not None:
converted_body["seconds"] = str(converted_body["seconds"])
if endpoint_body_rules:
converted_body = apply_body_rules(
converted_body,
endpoint_body_rules,
original_body=original_request_body,
)
# 构建 OpenAI 风格的 URL
upstream_url = self._build_openai_upstream_url(endpoint.base_url)
# 构建 OpenAI 风格的请求头
headers = self._build_openai_upstream_headers(
original_headers,
upstream_key,
endpoint,
body=converted_body,
original_body=original_request_body,
)
return await self._try_rust_sync_http_response(
method="POST",
url=upstream_url,
headers=headers,
body=converted_body,
provider_name=str(candidate.provider.name),
provider_id=str(candidate.provider.id),
endpoint_id=str(endpoint.id),
key_id=str(_key.id),
provider_api_format=provider_format,
client_api_format=self.FORMAT_ID,
model_name=internal_request.model,
content_type=str(headers.get("content-type") or "").strip()
or "application/json",
log_label="GeminiVideoCreate",
)
else:
# 原始 Gemini 格式
request_body = (
original_request_body.copy() if endpoint_body_rules else original_request_body
)
if endpoint_body_rules:
request_body = apply_body_rules(
request_body,
endpoint_body_rules,
original_body=original_request_body,
)
upstream_url = self._build_upstream_url(endpoint.base_url, internal_request.model)
headers = self._build_upstream_headers(
original_headers,
upstream_key,
endpoint,
auth_info,
body=request_body,
original_body=original_request_body,
)
return await self._try_rust_sync_http_response(
method="POST",
url=upstream_url,
headers=headers,
body=request_body,
provider_name=str(candidate.provider.name),
provider_id=str(candidate.provider.id),
endpoint_id=str(endpoint.id),
key_id=str(_key.id),
provider_api_format=provider_format,
client_api_format=self.FORMAT_ID,
model_name=internal_request.model,
content_type=str(headers.get("content-type") or "").strip()
or "application/json",
log_label="GeminiVideoCreate",
)
def _extract_task_id(payload: dict[str, Any]) -> str | None:
# 根据响应格式提取 task ID
# Gemini: {"name": "operations/..."}
# OpenAI: {"id": "..."}
if "name" in payload:
value = payload.get("name")
logger.debug(
"[GeminiVeoHandler] Upstream response name={}, keys={}",
value,
list(payload.keys()) if isinstance(payload, dict) else type(payload),
)
if not value:
return None
return normalize_gemini_operation_id(str(value))
if "id" in payload:
# OpenAI 格式
return str(payload["id"])
return None
outcome_or_response = await self._submit_with_failover(
api_format=self.FORMAT_ID,
model_name=internal_request.model,
task_type="video",
submit_func=_submit,
extract_external_task_id=_extract_task_id,
supported_auth_types={"api_key", "service_account", "vertex_ai"},
allow_format_conversion=True,
max_candidates=10,
)
if isinstance(outcome_or_response, JSONResponse):
return outcome_or_response
outcome = outcome_or_response
# 冻结 billing_rule 配置(用于异步任务的成本一致性)
# 复用 _select_candidate 中已查询的结果billing_require_rule=false 时需补查
rule_lookup = outcome.rule_lookup
if rule_lookup is None:
rule_lookup = BillingRuleService.find_rule(
self.db,
provider_id=outcome.candidate.provider.id,
model_name=internal_request.model,
task_type="video",
)
billing_rule_snapshot = self._build_billing_rule_snapshot(rule_lookup)
external_task_id = outcome.external_task_id
# 如果发生了格式转换,记录转换后的请求体
converted_request_body = original_request_body
if format_conversion_info["converted"]:
try:
converted_request_body = format_conversion_registry.convert_video_request(
original_request_body,
self.FORMAT_ID,
format_conversion_info["provider_format"],
)
except Exception as e:
logger.warning(
"[GeminiVeoHandler] Failed to record converted request: {}",
sanitize_error_message(str(e)),
)
task = self._create_task_record(
external_task_id=external_task_id,
candidate=outcome.candidate,
original_request_body=original_request_body,
converted_request_body=converted_request_body,
internal_request=internal_request,
candidate_keys=outcome.candidate_keys,
original_headers=original_headers,
billing_rule_snapshot=billing_rule_snapshot,
format_converted=format_conversion_info["converted"],
)
try:
self.db.add(task)
self.db.flush() # 先 flush 检测冲突
self.db.commit()
self.db.refresh(task)
logger.debug(
"[GeminiVeoHandler] Task created: id={}, external_task_id={}",
task.id,
task.external_task_id,
)
except IntegrityError:
self.db.rollback()
raise HTTPException(status_code=409, detail="Task already exists")
# 先构建返回给客户端的响应(使用短 ID 对外暴露)
internal_task = InternalVideoTask(
id=task.short_id,
external_id=external_task_id,
status=VideoStatus.SUBMITTED,
created_at=task.created_at,
original_request=internal_request,
)
base_url = self._get_request_base_url(http_request)
response_body = self._normalizer.video_task_from_internal(internal_task, base_url=base_url)
# 提交成功后补齐 Usage 的 provider 上下文,真正结算留到轮询完成时
response_time_ms = int((time.time() - self.start_time) * 1000)
try:
# 构建发送给上游的请求头(脱敏)
upstream_request_headers = self._build_upstream_headers(
original_headers,
"", # key 不重要,只是用于记录
outcome.candidate.endpoint,
None, # auth_info
body=converted_request_body,
original_body=original_request_body,
)
UsageService.finalize_submitted(
self.db,
request_id=self.request_id,
provider_name=outcome.candidate.provider.name,
provider_id=outcome.candidate.provider.id,
provider_endpoint_id=outcome.candidate.endpoint.id,
provider_api_key_id=outcome.candidate.key.id,
response_time_ms=response_time_ms,
status_code=outcome.upstream_status_code or 200,
endpoint_api_format=make_signature_key(
str(getattr(outcome.candidate.endpoint, "api_family", "")).strip().lower(),
str(getattr(outcome.candidate.endpoint, "endpoint_kind", "")).strip().lower(),
),
provider_request_headers=upstream_request_headers,
response_headers=outcome.upstream_headers,
response_body=response_body, # 使用我们转换后的响应(包含我们的 ID
)
self.db.commit()
except Exception as exc:
logger.warning(
"Failed to finalize submitted usage for video request_id={}: {}",
self.request_id,
sanitize_error_message(str(exc)),
)
return JSONResponse(response_body)
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:
# Gemini 使用 operations/{id} 格式,需要按 external_task_id 查找
task = self._get_task_by_external_id(task_id)
# 直接从数据库返回任务状态(后台轮询服务会持续更新状态)
internal_task = self._task_to_internal(task)
base_url = self._get_request_base_url(http_request)
response_body = self._normalizer.video_task_from_internal(internal_task, base_url=base_url)
return JSONResponse(response_body)
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:
tasks = (
self.db.query(VideoTask)
.filter(VideoTask.user_id == self.user.id)
.order_by(VideoTask.created_at.desc())
.limit(100)
.all()
)
base_url = self._get_request_base_url(http_request)
items = [
self._normalizer.video_task_from_internal(self._task_to_internal(t), base_url=base_url)
for t in tasks
]
return JSONResponse({"operations": items})
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:
from src.services.task.service import TaskService
_ = (http_request, query_params, path_params) # reserved for future extensions
err_resp = await TaskService(self.db).cancel(
task_id,
user_id=str(self.user.id),
original_headers=original_headers,
)
if err_resp is not None:
return self._build_error_response(err_resp)
return JSONResponse({})
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:
task = self._get_task_by_external_id(task_id)
# 根据任务状态返回不同的错误码
if not task.video_url:
if task.status in (
VideoStatus.PENDING.value,
VideoStatus.SUBMITTED.value,
VideoStatus.QUEUED.value,
VideoStatus.PROCESSING.value,
):
# 任务仍在处理中,返回 202 Accepted
raise HTTPException(
status_code=202,
detail=f"Video is still processing (status: {task.status})",
)
if task.status == VideoStatus.FAILED.value:
raise HTTPException(
status_code=422,
detail=f"Video generation failed: {task.error_message or 'Unknown error'}",
)
# 其他状态(如 CANCELLED
raise HTTPException(status_code=404, detail="Video not available")
# 检查视频是否已过期
if task.video_expires_at:
now = datetime.now(timezone.utc)
if task.video_expires_at < now:
raise HTTPException(status_code=410, detail="Video URL has expired")
# 获取 provider 的认证信息Gemini 下载视频需要带 API Key
endpoint, key = self._get_endpoint_and_key(task)
download_headers: dict[str, str] = {}
if key.api_key:
try:
upstream_key = crypto_service.decrypt(key.api_key)
# Gemini API 使用 x-goog-api-key 头进行认证
download_headers["x-goog-api-key"] = upstream_key
# 如果是 Vertex AI需要使用 OAuth Bearer token
auth_info = await get_provider_auth(endpoint, key)
if auth_info:
download_headers.pop("x-goog-api-key", None)
download_headers[auth_info.auth_header] = auth_info.auth_value
except Exception as exc:
logger.warning(
"[VideoDownload] Failed to get auth for download task={}: {}",
task.id,
sanitize_error_message(str(exc)),
)
# 继续尝试无认证下载(某些 URL 可能是预签名的)
# 代理下载而非直接重定向,避免暴露上游存储 URL
# 使用 httpx 支持重定向Gemini 视频 URL 会重定向到实际存储位置)
return await self._try_rust_download_stream(
url=task.video_url,
headers=download_headers,
task_id=str(task.id),
endpoint=endpoint,
key=key,
model_name=str(getattr(task, "model", "") or "") or None,
)
async def _try_rust_download_stream(
self,
*,
url: str,
headers: dict[str, str],
task_id: str,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
model_name: str | None = None,
) -> Response | StreamingResponse:
import httpx
from src.services.proxy_node.resolver import (
build_proxy_url_async,
get_system_proxy_config_async,
resolve_delegate_config_async,
resolve_effective_proxy,
resolve_proxy_info_async,
)
from src.services.request.execution_runtime_plan import (
ExecutionPlan,
ExecutionPlanBody,
ExecutionPlanTimeouts,
ExecutionProxySnapshot,
)
from src.services.request.execution_runtime_client import (
ExecutionRuntimeClient,
ExecutionRuntimeClientError,
)
if config.execution_runtime_backend != "rust":
raise ProviderNotAvailableException(
"Video 下载仅支持 Rust executor",
provider_name="gemini",
upstream_response=f"executor_backend={config.execution_runtime_backend}",
)
try:
effective_proxy = resolve_effective_proxy(
resolve_provider_proxy(endpoint=endpoint, key=key),
getattr(key, "proxy", None),
)
if not effective_proxy or not effective_proxy.get("enabled", True):
effective_proxy = await get_system_proxy_config_async()
delegate_cfg = await resolve_delegate_config_async(effective_proxy)
proxy_url: str | None = None
if effective_proxy and not (delegate_cfg and delegate_cfg.get("tunnel")):
proxy_url = await build_proxy_url_async(effective_proxy)
proxy_info = await resolve_proxy_info_async(effective_proxy)
proxy_snapshot = ExecutionProxySnapshot.from_proxy_info(
proxy_info,
proxy_url=proxy_url,
mode_override="tunnel" if delegate_cfg and delegate_cfg.get("tunnel") else None,
node_id_override=(
str(delegate_cfg.get("node_id") or "").strip() or None
if delegate_cfg and delegate_cfg.get("tunnel")
else None
),
)
plan = ExecutionPlan(
request_id=str(self.request_id or ""),
candidate_id=None,
provider_name="gemini",
provider_id=str(getattr(endpoint, "provider_id", "") or ""),
endpoint_id=str(getattr(endpoint, "id", "") or ""),
key_id=str(getattr(key, "id", "") or ""),
method="GET",
url=url,
headers=dict(headers),
body=ExecutionPlanBody(),
stream=True,
provider_api_format=self.FORMAT_ID,
client_api_format=self.FORMAT_ID,
model_name=str(model_name or ""),
proxy=proxy_snapshot,
timeouts=ExecutionPlanTimeouts(
connect_ms=30_000,
read_ms=300_000,
write_ms=300_000,
pool_ms=30_000,
total_ms=None,
),
)
except Exception as exc:
logger.warning(
"[VideoDownload] Rust plan build failed task={} url={}: {}",
task_id,
url,
sanitize_error_message(str(exc)),
)
raise ProviderNotAvailableException(
"Rust executor 请求计划构建失败",
provider_name="gemini",
upstream_response=sanitize_error_message(str(exc)),
) from exc
try:
rust_stream = await ExecutionRuntimeClient().execute_stream(plan)
except (ExecutionRuntimeClientError, httpx.HTTPError, ValueError) as exc:
logger.warning(
"[VideoDownload] Rust executor unavailable task={} url={}: {}",
task_id,
url,
sanitize_error_message(str(exc)),
)
raise ProviderNotAvailableException(
"执行器暂时不可用,请稍后重试",
provider_name="gemini",
upstream_response=sanitize_error_message(str(exc)),
) from exc
safe_headers = {
k: v for k, v in rust_stream.headers.items() if k.lower() not in HOP_BY_HOP_HEADERS
}
if rust_stream.status_code >= 400:
try:
async for _ in rust_stream.byte_iterator:
pass
finally:
await rust_stream.response_ctx.__aexit__(None, None, None)
raise HTTPException(status_code=rust_stream.status_code, detail="Upstream error")
async def _iter_bytes() -> AsyncIterator[bytes]:
try:
async for chunk in rust_stream.byte_iterator:
yield chunk
finally:
await rust_stream.response_ctx.__aexit__(None, None, None)
return StreamingResponse(
_iter_bytes(),
status_code=rust_stream.status_code,
headers=safe_headers,
media_type=safe_headers.get("content-type", "video/mp4"),
)
# ------------------------------------------------------------------
# Helpers
# ------------------------------------------------------------------
async def _resolve_upstream_key(
self, candidate: ProviderCandidate
) -> tuple[str, ProviderEndpoint, ProviderAPIKey, Any | None]:
try:
upstream_key = crypto_service.decrypt(candidate.key.api_key)
except Exception as exc:
logger.error(
"Failed to decrypt provider key id={}: {}",
candidate.key.id,
sanitize_error_message(str(exc)),
)
raise HTTPException(status_code=500, detail="Failed to decrypt provider key")
auth_info = await get_provider_auth(candidate.endpoint, candidate.key)
return upstream_key, candidate.endpoint, candidate.key, auth_info
def _build_upstream_url(self, base_url: str | None, model: str) -> str:
base = (base_url or self.DEFAULT_BASE_URL).rstrip("/")
if base.endswith("/v1beta"):
base = base[: -len("/v1beta")]
return f"{base}/v1beta/models/{model}:predictLongRunning"
def _build_cancel_url(self, base_url: str | None, operation_name: str) -> str:
base = (base_url or self.DEFAULT_BASE_URL).rstrip("/")
if base.endswith("/v1beta"):
base = base[: -len("/v1beta")]
return f"{base}/v1beta/{operation_name}:cancel"
def _build_upstream_headers(
self,
original_headers: dict[str, str],
upstream_key: str,
endpoint: ProviderEndpoint,
auth_info: Any | None,
*,
body: dict[str, Any] | None = None,
original_body: dict[str, Any] | None = None,
) -> dict[str, str]:
extra_headers = get_extra_headers_from_endpoint(endpoint)
endpoint_sig = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
headers = build_upstream_headers_for_endpoint(
original_headers,
endpoint_sig,
upstream_key,
endpoint_headers=extra_headers,
header_rules=getattr(endpoint, "header_rules", None),
body=body,
original_body=original_body,
condition_evaluator=evaluate_condition,
)
if auth_info:
# 覆盖为 OAuth2 BearerVertex AI
headers.pop("x-goog-api-key", None)
headers[auth_info.auth_header] = auth_info.auth_value
return headers
def _format_error_payload(self, error: dict[str, Any], status_code: int) -> dict[str, Any]:
"""Gemini 风格错误格式"""
return {
"code": error.get("code", status_code),
"message": sanitize_error_message(error.get("message", "Request failed")),
"status": error.get("status", "BAD_GATEWAY"),
}
# ------------------------------------------------------------------
# OpenAI format conversion helpers
# ------------------------------------------------------------------
def _build_openai_upstream_url(self, base_url: str | None) -> str:
"""构建 OpenAI Sora API 的上游 URL"""
base = (base_url or "https://api.openai.com").rstrip("/")
if base.endswith("/v1"):
return f"{base}/videos"
return f"{base}/v1/videos"
def _build_openai_upstream_headers(
self,
original_headers: dict[str, str],
upstream_key: str,
endpoint: ProviderEndpoint,
*,
body: dict[str, Any] | None = None,
original_body: dict[str, Any] | None = None,
) -> dict[str, str]:
"""构建 OpenAI 格式的请求头"""
extra_headers = get_extra_headers_from_endpoint(endpoint)
endpoint_sig = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
return build_upstream_headers_for_endpoint(
original_headers,
endpoint_sig,
upstream_key,
endpoint_headers=extra_headers,
header_rules=getattr(endpoint, "header_rules", None),
body=body,
original_body=original_body,
condition_evaluator=evaluate_condition,
)
def _create_task_record(
self,
*,
external_task_id: str,
candidate: ProviderCandidate,
original_request_body: dict[str, Any],
internal_request: Any,
candidate_keys: list[dict[str, Any]] | None = None,
original_headers: dict[str, str] | None = None,
billing_rule_snapshot: dict[str, Any] | None = None,
converted_request_body: dict[str, Any] | None = None,
format_converted: bool = False,
) -> VideoTask:
now = datetime.now(timezone.utc)
# 构建请求元数据(使用追踪信息)
request_metadata = {
"candidate_keys": candidate_keys or [],
"selected_key_id": candidate.key.id,
"selected_endpoint_id": candidate.endpoint.id,
"client_ip": self.client_ip,
"user_agent": self.user_agent,
"request_id": self.request_id,
"billing_rule_snapshot": billing_rule_snapshot,
}
# 记录请求头(脱敏处理)
if original_headers:
safe_headers = {
k: v
for k, v in original_headers.items()
if k.lower() not in {"authorization", "x-api-key", "x-goog-api-key", "cookie"}
}
request_metadata["request_headers"] = safe_headers
provider_api_format = make_signature_key(
str(getattr(candidate.endpoint, "api_family", "")).strip().lower(),
str(getattr(candidate.endpoint, "endpoint_kind", "")).strip().lower(),
)
return VideoTask(
id=str(uuid4()),
request_id=self.request_id,
external_task_id=external_task_id,
user_id=self.user.id,
api_key_id=self.api_key.id,
username=self.user.username,
api_key_name=self.api_key.name,
provider_id=candidate.provider.id,
endpoint_id=candidate.endpoint.id,
key_id=candidate.key.id,
client_api_format=self.FORMAT_ID,
provider_api_format=provider_api_format,
format_converted=format_converted,
model=internal_request.model,
prompt=internal_request.prompt,
original_request_body=original_request_body,
converted_request_body=converted_request_body or original_request_body,
duration_seconds=internal_request.duration_seconds,
resolution=internal_request.resolution,
aspect_ratio=internal_request.aspect_ratio,
status=VideoStatus.SUBMITTED.value,
progress_percent=0,
poll_interval_seconds=config.video_poll_interval_seconds,
next_poll_at=now + timedelta(seconds=config.video_poll_interval_seconds),
poll_count=0,
max_poll_count=config.video_max_poll_count,
submitted_at=now,
request_metadata=request_metadata,
)
def _task_to_internal(self, task: VideoTask) -> InternalVideoTask:
"""覆盖父类方法Gemini 使用 short_id 作为对外暴露的 ID"""
try:
status = VideoStatus(task.status)
except ValueError:
status = VideoStatus.PENDING
return InternalVideoTask(
id=task.short_id, # Gemini 使用短 ID
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 _get_task_by_external_id(self, external_id: str) -> VideoTask:
"""按 short_id 查找任务(我们对外暴露的 operation 格式是 models/{model}/operations/{short_id}"""
from src.api.handlers.base.video_handler_base import extract_short_id_from_operation
short_id = extract_short_id_from_operation(external_id)
# 通过 short_id 查找任务
task = (
self.db.query(VideoTask)
.filter(
VideoTask.short_id == short_id,
VideoTask.user_id == self.user.id,
)
.first()
)
if not task:
logger.debug("[GeminiVeoHandler] Task not found: short_id={}", short_id)
raise HTTPException(status_code=404, detail="Video task not found")
return task
__all__ = ["GeminiVeoHandler"]

View File

@@ -0,0 +1,25 @@
"""Gemini CLI handler package (lazy exports)."""
from __future__ import annotations
from importlib import import_module
from typing import Any
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
"GeminiCliAdapter": (".adapter", "GeminiCliAdapter"),
"build_gemini_cli_adapter": (".adapter", "build_gemini_cli_adapter"),
"GeminiCliMessageHandler": (".handler", "GeminiCliMessageHandler"),
}
__all__ = list(_LAZY_EXPORTS.keys())
def __getattr__(name: str) -> Any:
if name not in _LAZY_EXPORTS:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
module_name, attr_name = _LAZY_EXPORTS[name]
module = import_module(module_name, __name__)
value = getattr(module, attr_name)
globals()[name] = value
return value

View File

@@ -0,0 +1,171 @@
"""
Gemini CLI Adapter - 基于通用 CLI Adapter 基类的实现
继承 CliAdapterBase处理 Gemini CLI 格式的请求。
"""
from __future__ import annotations
from typing import Any
from fastapi import Request
from src.api.handlers.base.cli_adapter_base import CliAdapterBase, register_cli_adapter
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
from src.api.handlers.gemini.adapter import GeminiCapabilityDetector
from src.config.settings import config
from src.core.api_format import ApiFamily, get_auth_handler
from src.core.api_format.enums import AuthMethod
@register_cli_adapter
class GeminiCliAdapter(CliAdapterBase):
"""
Gemini CLI API 适配器
处理 Gemini CLI 格式的请求(透传模式,最小验证)。
"""
FORMAT_ID = "gemini:cli"
API_FAMILY = ApiFamily.GEMINI
name = "gemini.cli"
@property
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.gemini_cli.handler import GeminiCliMessageHandler
return GeminiCliMessageHandler
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats)
def extract_api_key(self, request: Request) -> str | None:
"""
从请求中提取 API 密钥 - Gemini CLI 支持 header 和 query 两种方式
优先级(与 Google SDK 行为一致):
1. URL 参数 ?key=
2. x-goog-api-key 请求头
"""
handler = get_auth_handler(AuthMethod.GOOG_API_KEY)
return handler.extract_credentials(request)
def detect_capability_requirements(
self,
headers: dict[str, str],
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""从请求体检测 Gemini 能力需求fileData.fileUri -> gemini_files"""
return GeminiCapabilityDetector.detect_from_request(headers, request_body)
def _merge_path_params(
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
) -> dict[str, Any]:
"""
合并 URL 路径参数到请求体 - Gemini CLI 特化版本
Gemini API 特点:
- model 不合并到请求体Gemini 原生请求体不含 model通过 URL 路径传递)
- stream 不合并到请求体Gemini API 通过 URL 端点区分流式/非流式)
基类已经从 path_params 获取 model 和 stream 用于日志和路由判断。
Args:
original_request_body: 原始请求体字典
path_params: URL 路径参数字典(包含 model、stream 等)
Returns:
原始请求体(不合并任何 path_params
"""
# Gemini: 不合并任何 path_params 到请求体
return original_request_body.copy()
def _extract_message_count(self, payload: dict[str, Any]) -> int:
"""Gemini CLI 使用 contents 字段"""
contents = payload.get("contents", [])
return len(contents) if isinstance(contents, list) else 0
def _build_audit_metadata(
self,
payload: dict[str, Any],
path_params: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""Gemini CLI 特定的审计元数据"""
# 从 path_params 获取 modelGemini 请求体不含 model
model = path_params.get("model", "unknown") if path_params else "unknown"
contents = payload.get("contents", [])
generation_config = payload.get("generation_config", {}) or {}
role_counts: dict[str, int] = {}
for content in contents:
role = content.get("role", "unknown") if isinstance(content, dict) else "unknown"
role_counts[role] = role_counts.get(role, 0) + 1
return {
"action": "gemini_cli_request",
"model": model,
"stream": bool(payload.get("stream", False)),
"max_output_tokens": generation_config.get("max_output_tokens"),
"contents_count": len(contents),
"content_roles": role_counts,
"temperature": generation_config.get("temperature"),
"top_p": generation_config.get("top_p"),
"top_k": generation_config.get("top_k"),
"tools_count": len(payload.get("tools") or []),
"system_instruction_present": bool(payload.get("system_instruction")),
"safety_settings_count": len(payload.get("safety_settings") or []),
}
@classmethod
def build_endpoint_url(
cls,
base_url: str,
request_data: dict[str, Any],
model_name: str | None = None,
*,
provider_type: str | None = None,
) -> str:
"""构建Gemini CLI API端点URL"""
effective_model_name = model_name or request_data.get("model", "")
if not effective_model_name:
raise ValueError("Model name is required for Gemini API")
base_url = base_url.rstrip("/")
if base_url.endswith("/v1beta"):
prefix = base_url
else:
prefix = f"{base_url}/v1beta"
return f"{prefix}/models/{effective_model_name}:generateContent"
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> GEMINI_CLI
@classmethod
def get_cli_user_agent(cls) -> str | None:
"""获取Gemini CLI User-Agent"""
return config.internal_user_agent_gemini_cli
@classmethod
def get_cli_extra_headers(
cls, *, base_url: str | None = None, provider_type: str | None = None
) -> dict[str, str]:
"""获取Gemini CLI额外请求头包含 x-app: cli 标识"""
headers = super().get_cli_extra_headers(base_url=base_url, provider_type=provider_type)
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的 adapter
return headers
def build_gemini_cli_adapter(x_app_header: str = "") -> GeminiCliAdapter:
"""
构建 Gemini CLI 适配器
Args:
x_app_header: X-App 请求头值(预留扩展)
Returns:
GeminiCliAdapter 实例
"""
return GeminiCliAdapter()
__all__ = ["GeminiCliAdapter", "build_gemini_cli_adapter"]

View File

@@ -0,0 +1,291 @@
"""
Gemini CLI Message Handler - 基于通用 CLI Handler 基类的实现
继承 CliMessageHandlerBase处理 Gemini CLI API 格式的请求。
"""
from typing import Any
from src.api.handlers.base.cli_handler_base import (
CliMessageHandlerBase,
StreamContext,
)
from src.core.api_format import ApiFamily, EndpointKind
class GeminiCliMessageHandler(CliMessageHandlerBase):
"""
Gemini CLI Message Handler - 处理 Gemini CLI API 格式
使用新三层架构 (Provider -> ProviderEndpoint -> ProviderAPIKey)
通过 TaskService/FailoverEngine 实现自动故障转移、健康监控和并发控制
响应格式特点:
- Gemini 使用 JSON 数组格式流式响应(非 SSE
- 每个 chunk 包含 candidates、usageMetadata 等字段
- finish_reason: STOP, MAX_TOKENS, SAFETY, RECITATION, OTHER
- Token 使用: promptTokenCount (输入), thoughtsTokenCount + candidatesTokenCount (输出), cachedContentTokenCount (缓存)
Gemini API 特殊处理:
- model 在 URL 路径中而非请求体,如 /v1beta/models/{model}:generateContent
- 请求体中的 model 字段用于内部路由,不发送给 API
"""
FORMAT_ID = "gemini:cli"
API_FAMILY = ApiFamily.GEMINI
ENDPOINT_KIND = EndpointKind.CLI
def extract_model_from_request(
self,
request_body: dict[str, Any], # noqa: ARG002 - 基类签名要求
path_params: dict[str, Any] | None = None,
) -> str:
"""
从请求中提取模型名 - Gemini 格式实现
Gemini API 的 model 在 URL 路径中而非请求体:
/v1beta/models/{model}:generateContent
Args:
request_body: 请求体Gemini 不包含 model
path_params: URL 路径参数(包含 model
Returns:
模型名,如果无法提取则返回 "unknown"
"""
# Gemini: model 从 URL 路径参数获取
if path_params and "model" in path_params:
return str(path_params["model"])
return "unknown"
def prepare_provider_request_body(
self,
request_body: dict[str, Any],
) -> dict[str, Any]:
"""
准备发送给 Gemini API 的请求体 - 移除 model 字段
Gemini API 要求 model 只在 URL 路径中,请求体中的 model 字段
会导致某些代理返回 404 错误。
Args:
request_body: 请求体
Returns:
不含 model 字段的请求体
"""
result = dict(request_body)
result.pop("model", None)
return result
def finalize_provider_request(
self,
request_body: dict[str, Any],
*,
mapped_model: str | None,
provider_api_format: str | None, # noqa: ARG002
) -> dict[str, Any]:
from src.api.handlers.gemini.image_gen import (
adapt_request_for_image_gen,
is_image_gen_model,
)
# Sanitize Gemini contents: strip parts without a valid data-oneof
# field and merge consecutive same-role entries. This catches cases
# missed by the normalizer (passthrough) or the antigravity envelope.
from src.core.api_format.conversion.normalizers.gemini import (
compact_gemini_contents,
)
contents = request_body.get("contents")
if isinstance(contents, list):
request_body["contents"] = compact_gemini_contents(contents)
if not is_image_gen_model(mapped_model):
return request_body
return adapt_request_for_image_gen(request_body)
def get_model_for_url(
self,
request_body: dict[str, Any],
mapped_model: str | None,
) -> str | None:
"""
Gemini 需要将 model 放入 URL 路径中
Args:
request_body: 请求体
mapped_model: 映射后的模型名(如果有)
Returns:
用于 URL 路径的模型名
"""
# 优先使用映射后的模型名,否则使用请求体中的
return mapped_model or request_body.get("model")
def _extract_usage_from_event(
self,
event: dict[str, Any],
*,
provider_type: str | None = None,
) -> dict[str, int]:
"""
从 Gemini 事件中提取 token 使用情况
调用 GeminiStreamParser.extract_usage 作为单一实现源
Args:
event: Gemini 流式响应事件
provider_type: Provider 类型(用于 Antigravity 特判)
Returns:
包含 input_tokens, output_tokens, cached_tokens 的字典
"""
from src.core.provider_types import ProviderType
if str(provider_type or "").lower() == ProviderType.ANTIGRAVITY:
return self._extract_antigravity_usage(event)
from src.api.handlers.gemini.stream_parser import GeminiStreamParser
usage = GeminiStreamParser().extract_usage(event)
if not usage:
return {
"input_tokens": 0,
"output_tokens": 0,
"cached_tokens": 0,
}
return {
"input_tokens": usage.get("input_tokens", 0),
"output_tokens": usage.get("output_tokens", 0),
"cached_tokens": usage.get("cached_tokens", 0),
}
def _extract_antigravity_usage(self, event: dict[str, Any]) -> dict[str, int]:
"""Antigravity 专用 usage 提取(宽松 + 边界保护)。
Antigravity 的 usageMetadata 可能缺少 totalTokenCount因此不能依赖
GeminiStreamParser.extract_usage 的“totalTokenCount 必须存在”的严格判断。
"""
usage_metadata = event.get("usageMetadata", {})
if not isinstance(usage_metadata, dict) or not usage_metadata:
return {"input_tokens": 0, "output_tokens": 0, "cached_tokens": 0}
def _as_int(v: Any) -> int:
try:
return int(v or 0)
except Exception:
return 0
prompt = _as_int(usage_metadata.get("promptTokenCount"))
cached = _as_int(usage_metadata.get("cachedContentTokenCount"))
candidates = _as_int(usage_metadata.get("candidatesTokenCount"))
thoughts = _as_int(usage_metadata.get("thoughtsTokenCount"))
return {
# 注意:计费层会根据 api_family(GEMINI) 扣除 cache_read_tokens
# 因此这里保持 Gemini 口径input_tokens=promptTokenCount含缓存
"input_tokens": max(0, prompt),
"output_tokens": max(0, candidates + thoughts),
"cached_tokens": max(0, cached),
}
def _process_event_data(
self,
ctx: StreamContext,
_event_type: str,
data: dict[str, Any],
) -> None:
"""
处理 Gemini CLI 格式的流式事件
Gemini 的流式响应是 JSON 数组格式,每个元素结构如下:
{
"candidates": [{
"content": {"parts": [{"text": "..."}], "role": "model"},
"finishReason": "STOP",
"safetyRatings": [...]
}],
"usageMetadata": {
"promptTokenCount": 10,
"candidatesTokenCount": 20,
"totalTokenCount": 30,
"cachedContentTokenCount": 5
},
"modelVersion": "gemini-1.5-pro"
}
注意: Gemini 流解析器会将每个 JSON 对象作为一个"事件"传递
event_type 在这里可能为空或是自定义的标记
跨格式转换时(如 provider=claude:chat原始事件数据是 Provider 格式而非 Gemini 格式。
此时委托基类方法通过 Provider 格式解析器提取 usage。
"""
# 跨格式转换时:原始事件是 Provider 格式,
# 基类 _process_event_data 会自动选择正确的 Provider 解析器提取 usage/text
if ctx.provider_api_format and ctx.provider_api_format != ctx.client_api_format:
super()._process_event_data(ctx, _event_type, data)
return
# 以下是同格式gemini:cli / gemini:chat的处理逻辑
# 提取候选响应
candidates = data.get("candidates", [])
if candidates:
candidate = candidates[0]
content = candidate.get("content", {})
# 提取文本内容
parts = content.get("parts", [])
for part in parts:
if "text" in part:
ctx.append_text(part["text"])
# 检查结束原因
finish_reason = candidate.get("finishReason")
if finish_reason in ("STOP", "MAX_TOKENS", "SAFETY", "RECITATION", "OTHER"):
ctx.has_completion = True
ctx.final_response = data
# 提取使用量信息(复用 GeminiStreamParser.extract_usage
usage = self._extract_usage_from_event(data, provider_type=ctx.provider_type)
if usage["input_tokens"] > 0 or usage["output_tokens"] > 0:
ctx.input_tokens = usage["input_tokens"]
ctx.output_tokens = usage["output_tokens"]
ctx.cached_tokens = usage["cached_tokens"]
# 提取模型版本作为响应 ID
model_version = data.get("modelVersion")
if model_version:
if not ctx.response_id:
ctx.response_id = f"gemini-{model_version}"
# 存储到 response_metadata 供 Usage 记录使用
ctx.response_metadata["model_version"] = model_version
# 检查错误
if "error" in data:
ctx.has_completion = True
ctx.final_response = data
def _extract_response_metadata(
self,
response: dict[str, Any],
) -> dict[str, Any]:
"""
从 Gemini 响应中提取元数据
提取 modelVersion 字段,记录实际使用的模型版本。
Args:
response: Gemini API 响应
Returns:
包含 model_version 的元数据字典
"""
metadata: dict[str, Any] = {}
model_version = response.get("modelVersion")
if model_version:
metadata["model_version"] = model_version
return metadata

View File

@@ -0,0 +1,26 @@
"""OpenAI handler package (lazy exports)."""
from __future__ import annotations
from importlib import import_module
from typing import Any
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
"OpenAIChatAdapter": (".adapter", "OpenAIChatAdapter"),
"OpenAIChatHandler": (".handler", "OpenAIChatHandler"),
"OpenAIVideoAdapter": (".video_adapter", "OpenAIVideoAdapter"),
"OpenAIVideoHandler": (".video_handler", "OpenAIVideoHandler"),
}
__all__ = list(_LAZY_EXPORTS.keys())
def __getattr__(name: str) -> Any:
if name not in _LAZY_EXPORTS:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
module_name, attr_name = _LAZY_EXPORTS[name]
module = import_module(module_name, __name__)
value = getattr(module, attr_name)
globals()[name] = value
return value

View File

@@ -0,0 +1,123 @@
"""
OpenAI Chat Adapter - 基于 ChatAdapterBase 的 OpenAI Chat API 适配器
处理 /v1/chat/completions 端点的 OpenAI Chat 格式请求。
"""
from __future__ import annotations
from typing import Any
from fastapi.responses import JSONResponse
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase, register_adapter
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.core.api_format import ApiFamily
from src.core.logger import logger
from src.models.openai import OpenAIRequest
@register_adapter
class OpenAIChatAdapter(ChatAdapterBase):
"""
OpenAI Chat Completions API 适配器
处理 OpenAI Chat 格式的请求(/v1/chat/completions 端点)。
"""
FORMAT_ID = "openai:chat"
API_FAMILY = ApiFamily.OPENAI
name = "openai.chat"
@property
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.openai.handler import OpenAIChatHandler
return OpenAIChatHandler
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats)
def _validate_request_body(
self, original_request_body: dict, path_params: dict | None = None
) -> None:
"""验证请求体"""
if not isinstance(original_request_body, dict):
return self._error_response(
400, "Request body must be a JSON object", "invalid_request_error"
)
required_fields = ["model", "messages"]
missing = [f for f in required_fields if f not in original_request_body]
if missing:
return self._error_response(
400,
f"Missing required fields: {', '.join(missing)}",
"invalid_request_error",
)
try:
return OpenAIRequest.model_validate(original_request_body, strict=False)
except ValueError as e:
return self._error_response(400, str(e), "invalid_request_error")
except Exception as e:
logger.warning(f"Pydantic验证警告(将继续处理): {str(e)}")
return OpenAIRequest.model_construct(
model=original_request_body.get("model"),
messages=original_request_body.get("messages", []),
stream=original_request_body.get("stream", False),
max_tokens=original_request_body.get("max_tokens"),
)
def _build_audit_metadata(self, payload: dict[str, Any], request_obj: Any) -> dict[str, Any]:
"""构建 OpenAI Chat 特定的审计元数据"""
role_counts = {}
for message in request_obj.messages:
role_counts[message.role] = role_counts.get(message.role, 0) + 1
return {
"action": "openai_chat_completion",
"model": request_obj.model,
"stream": bool(request_obj.stream),
"max_tokens": request_obj.max_tokens,
"temperature": request_obj.temperature,
"top_p": request_obj.top_p,
"messages_count": len(request_obj.messages),
"message_roles": role_counts,
"tools_count": len(request_obj.tools or []),
"response_format": bool(request_obj.response_format),
"user_identifier": request_obj.user,
}
def _error_response(self, status_code: int, message: str, error_type: str) -> JSONResponse:
"""生成 OpenAI 格式的错误响应"""
return JSONResponse(
status_code=status_code,
content={
"error": {
"type": error_type,
"message": message,
"code": status_code,
}
},
)
@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:
"""构建OpenAI API端点URL"""
base_url = base_url.rstrip("/")
if base_url.endswith("/v1"):
return f"{base_url}/chat/completions"
else:
return f"{base_url}/v1/chat/completions"
__all__ = ["OpenAIChatAdapter"]

View File

@@ -0,0 +1,127 @@
"""
OpenAI Chat Handler - 基于通用 Chat Handler 基类的简化实现
继承 ChatHandlerBase只需覆盖格式特定的方法。
代码量从原来的 ~1315 行减少到 ~100 行。
"""
from __future__ import annotations
from typing import Any
from starlette.requests import Request
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.core.api_format import ApiFamily, EndpointKind
class OpenAIChatHandler(ChatHandlerBase):
"""
OpenAI Chat Handler - 处理 OpenAI Chat Completions API 格式的请求
格式特点:
- 使用 prompt_tokens/completion_tokens
- 不支持 cache tokens
- 请求格式OpenAIRequest
"""
FORMAT_ID = "openai:chat"
API_FAMILY = ApiFamily.OPENAI
ENDPOINT_KIND = EndpointKind.CHAT
def extract_model_from_request(
self,
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> str:
"""
从请求中提取模型名 - OpenAI 格式实现
OpenAI API 的 model 在请求体顶级字段。
Args:
request_body: 请求体
path_params: URL 路径参数OpenAI 不使用)
Returns:
模型名
"""
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,
) -> dict[str, Any]:
"""
将映射后的模型名应用到请求体
OpenAI API 的 model 在请求体顶级字段。
Args:
request_body: 原始请求体
mapped_model: 映射后的模型名
Returns:
更新了 model 字段的请求体
"""
result = dict(request_body)
result["model"] = mapped_model
return result
async def _convert_request(self, request: Request) -> None:
"""
将请求转换为 OpenAI 格式的 Pydantic 对象
注意此方法只做类型转换dict → Pydantic不做跨格式转换。
跨格式转换由调度/执行层TaskService + RequestDispatcher在选中候选后、发送请求前执行
并受全局开关和端点配置控制。
Args:
request: 原始请求对象(应已是 OpenAI 格式)
Returns:
OpenAIRequest 对象
"""
from src.models.openai import OpenAIRequest
# 如果已经是 OpenAI 格式 Pydantic 对象,直接返回
if isinstance(request, OpenAIRequest):
return request
# 如果是字典,转换为 Pydantic 对象(假设已是 OpenAI 格式)
if isinstance(request, dict):
return OpenAIRequest(**request)
return request
def _extract_usage(self, response: dict) -> dict[str, int]:
"""
从 OpenAI 响应中提取 token 使用情况
OpenAI 格式使用:
- prompt_tokens / completion_tokens
- 不支持 cache tokens
"""
usage = response.get("usage", {})
return {
"input_tokens": usage.get("prompt_tokens", 0),
"output_tokens": usage.get("completion_tokens", 0),
"cache_creation_input_tokens": 0,
"cache_read_input_tokens": 0,
}
def _normalize_response(self, response: dict) -> dict:
"""
规范化 OpenAI 响应
Args:
response: 原始响应
Returns:
规范化后的响应
"""
# 作为中转站,直接透传响应,不做标准化处理
return response

View File

@@ -0,0 +1,187 @@
"""
OpenAI SSE 流解析器
解析 OpenAI Chat Completions API 的 Server-Sent Events 流。
"""
import json
from typing import Any
class OpenAIStreamParser:
"""
OpenAI SSE 流解析器
解析 OpenAI Chat Completions API 的 SSE 事件流。
OpenAI 流格式:
- 每个 chunk 是一个 JSON 对象,包含 choices 数组
- choices[0].delta 包含增量内容
- choices[0].finish_reason 表示结束原因
- 流结束时发送 data: [DONE]
"""
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
"""
解析 SSE 数据块
Args:
chunk: 原始 SSE 数据bytes 或 str
Returns:
解析后的 chunk 列表
"""
if isinstance(chunk, bytes):
text = chunk.decode("utf-8")
else:
text = chunk
chunks: list[dict[str, Any]] = []
lines = text.strip().split("\n")
for line in lines:
line = line.strip()
if not line:
continue
# 解析数据行
if line.startswith("data: "):
data_str = line[6:]
# 处理 [DONE] 标记
if data_str == "[DONE]":
chunks.append({"__done__": True})
continue
try:
data = json.loads(data_str)
chunks.append(data)
except json.JSONDecodeError:
# 无法解析的数据,跳过
pass
return chunks
def parse_line(self, line: str) -> dict[str, Any] | None:
"""
解析单行 SSE 数据
Args:
line: SSE 数据行(已去除 "data: " 前缀)
Returns:
解析后的 chunk 字典,如果无法解析返回 None
"""
if not line or line == "[DONE]":
return None
try:
result = json.loads(line)
if isinstance(result, dict):
return result
return None
except json.JSONDecodeError:
return None
def is_done_chunk(self, chunk: dict[str, Any]) -> bool:
"""
判断是否为结束 chunk
Args:
chunk: chunk 字典
Returns:
True 如果是结束 chunk
"""
# 内部标记
if chunk.get("__done__"):
return True
# 检查 finish_reason
choices = chunk.get("choices", [])
if choices:
finish_reason = choices[0].get("finish_reason")
return finish_reason is not None
return False
def get_finish_reason(self, chunk: dict[str, Any]) -> str | None:
"""
获取结束原因
Args:
chunk: chunk 字典
Returns:
结束原因字符串
"""
choices = chunk.get("choices", [])
if choices:
reason = choices[0].get("finish_reason")
return str(reason) if reason is not None else None
return None
def extract_text_delta(self, chunk: dict[str, Any]) -> str | None:
"""
从 chunk 中提取文本增量
Args:
chunk: chunk 字典
Returns:
文本增量,如果没有返回 None
"""
choices = chunk.get("choices", [])
if not choices:
return None
delta = choices[0].get("delta", {})
content = delta.get("content")
if isinstance(content, str):
return content
return None
def extract_tool_calls_delta(self, chunk: dict[str, Any]) -> list[dict[str, Any]] | None:
"""
从 chunk 中提取工具调用增量
Args:
chunk: chunk 字典
Returns:
工具调用列表,如果没有返回 None
"""
choices = chunk.get("choices", [])
if not choices:
return None
delta = choices[0].get("delta", {})
tool_calls = delta.get("tool_calls")
if isinstance(tool_calls, list):
return tool_calls
return None
def extract_role(self, chunk: dict[str, Any]) -> str | None:
"""
从 chunk 中提取角色
通常只在第一个 chunk 中出现。
Args:
chunk: chunk 字典
Returns:
角色字符串
"""
choices = chunk.get("choices", [])
if not choices:
return None
delta = choices[0].get("delta", {})
role = delta.get("role")
return str(role) if role is not None else None
__all__ = ["OpenAIStreamParser"]

View File

@@ -0,0 +1,24 @@
"""
OpenAI Video Adapter - 基于 VideoAdapterBase 的 Sora 适配器
"""
from __future__ import annotations
from src.api.handlers.base.video_adapter_base import VideoAdapterBase
from src.api.handlers.base.video_handler_base import VideoHandlerBase
from src.core.api_format import ApiFamily
class OpenAIVideoAdapter(VideoAdapterBase):
FORMAT_ID = "openai:video"
API_FAMILY = ApiFamily.OPENAI
name = "openai.video"
@property
def HANDLER_CLASS(self) -> type[VideoHandlerBase]:
from src.api.handlers.openai.video_handler import OpenAIVideoHandler
return OpenAIVideoHandler
__all__ = ["OpenAIVideoAdapter"]

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,25 @@
"""OpenAI CLI handler package (lazy exports)."""
from __future__ import annotations
from importlib import import_module
from typing import Any
_LAZY_EXPORTS: dict[str, tuple[str, str]] = {
"OpenAICliAdapter": (".adapter", "OpenAICliAdapter"),
"OpenAICompactAdapter": (".adapter", "OpenAICompactAdapter"),
"OpenAICliMessageHandler": (".handler", "OpenAICliMessageHandler"),
}
__all__ = list(_LAZY_EXPORTS.keys())
def __getattr__(name: str) -> Any:
if name not in _LAZY_EXPORTS:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
module_name, attr_name = _LAZY_EXPORTS[name]
module = import_module(module_name, __name__)
value = getattr(module, attr_name)
globals()[name] = value
return value

View File

@@ -0,0 +1,138 @@
"""
OpenAI CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
继承 CliAdapterBase只需配置 FORMAT_ID 和 HANDLER_CLASS。
"""
from __future__ import annotations
from typing import Any
from src.api.base.context import ApiRequestContext
from src.api.handlers.base.cli_adapter_base import CliAdapterBase, register_cli_adapter
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
from src.config.settings import config
from src.core.api_format import ApiFamily, EndpointKind
from src.core.provider_types import ProviderType
from src.utils.url_utils import is_codex_url
@register_cli_adapter
class OpenAICliAdapter(CliAdapterBase):
"""
OpenAI CLI API 适配器
处理 /v1/responses 端点的请求。
"""
FORMAT_ID = "openai:cli"
API_FAMILY = ApiFamily.OPENAI
name = "openai.cli"
@property
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.openai_cli.handler import OpenAICliMessageHandler
return OpenAICliMessageHandler
def __init__(
self,
allowed_api_formats: list[str] | None = None,
*,
compact: bool = False,
):
super().__init__(allowed_api_formats)
self._compact = compact
async def handle(self, context: ApiRequestContext) -> Any:
"""处理 CLI API 请求。"""
if self._compact:
from src.services.provider.adapters.codex.context import (
CodexRequestContext,
set_codex_request_context,
)
# Keep compact routing state out of the request body. Transport/policy layers
# read this request-scoped flag directly when legacy compact fallback is needed.
set_codex_request_context(CodexRequestContext(is_compact=True))
body = await context.ensure_json_body_async()
# compact 端点永远非流式
body.pop("stream", None)
return await super().handle(context)
@classmethod
def build_endpoint_url(
cls,
base_url: str,
request_data: dict[str, Any],
model_name: str | None = None,
*,
compact: bool = False,
provider_type: str | None = None,
) -> str:
"""构建OpenAI CLI API端点URL使用 Responses API
对于 Codex OAuth 端点(如 chatgpt.com/backend-api/codex直接追加 /responses
对于标准 OpenAI API使用 /v1/responses。
compact=True 时追加 /compact 后缀。
provider_type 优先:仅当 provider_type 为 codex 时才使用 Codex 路由规则;
未传入 provider_type 时回退到 URL 模式匹配(兼容旧调用方)。
"""
suffix = "/responses/compact" if compact else "/responses"
base_url = base_url.rstrip("/")
# 判断是否按 Codex 规则构建 URL
is_codex = (
(provider_type or "").lower() == ProviderType.CODEX
if provider_type
else is_codex_url(base_url)
)
if is_codex:
return f"{base_url}{suffix}"
# 标准 OpenAI API
if base_url.endswith("/v1"):
return f"{base_url}{suffix}"
else:
return f"{base_url}/v1{suffix}"
# build_request_body 使用基类实现
# OpenAI CLI normalizer 会自动添加 instructions 字段
@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
del base_url, provider_type
return build_test_request_body(cls.FORMAT_ID, request_data)
@classmethod
def get_cli_user_agent(cls) -> str | None:
"""获取OpenAI CLI User-Agent"""
return config.internal_user_agent_openai_cli
__all__ = ["OpenAICliAdapter"]
@register_cli_adapter
class OpenAICompactAdapter(OpenAICliAdapter):
"""OpenAI Compact Responses adapter (/v1/responses/compact)."""
FORMAT_ID = "openai:compact"
ENDPOINT_KIND = EndpointKind.COMPACT
name = "openai.compact"
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats=allowed_api_formats, compact=True)
__all__.append("OpenAICompactAdapter")

View File

@@ -0,0 +1,226 @@
"""
OpenAI CLI Message Handler - 基于通用 CLI Handler 基类的简化实现
继承 CliMessageHandlerBase只需覆盖格式特定的配置和事件处理逻辑。
代码量从原来的 900+ 行减少到 ~100 行。
"""
from typing import Any
from src.api.handlers.base.cli_handler_base import (
CliMessageHandlerBase,
StreamContext,
)
from src.core.api_format import ApiFamily, EndpointKind
class OpenAICliMessageHandler(CliMessageHandlerBase):
"""
OpenAI CLI Message Handler - 处理 OpenAI CLI Responses API 格式
使用新三层架构 (Provider -> ProviderEndpoint -> ProviderAPIKey)
通过 TaskService/FailoverEngine 实现自动故障转移、健康监控和并发控制
响应格式特点:
- 使用 output[] 数组而非 content[]
- 使用 output_text 类型而非普通 text
- 流式事件response.output_text.delta, response.output_text.done
模型字段:请求体顶级 model 字段
"""
FORMAT_ID = "openai:cli"
API_FAMILY = ApiFamily.OPENAI
ENDPOINT_KIND = EndpointKind.CLI
def extract_model_from_request(
self,
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> str:
"""
从请求中提取模型名 - OpenAI 格式实现
OpenAI API 的 model 在请求体顶级字段。
Args:
request_body: 请求体
path_params: URL 路径参数OpenAI 不使用)
Returns:
模型名
"""
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,
) -> dict[str, Any]:
"""
OpenAI CLI (Responses API) 的 model 在请求体顶级字段。
Args:
request_body: 原始请求体
mapped_model: 映射后的模型名
Returns:
更新了 model 字段的请求体
"""
result = dict(request_body)
result["model"] = mapped_model
return result
def _process_event_data(
self,
ctx: StreamContext,
event_type: str,
data: dict[str, Any],
) -> None:
"""
处理 OpenAI CLI 格式的 SSE 事件
事件类型:
- response.output_text.delta: 文本增量
- response.completed: 响应完成(包含 usage
跨格式转换时(如 provider=claude:chat原始事件数据是 Provider 格式而非 OpenAI CLI 格式。
此时先调用基类方法通过 Provider 格式解析器提取 usage再执行 OpenAI CLI 特定的处理逻辑。
"""
# 跨格式转换时:原始事件是 Provider 格式(如 Claude
# 基类 _process_event_data 会自动选择正确的 Provider 解析器提取 usage/text
if ctx.provider_api_format and ctx.provider_api_format != ctx.client_api_format:
super()._process_event_data(ctx, event_type, data)
return
# 以下是同格式openai:cli的处理逻辑
# 提取 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"]
# 处理文本增量
if event_type in ["response.output_text.delta", "response.outtext.delta"]:
delta = data.get("delta")
if isinstance(delta, str):
ctx.append_text(delta)
elif isinstance(delta, dict) and "text" in delta:
ctx.append_text(delta["text"])
# 处理完成事件
elif event_type == "response.completed":
ctx.has_completion = True
response_obj = data.get("response")
if isinstance(response_obj, dict):
ctx.final_response = response_obj
usage_obj = response_obj.get("usage")
if isinstance(usage_obj, dict):
ctx.final_usage = usage_obj
ctx.input_tokens = usage_obj.get("input_tokens", 0)
ctx.output_tokens = usage_obj.get("output_tokens", 0)
details = usage_obj.get("input_tokens_details")
if isinstance(details, dict):
ctx.cached_tokens = details.get("cached_tokens", 0)
# 如果没有收集到文本,从 output 中提取
if not ctx.collected_text and "output" in response_obj:
for output_item in response_obj.get("output", []):
if output_item.get("type") != "message":
continue
for content_item in output_item.get("content", []):
if content_item.get("type") == "output_text":
text = content_item.get("text", "")
if text:
ctx.append_text(text)
# 备用:从顶层 usage 提取
usage_obj = data.get("usage")
if isinstance(usage_obj, dict) and not ctx.final_usage:
ctx.final_usage = usage_obj
ctx.input_tokens = usage_obj.get("input_tokens", 0)
ctx.output_tokens = usage_obj.get("output_tokens", 0)
details = usage_obj.get("input_tokens_details")
if isinstance(details, dict):
ctx.cached_tokens = details.get("cached_tokens", 0)
# 备用:从 response 字段提取
response_obj = data.get("response")
if isinstance(response_obj, dict) and not ctx.final_response:
ctx.final_response = response_obj
def _extract_response_metadata(
self,
response: dict[str, Any],
) -> dict[str, Any]:
"""
从 OpenAI 响应中提取元数据
提取 model、status、response_id 等字段作为元数据。
Args:
response: OpenAI API 响应
Returns:
提取的元数据字典
"""
metadata: dict[str, Any] = {}
# 提取模型名称(实际使用的模型)
if "model" in response:
metadata["model"] = response["model"]
# 提取响应 ID
if "id" in response:
metadata["response_id"] = response["id"]
# 提取状态
if "status" in response:
metadata["status"] = response["status"]
# 提取对象类型
if "object" in response:
metadata["object"] = response["object"]
# 提取系统指纹(如果存在)
if "system_fingerprint" in response:
metadata["system_fingerprint"] = response["system_fingerprint"]
return metadata
def _finalize_stream_metadata(self, ctx: StreamContext) -> None:
"""
从流上下文中提取最终元数据
在流传输完成后调用,从收集的事件中提取元数据。
Args:
ctx: 流上下文
"""
# 从 response_id 提取响应 ID
if ctx.response_id:
ctx.response_metadata["response_id"] = ctx.response_id
# 从 final_response 提取更多元数据
if ctx.final_response and isinstance(ctx.final_response, dict):
if "model" in ctx.final_response:
ctx.response_metadata["model"] = ctx.final_response["model"]
if "status" in ctx.final_response:
ctx.response_metadata["status"] = ctx.final_response["status"]
if "object" in ctx.final_response:
ctx.response_metadata["object"] = ctx.final_response["object"]
if "system_fingerprint" in ctx.final_response:
ctx.response_metadata["system_fingerprint"] = ctx.final_response[
"system_fingerprint"
]
# 如果没有从响应中获取到 model使用上下文中的
if "model" not in ctx.response_metadata and ctx.model:
ctx.response_metadata["model"] = ctx.model