mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
68
_deprecated_py_src/api/handlers/base/__init__.py
Normal file
68
_deprecated_py_src/api/handlers/base/__init__.py
Normal 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",
|
||||
]
|
||||
740
_deprecated_py_src/api/handlers/base/base_handler.py
Normal file
740
_deprecated_py_src/api/handlers/base/base_handler.py
Normal 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 key(family: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
|
||||
344
_deprecated_py_src/api/handlers/base/chat_adapter_base.py
Normal file
344
_deprecated_py_src/api/handlers/base/chat_adapter_base.py
Normal 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())
|
||||
156
_deprecated_py_src/api/handlers/base/chat_error_utils.py
Normal file
156
_deprecated_py_src/api/handlers/base/chat_error_utils.py
Normal 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)
|
||||
1344
_deprecated_py_src/api/handlers/base/chat_handler_base.py
Normal file
1344
_deprecated_py_src/api/handlers/base/chat_handler_base.py
Normal file
File diff suppressed because it is too large
Load Diff
971
_deprecated_py_src/api/handlers/base/chat_sync_executor.py
Normal file
971
_deprecated_py_src/api/handlers/base/chat_sync_executor.py
Normal 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}"
|
||||
313
_deprecated_py_src/api/handlers/base/cli_adapter_base.py
Normal file
313
_deprecated_py_src/api/handlers/base/cli_adapter_base.py
Normal 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())
|
||||
587
_deprecated_py_src/api/handlers/base/cli_event_mixin.py
Normal file
587
_deprecated_py_src/api/handlers/base/cli_event_mixin.py
Normal 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=False),data_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"
|
||||
119
_deprecated_py_src/api/handlers/base/cli_handler_base.py
Normal file
119
_deprecated_py_src/api/handlers/base/cli_handler_base.py
Normal 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
|
||||
723
_deprecated_py_src/api/handlers/base/cli_monitor_mixin.py
Normal file
723
_deprecated_py_src/api/handlers/base/cli_monitor_mixin.py
Normal 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()
|
||||
634
_deprecated_py_src/api/handlers/base/cli_prefetch_mixin.py
Normal file
634
_deprecated_py_src/api/handlers/base/cli_prefetch_mixin.py
Normal 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
|
||||
293
_deprecated_py_src/api/handlers/base/cli_protocol.py
Normal file
293
_deprecated_py_src/api/handlers/base/cli_protocol.py
Normal 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: ...
|
||||
505
_deprecated_py_src/api/handlers/base/cli_request_mixin.py
Normal file
505
_deprecated_py_src/api/handlers/base/cli_request_mixin.py
Normal 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 {}
|
||||
106
_deprecated_py_src/api/handlers/base/cli_sse_helpers.py
Normal file
106
_deprecated_py_src/api/handlers/base/cli_sse_helpers.py
Normal 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
|
||||
946
_deprecated_py_src/api/handlers/base/cli_stream_mixin.py
Normal file
946
_deprecated_py_src/api/handlers/base/cli_stream_mixin.py
Normal 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
|
||||
634
_deprecated_py_src/api/handlers/base/cli_sync_mixin.py
Normal file
634
_deprecated_py_src/api/handlers/base/cli_sync_mixin.py
Normal 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}"
|
||||
271
_deprecated_py_src/api/handlers/base/content_extractors.py
Normal file
271
_deprecated_py_src/api/handlers/base/content_extractors.py
Normal 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())
|
||||
1769
_deprecated_py_src/api/handlers/base/endpoint_checker.py
Normal file
1769
_deprecated_py_src/api/handlers/base/endpoint_checker.py
Normal file
File diff suppressed because it is too large
Load Diff
625
_deprecated_py_src/api/handlers/base/handler_adapter_base.py
Normal file
625
_deprecated_py_src/api/handlers/base/handler_adapter_base.py
Normal 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
|
||||
717
_deprecated_py_src/api/handlers/base/parsers.py
Normal file
717
_deprecated_py_src/api/handlers/base/parsers.py
Normal 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 可能在首个 chunk(message_start)或最后一个 chunk(message_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",
|
||||
]
|
||||
1416
_deprecated_py_src/api/handlers/base/request_builder.py
Normal file
1416
_deprecated_py_src/api/handlers/base/request_builder.py
Normal file
File diff suppressed because it is too large
Load Diff
19
_deprecated_py_src/api/handlers/base/response_parser.py
Normal file
19
_deprecated_py_src/api/handlers/base/response_parser.py
Normal 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",
|
||||
]
|
||||
473
_deprecated_py_src/api/handlers/base/stream_context.py
Normal file
473
_deprecated_py_src/api/handlers/base/stream_context.py
Normal 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 而非 and:CancelledError 路径中 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}"
|
||||
1584
_deprecated_py_src/api/handlers/base/stream_processor.py
Normal file
1584
_deprecated_py_src/api/handlers/base/stream_processor.py
Normal file
File diff suppressed because it is too large
Load Diff
674
_deprecated_py_src/api/handlers/base/stream_telemetry.py
Normal file
674
_deprecated_py_src/api/handlers/base/stream_telemetry.py
Normal 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)
|
||||
246
_deprecated_py_src/api/handlers/base/upstream_stream_bridge.py
Normal file
246
_deprecated_py_src/api/handlers/base/upstream_stream_bridge.py
Normal 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",
|
||||
]
|
||||
255
_deprecated_py_src/api/handlers/base/utils.py
Normal file
255
_deprecated_py_src/api/handlers/base/utils.py
Normal 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]:
|
||||
"""
|
||||
构建 SSE(text/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
|
||||
160
_deprecated_py_src/api/handlers/base/video_adapter_base.py
Normal file
160
_deprecated_py_src/api/handlers/base/video_adapter_base.py
Normal 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"]
|
||||
495
_deprecated_py_src/api/handlers/base/video_handler_base.py
Normal file
495
_deprecated_py_src/api/handlers/base/video_handler_base.py
Normal 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"]
|
||||
Reference in New Issue
Block a user