Merge branch 'fix/python314-upgrade'

# Conflicts:
#	src/api/handlers/base/base_handler.py
#	src/api/handlers/base/request_builder.py
#	src/models/endpoint_models.py
#	src/services/orchestration/candidate_resolver.py
#	src/services/orchestration/fallback_orchestrator.py
This commit is contained in:
fawney19
2026-01-30 12:59:52 +08:00
257 changed files with 4115 additions and 5236 deletions

View File

@@ -27,21 +27,20 @@
from __future__ import annotations
import asyncio
import time
from typing import (
TYPE_CHECKING,
Any,
Awaitable,
Callable,
Coroutine,
Dict,
Optional,
Protocol,
TypeVar,
runtime_checkable,
)
from collections.abc import Callable
from collections.abc import Awaitable, Coroutine
from fastapi import Request
from fastapi.responses import JSONResponse, StreamingResponse
from sqlalchemy.orm import Session
@@ -57,6 +56,9 @@ from src.services.usage.service import UsageService
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]]
class MessageTelemetry:
"""
@@ -105,29 +107,29 @@ class MessageTelemetry:
output_tokens: int,
response_time_ms: int,
status_code: int,
request_body: Dict[str, Any],
request_headers: Dict[str, Any],
request_body: dict[str, Any],
request_headers: dict[str, Any],
response_body: Any,
response_headers: Dict[str, Any],
client_response_headers: Optional[Dict[str, Any]] = None,
response_headers: dict[str, Any],
client_response_headers: dict[str, Any] | None = None,
cache_creation_tokens: int = 0,
cache_read_tokens: int = 0,
is_stream: bool = False,
provider_request_headers: Optional[Dict[str, Any]] = None,
provider_request_headers: dict[str, Any] | None = None,
# 时间指标
first_byte_time_ms: Optional[int] = None, # 首字时间/TTFB
first_byte_time_ms: int | None = None, # 首字时间/TTFB
# Provider 侧追踪信息(用于记录真实成本)
provider_id: Optional[str] = None,
provider_endpoint_id: Optional[str] = None,
provider_api_key_id: Optional[str] = None,
api_format: Optional[str] = None,
provider_id: str | None = None,
provider_endpoint_id: str | None = None,
provider_api_key_id: str | None = None,
api_format: str | None = None,
# 格式转换追踪
endpoint_api_format: Optional[str] = None, # 端点原生 API 格式
endpoint_api_format: str | None = None, # 端点原生 API 格式
has_format_conversion: bool = False, # 是否发生了格式转换
# 模型映射信息
target_model: Optional[str] = None,
target_model: str | None = None,
# Provider 响应元数据(如 Gemini 的 modelVersion
response_metadata: Optional[Dict[str, Any]] = None,
response_metadata: dict[str, Any] | None = None,
) -> float:
total_cost = await self.calculate_cost(
provider,
@@ -199,24 +201,24 @@ class MessageTelemetry:
response_time_ms: int,
status_code: int,
error_message: str,
request_body: Dict[str, Any],
request_headers: Dict[str, Any],
request_body: dict[str, Any],
request_headers: dict[str, Any],
is_stream: bool,
api_format: Optional[str] = None,
provider_request_headers: Optional[Dict[str, Any]] = None,
api_format: str | None = None,
provider_request_headers: dict[str, Any] | None = None,
# 预估 token 信息(来自 message_start 事件,用于中断请求的成本估算)
input_tokens: int = 0,
output_tokens: int = 0,
cache_creation_tokens: int = 0,
cache_read_tokens: int = 0,
response_body: Optional[Dict[str, Any]] = None,
response_headers: Optional[Dict[str, Any]] = None,
client_response_headers: Optional[Dict[str, Any]] = None,
response_body: dict[str, Any] | None = None,
response_headers: dict[str, Any] | None = None,
client_response_headers: dict[str, Any] | None = None,
# 格式转换追踪
endpoint_api_format: Optional[str] = None,
endpoint_api_format: str | None = None,
has_format_conversion: bool = False,
# 模型映射信息
target_model: Optional[str] = None,
target_model: str | None = None,
) -> None:
"""
记录失败请求
@@ -273,24 +275,24 @@ class MessageTelemetry:
provider: str,
model: str,
response_time_ms: int,
first_byte_time_ms: Optional[int],
first_byte_time_ms: int | None,
status_code: int,
request_body: Dict[str, Any],
request_headers: Dict[str, Any],
request_body: dict[str, Any],
request_headers: dict[str, Any],
is_stream: bool,
api_format: Optional[str] = None,
provider_request_headers: Optional[Dict[str, Any]] = None,
api_format: str | None = None,
provider_request_headers: dict[str, Any] | None = None,
input_tokens: int = 0,
output_tokens: int = 0,
cache_creation_tokens: int = 0,
cache_read_tokens: int = 0,
response_body: Optional[Dict[str, Any]] = None,
response_headers: Optional[Dict[str, Any]] = None,
client_response_headers: Optional[Dict[str, Any]] = None,
response_body: dict[str, Any] | None = None,
response_headers: dict[str, Any] | None = None,
client_response_headers: dict[str, Any] | None = None,
# 格式转换追踪
endpoint_api_format: Optional[str] = None,
endpoint_api_format: str | None = None,
has_format_conversion: bool = False,
target_model: Optional[str] = None,
target_model: str | None = None,
) -> None:
"""
记录客户端取消的请求
@@ -341,9 +343,9 @@ class MessageHandlerProtocol(Protocol):
self,
request: Any,
http_request: Request,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
query_params: Optional[Dict[str, str]] = None,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
) -> StreamingResponse:
"""处理流式请求"""
...
@@ -352,9 +354,9 @@ class MessageHandlerProtocol(Protocol):
self,
request: Any,
http_request: Request,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
query_params: Optional[Dict[str, str]] = None,
original_headers: dict[str, str],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
) -> JSONResponse:
"""处理非流式请求"""
...
@@ -371,9 +373,6 @@ class BaseMessageHandler:
推荐使用 MessageHandlerProtocol 中定义的签名。
"""
# Adapter 检测器类型
AdapterDetectorType = Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
def __init__(
self,
*,
@@ -384,8 +383,8 @@ class BaseMessageHandler:
client_ip: str,
user_agent: str,
start_time: float,
allowed_api_formats: Optional[list[str]] = None,
adapter_detector: Optional[AdapterDetectorType] = None,
allowed_api_formats: list[str] | None = None,
adapter_detector: AdapterDetectorType | None = None,
) -> None:
self.db = db
self.user = user
@@ -408,9 +407,9 @@ class BaseMessageHandler:
def _resolve_capability_requirements(
self,
model_name: str,
request_headers: Optional[Dict[str, str]] = None,
request_body: Optional[Dict[str, Any]] = None,
) -> Dict[str, bool]:
request_headers: dict[str, str] | None = None,
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""
解析请求的能力需求
@@ -442,12 +441,12 @@ class BaseMessageHandler:
async def _resolve_preferred_key_ids(
self,
model_name: str,
request_body: Optional[Dict[str, Any]] = None,
) -> Optional[list[str]]:
request_body: dict[str, Any] | None = None,
) -> list[str] | None:
"""可选的 Key 优先级解析钩子(默认不启用)。"""
return None
def get_api_format(self, provider_type: Optional[str] = None) -> APIFormat:
def get_api_format(self, provider_type: str | None = None) -> APIFormat:
"""根据 provider_type 解析 API 格式,未知类型默认 OPENAI"""
if provider_type:
result = resolve_api_format(provider_type, default=APIFormat.OPENAI)
@@ -456,17 +455,17 @@ class BaseMessageHandler:
def build_provider_payload(
self,
original_body: Dict[str, Any],
original_body: dict[str, Any],
*,
mapped_model: Optional[str] = None,
) -> 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 _update_usage_to_streaming(self, request_id: Optional[str] = None) -> None:
def _update_usage_to_streaming(self, request_id: str | None = None) -> None:
"""更新 Usage 状态为 streaming流式传输开始时调用
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输
@@ -500,7 +499,7 @@ class BaseMessageHandler:
# 创建后台任务,不阻塞当前流
asyncio.create_task(_do_update())
def _update_usage_to_streaming_with_ctx(self, ctx: "StreamContext") -> None:
def _update_usage_to_streaming_with_ctx(self, ctx: StreamContext) -> None:
"""更新 Usage 状态为 streaming同时更新 provider 相关信息
使用 asyncio 后台任务执行数据库更新,避免阻塞流式传输

View File

@@ -19,7 +19,7 @@ Chat Adapter 通用基类
import time
import traceback
from abc import abstractmethod
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
from fastapi import HTTPException, Request
@@ -65,7 +65,7 @@ class ChatAdapterBase(ApiAdapter):
# 子类必须覆盖
FORMAT_ID: str = "UNKNOWN"
HANDLER_CLASS: Type[ChatHandlerBase]
HANDLER_CLASS: type[ChatHandlerBase]
# 适配器配置
name: str = "chat.base"
@@ -90,7 +90,7 @@ class ChatAdapterBase(ApiAdapter):
return base_url
@classmethod
def build_base_headers(cls, api_key: str) -> Dict[str, str]:
def build_base_headers(cls, api_key: str) -> dict[str, str]:
"""构建基础请求头,使用统一的 headers.py 实现"""
return build_adapter_base_headers(cls._get_api_format(), api_key)
@@ -101,13 +101,13 @@ class ChatAdapterBase(ApiAdapter):
@classmethod
def build_headers_with_extra(
cls, api_key: str, extra_headers: Optional[Dict[str, str]] = None
) -> Dict[str, str]:
cls, api_key: str, extra_headers: dict[str, str] | None = None
) -> dict[str, str]:
"""构建完整请求头(包含 extra_headers使用统一的 headers.py 实现"""
return build_adapter_headers(cls._get_api_format(), api_key, extra_headers)
@classmethod
def build_request_body(cls, request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
def build_request_body(cls, request_data: dict[str, Any] | None = None) -> dict[str, Any]:
"""构建测试请求体,使用转换器注册表自动处理格式转换
Args:
@@ -120,11 +120,11 @@ class ChatAdapterBase(ApiAdapter):
return build_test_request_body(cls.FORMAT_ID, request_data)
def extract_api_key(self, request: Request) -> Optional[str]:
def extract_api_key(self, request: Request) -> str | None:
"""从请求中提取 API 密钥,使用统一的 headers.py 实现"""
return extract_client_api_key(dict(request.headers), self._get_api_format())
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
async def handle(self, context: ApiRequestContext):
@@ -282,8 +282,8 @@ class ChatAdapterBase(ApiAdapter):
)
def _merge_path_params(
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any]
) -> Dict[str, Any]:
self, original_request_body: dict[str, Any], path_params: dict[str, Any]
) -> dict[str, Any]:
"""
合并 URL 路径参数到请求体 - 子类可覆盖
@@ -316,7 +316,7 @@ class ChatAdapterBase(ApiAdapter):
"""
pass
def _extract_message_count(self, payload: Dict[str, Any], request_obj) -> int:
def _extract_message_count(self, payload: dict[str, Any], request_obj) -> int:
"""
提取消息数量 - 子类可覆盖
@@ -327,7 +327,7 @@ class ChatAdapterBase(ApiAdapter):
messages = request_obj.messages
return len(messages) if isinstance(messages, list) else 0
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
"""
构建审计日志元数据 - 子类可覆盖
"""
@@ -355,8 +355,8 @@ class ChatAdapterBase(ApiAdapter):
model: str,
stream: bool,
start_time: float,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
client_ip: str,
request_id: str,
) -> JSONResponse:
@@ -426,8 +426,8 @@ class ChatAdapterBase(ApiAdapter):
model: str,
stream: bool,
start_time: float,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
client_ip: str,
request_id: str,
) -> JSONResponse:
@@ -527,12 +527,12 @@ class ChatAdapterBase(ApiAdapter):
cache_read_input_tokens: int,
input_price_per_1m: float,
output_price_per_1m: float,
cache_creation_price_per_1m: Optional[float],
cache_read_price_per_1m: Optional[float],
price_per_request: Optional[float],
tiered_pricing: Optional[dict] = None,
cache_ttl_minutes: Optional[int] = None,
) -> Dict[str, Any]:
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]:
"""
计算请求成本
@@ -597,8 +597,8 @@ class ChatAdapterBase(ApiAdapter):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""
查询上游 API 支持的模型列表
@@ -626,16 +626,16 @@ class ChatAdapterBase(ApiAdapter):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
request_data: Dict[str, Any],
extra_headers: Optional[Dict[str, str]] = None,
request_data: dict[str, Any],
extra_headers: dict[str, str] | None = None,
# 用量计算参数(现在强制记录)
db: Optional[Any] = None,
user: Optional[Any] = None,
provider_name: Optional[str] = None,
provider_id: Optional[str] = None,
api_key_id: Optional[str] = None,
model_name: Optional[str] = None,
) -> Dict[str, Any]:
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,
) -> dict[str, Any]:
"""
测试模型连接性(非流式)
@@ -682,11 +682,11 @@ class ChatAdapterBase(ApiAdapter):
# Adapter 注册表 - 用于根据 API format 获取 Adapter 实例
# =========================================================================
_ADAPTER_REGISTRY: Dict[str, Type["ChatAdapterBase"]] = {}
_ADAPTER_REGISTRY: dict[str, type[ChatAdapterBase]] = {}
_ADAPTERS_LOADED = False
def register_adapter(adapter_class: Type["ChatAdapterBase"]) -> Type["ChatAdapterBase"]:
def register_adapter(adapter_class: type[ChatAdapterBase]) -> type[ChatAdapterBase]:
"""
注册 Adapter 类到注册表
@@ -731,7 +731,7 @@ def _ensure_adapters_loaded():
_ADAPTERS_LOADED = True
def get_adapter_class(api_format: str) -> Optional[Type["ChatAdapterBase"]]:
def get_adapter_class(api_format: str) -> type[ChatAdapterBase] | None:
"""
根据 API format 获取 Adapter 类
@@ -745,7 +745,7 @@ def get_adapter_class(api_format: str) -> Optional[Type["ChatAdapterBase"]]:
return _ADAPTER_REGISTRY.get(api_format.upper()) if api_format else None
def get_adapter_instance(api_format: str) -> Optional["ChatAdapterBase"]:
def get_adapter_instance(api_format: str) -> ChatAdapterBase | None:
"""
根据 API format 获取 Adapter 实例

View File

@@ -22,7 +22,10 @@ Chat Handler Base - Chat API 格式的通用基类
import asyncio
import json
from abc import ABC, abstractmethod
from typing import Any, AsyncGenerator, Awaitable, Callable, Dict, Optional, Union
from typing import Any
from collections.abc import Callable
from collections.abc import AsyncGenerator, Awaitable
import httpx
from fastapi import BackgroundTasks, Request
@@ -75,10 +78,10 @@ def _get_error_status_code(e: Exception, default: int = 400) -> int:
def _convert_error_response_best_effort(
error_response: Dict[str, Any],
error_response: dict[str, Any],
source_format: str,
target_format: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
将上游错误响应 best-effort 转换为客户端格式。
@@ -97,7 +100,7 @@ def _convert_error_response_best_effort(
def _build_client_error_response_best_effort(
message: str,
target_format: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
当无法解析上游错误 body 时构造一个目标格式的错误响应best-effort
"""
@@ -117,11 +120,11 @@ def _build_client_error_response_best_effort(
def _build_error_json_payload(
e: Union[ThinkingSignatureException, UpstreamClientException],
e: ThinkingSignatureException | UpstreamClientException,
client_format: str,
provider_format: str,
needs_conversion: bool = True,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
构建错误 JSON 响应 payload公共逻辑
@@ -185,10 +188,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
client_ip: str,
user_agent: str,
start_time: float,
allowed_api_formats: Optional[list] = None,
adapter_detector: Optional[
Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
] = None,
allowed_api_formats: list | None = None,
adapter_detector: None | (
Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
) = None,
):
allowed = allowed_api_formats or [self.FORMAT_ID]
super().__init__(
@@ -202,7 +205,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
allowed_api_formats=allowed,
adapter_detector=adapter_detector,
)
self._parser: Optional[ResponseParser] = None
self._parser: ResponseParser | None = None
self._request_builder = PassthroughRequestBuilder()
@property
@@ -228,7 +231,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
pass
@abstractmethod
def _extract_usage(self, response: Dict) -> Dict[str, int]:
def _extract_usage(self, response: dict) -> dict[str, int]:
"""
从响应中提取 token 使用情况
@@ -241,7 +244,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
"""
pass
def _normalize_response(self, response: Dict) -> Dict:
def _normalize_response(self, response: dict) -> dict:
"""
规范化响应(可选覆盖)
@@ -257,8 +260,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002 - 子类使用
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002 - 子类使用
) -> str:
"""
从请求中提取模型名 - 子类可覆盖
@@ -282,9 +285,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
def apply_mapped_model(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
mapped_model: str, # noqa: ARG002 - 子类使用
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
将映射后的模型名应用到请求体
@@ -303,9 +306,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
def get_model_for_url(
self,
request_body: Dict[str, Any],
mapped_model: Optional[str],
) -> Optional[str]:
request_body: dict[str, Any],
mapped_model: str | None,
) -> str | None:
"""
获取用于 URL 路径的模型名
@@ -323,8 +326,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
def prepare_provider_request_body(
self,
request_body: Dict[str, Any],
) -> Dict[str, Any]:
request_body: dict[str, Any],
) -> dict[str, Any]:
"""
准备发送给 Provider 的请求体 - 子类可覆盖
@@ -341,9 +344,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
def _set_model_after_conversion(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
provider_api_format: str,
mapped_model: Optional[str],
mapped_model: str | None,
fallback_model: str,
) -> None:
"""
@@ -372,7 +375,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
def _set_stream_after_conversion(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
client_api_format: str,
provider_api_format: str,
is_stream: bool,
@@ -414,8 +417,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
self,
source_model: str,
provider_id: str,
api_format: Optional[str] = None,
) -> Optional[str]:
api_format: str | None = None,
) -> str | None:
"""
获取模型映射后的实际模型名
@@ -452,10 +455,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
self,
request: Any,
http_request: Request,
original_headers: Dict[str, Any],
original_request_body: Dict[str, Any],
query_params: Optional[Dict[str, str]] = None,
) -> Union[StreamingResponse, JSONResponse]:
original_headers: dict[str, Any],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
) -> StreamingResponse | JSONResponse:
"""处理流式响应"""
logger.debug(f"开始流式响应处理 ({self.FORMAT_ID})")
@@ -466,7 +469,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: Dict[str, Any] = {"body": original_request_body}
request_body_ref: dict[str, Any] = {"body": original_request_body}
# 创建类型安全的流式上下文
ctx = StreamContext(model=model, api_format=api_format)
@@ -492,7 +495,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> AsyncGenerator[bytes, None]:
) -> AsyncGenerator[bytes]:
return await self._execute_stream_request(
ctx,
stream_processor,
@@ -615,12 +618,12 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
original_request_body: Dict[str, Any],
original_headers: Dict[str, str],
query_params: Optional[Dict[str, str]] = None,
candidate: Optional[ProviderCandidate] = None,
is_disconnected: Optional[Callable[[], Awaitable[bool]]] = None,
) -> AsyncGenerator[bytes, None]:
original_request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
candidate: ProviderCandidate | None = None,
is_disconnected: Callable[[], Awaitable[bool]] | None = None,
) -> AsyncGenerator[bytes]:
"""执行流式请求并返回流生成器"""
# 重置上下文状态(重试时清除之前的数据)
ctx.reset_for_retry()
@@ -799,7 +802,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
ctx.error_message = "client_disconnected_during_prefetch"
raise
except asyncio.TimeoutError:
except TimeoutError:
# 整体请求超时(建立连接 + 获取首字节)
# 清理可能已建立的连接上下文
if response_ctx is not None:
@@ -856,8 +859,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
self,
ctx: StreamContext,
error: Exception,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None:
"""记录流式请求失败"""
response_time_ms = self.elapsed_ms()
@@ -904,9 +907,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
self,
request: Any,
http_request: Request,
original_headers: Dict[str, Any],
original_request_body: Dict[str, Any],
query_params: Optional[Dict[str, str]] = None,
original_headers: dict[str, Any],
original_request_body: dict[str, Any],
query_params: dict[str, str] | None = None,
) -> JSONResponse:
"""处理非流式响应"""
logger.debug(f"开始非流式响应处理 ({self.FORMAT_ID})")
@@ -918,29 +921,29 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: Dict[str, Any] = {"body": original_request_body}
request_body_ref: dict[str, Any] = {"body": original_request_body}
# 用于跟踪的变量
provider_name: Optional[str] = None
response_json: Optional[Dict[str, Any]] = None
provider_name: str | None = None
response_json: dict[str, Any] | None = None
status_code = 200
response_headers: Dict[str, str] = {}
provider_request_headers: Dict[str, str] = {}
provider_request_body: Optional[Dict[str, Any]] = None
provider_api_format_for_error: Optional[str] = None
client_api_format_for_error: Optional[str] = None
response_headers: dict[str, str] = {}
provider_request_headers: dict[str, str] = {}
provider_request_body: dict[str, Any] | None = None
provider_api_format_for_error: str | None = None
client_api_format_for_error: str | None = None
needs_conversion_for_error: bool = False
provider_id: Optional[str] = None # Provider ID用于失败记录
endpoint_id: Optional[str] = None # Endpoint ID用于失败记录
key_id: Optional[str] = None # Key ID用于失败记录
mapped_model_result: Optional[str] = None # 映射后的目标模型名(用于 Usage 记录)
provider_id: str | None = None # Provider ID用于失败记录
endpoint_id: str | None = None # Endpoint ID用于失败记录
key_id: str | None = None # Key ID用于失败记录
mapped_model_result: str | None = None # 映射后的目标模型名(用于 Usage 记录)
async def sync_request_func(
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> Dict[str, Any]:
) -> dict[str, Any]:
nonlocal provider_name, response_json, status_code, response_headers
nonlocal provider_request_headers, provider_request_body, mapped_model_result
nonlocal provider_api_format_for_error, client_api_format_for_error, needs_conversion_for_error
@@ -1293,7 +1296,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
actual_request_body = provider_request_body or original_request_body
# 尝试从异常中提取响应头
error_response_headers: Dict[str, str] = {}
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"):

View File

@@ -17,7 +17,7 @@ CLI Adapter 通用基类
import time
import traceback
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
from fastapi import HTTPException, Request
@@ -63,7 +63,7 @@ class CliAdapterBase(ApiAdapter):
# 子类必须覆盖
FORMAT_ID: str = "UNKNOWN"
HANDLER_CLASS: Type[CliMessageHandlerBase]
HANDLER_CLASS: type[CliMessageHandlerBase]
# 适配器配置
name: str = "cli.base"
@@ -72,7 +72,7 @@ class CliAdapterBase(ApiAdapter):
# 计费模板配置(子类可覆盖,如 "claude", "openai", "gemini"
BILLING_TEMPLATE: str = "claude"
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
self.allowed_api_formats = allowed_api_formats or [self.FORMAT_ID]
# =========================================================================
@@ -87,7 +87,7 @@ class CliAdapterBase(ApiAdapter):
except KeyError:
return APIFormat.OPENAI
def extract_api_key(self, request: Request) -> Optional[str]:
def extract_api_key(self, request: Request) -> str | None:
"""
从请求中提取 API 密钥
@@ -96,7 +96,7 @@ class CliAdapterBase(ApiAdapter):
return extract_client_api_key(dict(request.headers), self._get_api_format())
@classmethod
def build_base_headers(cls, api_key: str) -> Dict[str, str]:
def build_base_headers(cls, api_key: str) -> dict[str, str]:
"""
构建 CLI API 认证头
@@ -106,8 +106,8 @@ class CliAdapterBase(ApiAdapter):
@classmethod
def build_headers_with_extra(
cls, api_key: str, extra_headers: Optional[Dict[str, str]] = None
) -> Dict[str, str]:
cls, api_key: str, extra_headers: dict[str, str] | None = None
) -> dict[str, str]:
"""
构建带额外头部的完整请求头
@@ -260,8 +260,8 @@ class CliAdapterBase(ApiAdapter):
)
def _merge_path_params(
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any]
) -> Dict[str, Any]:
self, original_request_body: dict[str, Any], path_params: dict[str, Any]
) -> dict[str, Any]:
"""
合并 URL 路径参数到请求体 - 子类可覆盖
@@ -280,7 +280,7 @@ class CliAdapterBase(ApiAdapter):
merged[key] = value
return merged
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
def _extract_message_count(self, payload: dict[str, Any]) -> int:
"""
提取消息数量 - 子类可覆盖
@@ -297,9 +297,9 @@ class CliAdapterBase(ApiAdapter):
def _build_audit_metadata(
self,
payload: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
payload: dict[str, Any],
path_params: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""
构建审计日志元数据 - 子类可覆盖
@@ -338,8 +338,8 @@ class CliAdapterBase(ApiAdapter):
model: str,
stream: bool,
start_time: float,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
client_ip: str,
request_id: str,
) -> JSONResponse:
@@ -409,8 +409,8 @@ class CliAdapterBase(ApiAdapter):
model: str,
stream: bool,
start_time: float,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
client_ip: str,
request_id: str,
) -> JSONResponse:
@@ -507,12 +507,12 @@ class CliAdapterBase(ApiAdapter):
cache_read_input_tokens: int,
input_price_per_1m: float,
output_price_per_1m: float,
cache_creation_price_per_1m: Optional[float],
cache_read_price_per_1m: Optional[float],
price_per_request: Optional[float],
tiered_pricing: Optional[dict] = None,
cache_ttl_minutes: Optional[int] = None,
) -> Dict[str, Any]:
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]:
"""
计算请求成本
@@ -567,8 +567,8 @@ class CliAdapterBase(ApiAdapter):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""
查询上游 API 支持的模型列表
@@ -596,16 +596,16 @@ class CliAdapterBase(ApiAdapter):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
request_data: Dict[str, Any],
extra_headers: Optional[Dict[str, str]] = None,
request_data: dict[str, Any],
extra_headers: dict[str, str] | None = None,
# 用量计算参数
db: Optional[Any] = None,
user: Optional[Any] = None,
provider_name: Optional[str] = None,
provider_id: Optional[str] = None,
api_key_id: Optional[str] = None,
model_name: Optional[str] = None,
) -> Dict[str, Any]:
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,
) -> dict[str, Any]:
"""
测试模型连接性(非流式)
@@ -669,7 +669,7 @@ class CliAdapterBase(ApiAdapter):
# =========================================================================
@classmethod
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
"""
构建CLI API端点URL - 子类应覆盖
@@ -684,7 +684,7 @@ class CliAdapterBase(ApiAdapter):
raise NotImplementedError(f"{cls.FORMAT_ID} adapter must implement build_endpoint_url")
@classmethod
def build_request_body(cls, request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
def build_request_body(cls, request_data: dict[str, Any] | None = None) -> dict[str, Any]:
"""构建测试请求体,使用转换器注册表自动处理格式转换
Args:
@@ -698,7 +698,7 @@ class CliAdapterBase(ApiAdapter):
return build_test_request_body(cls.FORMAT_ID, request_data)
@classmethod
def get_cli_user_agent(cls) -> Optional[str]:
def get_cli_user_agent(cls) -> str | None:
"""
获取CLI User-Agent - 子类可覆盖
@@ -708,7 +708,7 @@ class CliAdapterBase(ApiAdapter):
return None
@classmethod
def get_cli_extra_headers(cls) -> Dict[str, str]:
def get_cli_extra_headers(cls) -> dict[str, str]:
"""
获取CLI额外请求头 - 子类可覆盖
@@ -718,7 +718,7 @@ class CliAdapterBase(ApiAdapter):
Returns:
额外请求头字典
"""
headers: 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
@@ -728,11 +728,11 @@ class CliAdapterBase(ApiAdapter):
# CLI Adapter 注册表 - 用于根据 API format 获取 CLI Adapter 实例
# =========================================================================
_CLI_ADAPTER_REGISTRY: Dict[str, Type["CliAdapterBase"]] = {}
_CLI_ADAPTER_REGISTRY: dict[str, type[CliAdapterBase]] = {}
_CLI_ADAPTERS_LOADED = False
def register_cli_adapter(adapter_class: Type["CliAdapterBase"]) -> Type["CliAdapterBase"]:
def register_cli_adapter(adapter_class: type[CliAdapterBase]) -> type[CliAdapterBase]:
"""
注册 CLI Adapter 类到注册表
@@ -771,13 +771,13 @@ def _ensure_cli_adapters_loaded():
_CLI_ADAPTERS_LOADED = True
def get_cli_adapter_class(api_format: str) -> Optional[Type["CliAdapterBase"]]:
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) -> Optional["CliAdapterBase"]:
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:

View File

@@ -10,6 +10,8 @@ CLI Message Handler 通用基类
3. 简化新格式接入 - 只需实现 ResponseParser 和少量钩子方法
"""
from __future__ import annotations
import asyncio
import codecs
import json
@@ -17,14 +19,11 @@ import time
from typing import (
TYPE_CHECKING,
Any,
AsyncGenerator,
Callable,
Dict,
List,
Optional,
Tuple,
)
from collections.abc import Callable
from collections.abc import AsyncGenerator
import httpx
from fastapi import BackgroundTasks, Request
from fastapi.responses import JSONResponse, StreamingResponse
@@ -45,7 +44,6 @@ from src.api.handlers.base.request_builder import PassthroughRequestBuilder, get
# 直接从具体模块导入,避免循环依赖
from src.api.handlers.base.response_parser import (
ResponseParser,
StreamStats,
)
from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.utils import (
@@ -86,7 +84,7 @@ from src.utils.timeout import read_first_chunk_with_ttfb_timeout
# ==============================================================================
def _parse_sse_data_line(line: str) -> Tuple[Optional[Any], str]:
def _parse_sse_data_line(line: str) -> tuple[Any | None, str]:
"""
解析标准 SSE data 行
@@ -108,7 +106,7 @@ def _parse_sse_data_line(line: str) -> Tuple[Optional[Any], str]:
return None, "invalid"
def _parse_sse_event_data_line(line: str) -> Tuple[Optional[Any], str]:
def _parse_sse_event_data_line(line: str) -> tuple[Any | None, str]:
"""
解析 event + data 同行格式(如 "event: xxx data: {...}"
@@ -126,7 +124,7 @@ def _parse_sse_event_data_line(line: str) -> Tuple[Optional[Any], str]:
return None, "invalid"
def _parse_gemini_json_array_line(line: str) -> Tuple[Optional[Any], str]:
def _parse_gemini_json_array_line(line: str) -> tuple[Any | None, str]:
"""
解析 Gemini JSON-array 格式的裸 JSON 行
@@ -151,9 +149,9 @@ def _parse_gemini_json_array_line(line: str) -> Tuple[Optional[Any], str]:
def _format_converted_events_to_sse(
converted_events: List[Dict[str, Any]],
converted_events: list[dict[str, Any]],
client_format: str,
) -> List[str]:
) -> list[str]:
"""
将转换后的事件格式化为 SSE 行
@@ -164,7 +162,7 @@ def _format_converted_events_to_sse(
Returns:
SSE 行列表(每个元素是完整的 SSE 事件,包含尾部空行)
"""
result: List[str] = []
result: list[str] = []
needs_event_line = client_format.upper() in ("CLAUDE", "CLAUDE_CLI")
for evt in converted_events:
@@ -213,10 +211,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
client_ip: str,
user_agent: str,
start_time: float,
allowed_api_formats: Optional[list] = None,
adapter_detector: Optional[
Callable[[Dict[str, str], Optional[Dict[str, Any]]], Dict[str, bool]]
] = None,
allowed_api_formats: list | None = None,
adapter_detector: None | (
Callable[[dict[str, str], dict[str, Any] | None], dict[str, bool]]
) = None,
):
allowed = allowed_api_formats or [self.FORMAT_ID]
super().__init__(
@@ -230,7 +228,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
allowed_api_formats=allowed,
adapter_detector=adapter_detector,
)
self._parser: Optional[ResponseParser] = None
self._parser: ResponseParser | None = None
self._request_builder = PassthroughRequestBuilder()
@property
@@ -253,7 +251,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
self,
source_model: str,
provider_id: str,
) -> Optional[str]:
) -> str | None:
"""
获取模型映射后的实际模型名
@@ -296,8 +294,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002 - 子类使用
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002 - 子类使用
) -> str:
"""
从请求中提取模型名 - 子类可覆盖
@@ -321,9 +319,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
def apply_mapped_model(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
mapped_model: str, # noqa: ARG002 - 子类使用
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
将映射后的模型名应用到请求体
@@ -342,8 +340,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
def prepare_provider_request_body(
self,
request_body: Dict[str, Any],
) -> Dict[str, Any]:
request_body: dict[str, Any],
) -> dict[str, Any]:
"""
准备发送给 Provider 的请求体 - 子类可覆盖
@@ -359,7 +357,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
return request_body
@staticmethod
def _get_format_metadata(format_id: str) -> Optional["ApiFormatDefinition"]:
def _get_format_metadata(format_id: str) -> ApiFormatDefinition | None:
"""获取格式元数据(解析失败返回 None"""
from src.core.api_format import APIFormat
from src.core.api_format.metadata import API_FORMAT_DEFINITIONS
@@ -372,10 +370,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _finalize_converted_request(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
client_api_format: str,
provider_api_format: str,
mapped_model: Optional[str],
mapped_model: str | None,
fallback_model: str,
is_stream: bool,
) -> None:
@@ -418,13 +416,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _convert_request_for_cross_format(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
client_api_format: str,
provider_api_format: str,
mapped_model: Optional[str],
mapped_model: str | None,
fallback_model: str,
is_stream: bool,
) -> Tuple[Dict[str, Any], str]:
) -> tuple[dict[str, Any], str]:
"""
跨格式请求转换的公共逻辑
@@ -465,9 +463,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
def get_model_for_url(
self,
request_body: Dict[str, Any],
mapped_model: Optional[str],
) -> Optional[str]:
request_body: dict[str, Any],
mapped_model: str | None,
) -> str | None:
"""
获取用于 URL 路径的模型名
@@ -485,8 +483,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _extract_response_metadata(
self,
response: Dict[str, Any],
) -> Dict[str, Any]:
response: dict[str, Any],
) -> dict[str, Any]:
"""
从响应中提取 Provider 特有的元数据 - 子类可覆盖
@@ -503,11 +501,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
async def process_stream(
self,
original_request_body: Dict[str, Any],
original_headers: Dict[str, str],
query_params: Optional[Dict[str, str]] = None,
path_params: Optional[Dict[str, Any]] = None,
http_request: Optional[Request] = None,
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,
) -> StreamingResponse:
"""
处理流式请求
@@ -529,7 +527,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: Dict[str, Any] = {"body": original_request_body}
request_body_ref: dict[str, Any] = {"body": original_request_body}
# 使用子类实现的方法提取 model不同 API 格式的 model 位置不同)
# 注意:使用 original_request_body因为整流只修改 messages不影响 model 字段
@@ -550,7 +548,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> AsyncGenerator[bytes, None]:
) -> AsyncGenerator[bytes]:
return await self._execute_stream_request(
ctx,
provider,
@@ -653,12 +651,12 @@ class CliMessageHandlerBase(BaseMessageHandler):
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
original_request_body: Dict[str, Any],
original_headers: Dict[str, str],
query_params: Optional[Dict[str, str]] = None,
candidate: Optional[ProviderCandidate] = None,
http_request: Optional[Request] = None,
) -> AsyncGenerator[bytes, None]:
original_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,
) -> AsyncGenerator[bytes]:
"""执行流式请求并返回流生成器"""
# 重置上下文状态(重试时清除之前的数据,避免累积)
ctx.parsed_chunks = []
@@ -824,7 +822,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
else:
await asyncio.wait_for(_connect_and_prefetch(), timeout=request_timeout)
except asyncio.TimeoutError:
except TimeoutError:
# 整体请求超时(建立连接 + 获取首字节)
# 清理可能已建立的连接上下文
if response_ctx is not None:
@@ -898,7 +896,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
stream_response: httpx.Response,
response_ctx: Any,
http_client: httpx.AsyncClient,
) -> AsyncGenerator[bytes, None]:
) -> AsyncGenerator[bytes]:
"""创建响应流生成器(使用字节流)"""
try:
sse_parser = SSEEventParser()
@@ -961,9 +959,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
},
}
self._mark_first_output(ctx, output_state)
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode(
"utf-8"
)
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
return # 结束生成器
# 格式转换或直接透传
@@ -1015,7 +1011,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
logger.debug("流式数据转发完成")
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
@@ -1040,7 +1036,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
except httpx.RemoteProtocolError:
if ctx.data_count > 0:
error_event = {
@@ -1050,7 +1046,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
"message": "上游连接意外关闭,部分响应已成功传输",
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
raise
finally:
@@ -1241,7 +1237,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
except (EmbeddedErrorException, ProviderTimeoutException, ProviderNotAvailableException):
# 重新抛出可重试的 Provider 异常,触发故障转移
raise
except (OSError, IOError) as e:
except OSError as e:
# 网络 I/O 异常:记录警告,可能需要重试
logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}")
except Exception as e:
@@ -1261,7 +1257,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
response_ctx: Any,
http_client: httpx.AsyncClient,
prefetched_chunks: list,
) -> AsyncGenerator[bytes, None]:
) -> AsyncGenerator[bytes]:
"""创建响应流生成器(带预读数据,使用字节流)"""
try:
sse_parser = SSEEventParser()
@@ -1382,9 +1378,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
},
}
self._mark_first_output(ctx, output_state)
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode(
"utf-8"
)
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
return
# 格式转换或直接透传
@@ -1439,7 +1433,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
logger.debug("流式数据转发完成")
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
@@ -1463,7 +1457,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
"message": ctx.error_message,
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
except httpx.RemoteProtocolError:
if ctx.data_count > 0:
error_event = {
@@ -1473,7 +1467,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
"message": "上游连接意外关闭,部分响应已成功传输",
},
}
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode()
else:
raise
finally:
@@ -1489,7 +1483,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _handle_sse_event(
self,
ctx: StreamContext,
event_name: Optional[str],
event_name: str | None,
data_str: str,
record_chunk: bool = False,
) -> None:
@@ -1538,7 +1532,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
self,
ctx: StreamContext,
event_type: str,
data: Dict[str, Any],
data: dict[str, Any],
) -> None:
"""
处理解析后的事件数据 - 子类应覆盖此方法
@@ -1612,7 +1606,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _record_converted_chunks(
self,
ctx: StreamContext,
converted_events: List[Dict[str, Any]],
converted_events: list[dict[str, Any]],
) -> None:
"""
记录转换后的 chunk 数据到 parsed_chunks并更新统计信息
@@ -1656,7 +1650,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _extract_usage_from_converted_event(
self,
ctx: StreamContext,
evt: Dict[str, Any],
evt: dict[str, Any],
event_type: str,
) -> None:
"""
@@ -1672,7 +1666,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
evt: 转换后的事件
event_type: 事件类型
"""
usage: Optional[Dict[str, Any]] = None
usage: dict[str, Any] | None = None
# Claude 格式: message_delta 或 message_start
if event_type == "message_delta":
@@ -1737,9 +1731,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
async def _create_monitored_stream(
self,
ctx: StreamContext,
stream_generator: AsyncGenerator[bytes, None],
http_request: Optional[Request] = None,
) -> AsyncGenerator[bytes, None]:
stream_generator: AsyncGenerator[bytes],
http_request: Request | None = None,
) -> AsyncGenerator[bytes]:
"""
创建带监控的流生成器
@@ -1833,8 +1827,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
async def _record_stream_stats(
self,
ctx: StreamContext,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None:
"""在流完成后记录统计信息"""
try:
@@ -1996,7 +1990,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
from src.services.request.candidate import RequestCandidateService
# 计算候选自身的 TTFB
candidate_first_byte_time_ms: Optional[int] = None
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(
@@ -2061,8 +2055,8 @@ class CliMessageHandlerBase(BaseMessageHandler):
self,
ctx: StreamContext,
error: Exception,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
) -> None:
"""记录流式请求失败"""
# 使用 self.start_time 作为时间基准,与首字时间保持一致
@@ -2111,10 +2105,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
async def process_sync(
self,
original_request_body: Dict[str, Any],
original_headers: Dict[str, str],
query_params: Optional[Dict[str, str]] = None,
path_params: Optional[Dict[str, Any]] = None,
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,
) -> JSONResponse:
"""
处理非流式请求
@@ -2142,19 +2136,19 @@ class CliMessageHandlerBase(BaseMessageHandler):
endpoint_id = None # Endpoint ID用于失败记录
key_id = None # Key ID用于失败记录
mapped_model_result = None # 映射后的目标模型名(用于 Usage 记录)
response_metadata_result: Dict[str, Any] = {} # Provider 响应元数据
response_metadata_result: dict[str, Any] = {} # Provider 响应元数据
needs_conversion = False # 是否需要格式转换(由 candidate 决定)
# 可变请求体容器:允许 orchestrator 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: Dict[str, Any] = {"body": original_request_body}
request_body_ref: dict[str, Any] = {"body": original_request_body}
async def sync_request_func(
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
candidate: ProviderCandidate,
) -> Dict[str, Any]:
) -> 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
provider_name = str(provider.name)
provider_api_format = str(endpoint.api_format) if endpoint.api_format else ""
@@ -2470,7 +2464,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
actual_request_body = provider_request_body or original_request_body
# 尝试从异常中提取响应头
error_response_headers: Dict[str, str] = {}
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"):
@@ -2581,7 +2575,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
)
return True
def _mark_first_output(self, ctx: StreamContext, state: Dict[str, bool]) -> None:
def _mark_first_output(self, ctx: StreamContext, state: dict[str, bool]) -> None:
"""
标记首次输出:记录 TTFB 并更新 streaming 状态
@@ -2605,7 +2599,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
ctx: StreamContext,
line: str,
events: list, # noqa: ARG002 - 预留给上下文感知转换
) -> Tuple[List[str], List[Dict[str, Any]]]:
) -> tuple[list[str], list[dict[str, Any]]]:
"""
将 SSE 行从 Provider 格式转换为客户端格式
@@ -2690,7 +2684,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
def _parse_sse_line_to_json(
self, line: str, provider_format: str
) -> Tuple[Optional[Any], str]:
) -> tuple[Any | None, str]:
"""
解析 SSE 行为 JSON 对象

View File

@@ -8,7 +8,6 @@ StreamSmoother 使用这些提取器来处理不同格式的 SSE 事件。
import copy
import json
from abc import ABC, abstractmethod
from typing import Optional
class ContentExtractor(ABC):
@@ -20,7 +19,7 @@ class ContentExtractor(ABC):
"""
@abstractmethod
def extract_content(self, data: dict) -> Optional[str]:
def extract_content(self, data: dict) -> str | None:
"""
从 SSE 数据中提取可拆分的文本内容
@@ -64,7 +63,7 @@ class OpenAIContentExtractor(ContentExtractor):
- 只在 delta 仅包含 role/content 时允许拆分,避免破坏 tool_calls 等结构
"""
def extract_content(self, data: dict) -> Optional[str]:
def extract_content(self, data: dict) -> str | None:
if not isinstance(data, dict):
return None
@@ -115,7 +114,7 @@ class OpenAIContentExtractor(ContentExtractor):
new_choices.append(new_choice)
new_data["choices"] = new_choices
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode("utf-8")
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
class ClaudeContentExtractor(ContentExtractor):
@@ -127,7 +126,7 @@ class ClaudeContentExtractor(ContentExtractor):
- 数据结构: delta.type=text_delta, delta.text
"""
def extract_content(self, data: dict) -> Optional[str]:
def extract_content(self, data: dict) -> str | None:
if not isinstance(data, dict):
return None
@@ -165,9 +164,7 @@ class ClaudeContentExtractor(ContentExtractor):
# 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(
"utf-8"
)
return f"event: {event_name}\ndata: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
class GeminiContentExtractor(ContentExtractor):
@@ -179,7 +176,7 @@ class GeminiContentExtractor(ContentExtractor):
- 只有纯文本块才拆分
"""
def extract_content(self, data: dict) -> Optional[str]:
def extract_content(self, data: dict) -> str | None:
if not isinstance(data, dict):
return None
@@ -226,7 +223,7 @@ class GeminiContentExtractor(ContentExtractor):
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("utf-8")
return f"data: {json.dumps(new_data, ensure_ascii=False)}\n\n".encode()
# 提取器注册表
@@ -237,7 +234,7 @@ _EXTRACTORS: dict[str, type[ContentExtractor]] = {
}
def get_extractor(format_name: str) -> Optional[ContentExtractor]:
def get_extractor(format_name: str) -> ContentExtractor | None:
"""
根据格式名获取对应的内容提取器实例

View File

@@ -14,15 +14,14 @@
- EndpointCheckOrchestrator: 协调整个流程
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, AsyncIterator, Dict, Iterable, Optional, Union, List
from abc import ABC, abstractmethod
from typing import Any
from collections.abc import Iterable
import time
import uuid
import json
from functools import lru_cache
import asyncio
from collections import defaultdict
import httpx
@@ -31,7 +30,7 @@ from src.core.api_format import CORE_REDACT_HEADERS, merge_headers_with_protecti
from src.utils.ssl_utils import get_ssl_context
def _redact_headers(headers: Dict[str, str]) -> Dict[str, str]:
def _redact_headers(headers: dict[str, str]) -> dict[str, str]:
return redact_headers_for_log(headers, CORE_REDACT_HEADERS)
@@ -46,10 +45,10 @@ def _truncate_repr(value: Any, limit: int = 1200) -> str:
def build_safe_headers(
base_headers: Dict[str, str],
extra_headers: Optional[Dict[str, str]],
base_headers: dict[str, str],
extra_headers: dict[str, str] | None,
protected_keys: Iterable[str],
) -> Dict[str, str]:
) -> dict[str, str]:
"""
合并 extra_headers但防止覆盖 protected_keys大小写不敏感
"""
@@ -60,16 +59,16 @@ async def run_endpoint_check(
*,
client: httpx.AsyncClient, # 保持兼容性,但内部不使用
url: str,
headers: Dict[str, str],
json_body: Dict[str, Any],
headers: dict[str, str],
json_body: dict[str, Any],
api_format: str,
provider_name: Optional[str] = None,
model_name: Optional[str] = None,
api_key_id: Optional[str] = None,
provider_id: Optional[str] = None,
db: Optional[Any] = None, # Session对象需要时才导入
user: Optional[Any] = None, # User对象
) -> Dict[str, Any]:
provider_name: str | None = None,
model_name: str | None = None,
api_key_id: str | None = None,
provider_id: str | None = None,
db: Any | None = None, # Session对象需要时才导入
user: Any | None = None, # User对象
) -> dict[str, Any]:
"""
执行端点检查(重构版本,使用新的架构):
- 使用新的架构类来分离关注点
@@ -123,21 +122,21 @@ async def _calculate_and_record_usage(
provider_id: str,
api_key_id: str,
model_name: str,
request_data: Dict[str, Any],
response_data: Optional[Dict[str, Any]],
request_data: dict[str, Any],
response_data: dict[str, Any] | None,
request_id: str,
response_time_ms: int,
request_headers: Dict[str, str],
response_headers: Optional[Dict[str, str]] = None,
request_headers: dict[str, str],
response_headers: dict[str, str] | None = None,
status_code: int = 0,
error_message: Optional[str] = None,
error_message: str | None = None,
# 新增支持直接传递token数据
input_tokens: Optional[int] = None,
output_tokens: Optional[int] = None,
cache_creation_input_tokens: Optional[int] = None,
cache_read_input_tokens: Optional[int] = None,
api_format: Optional[str] = None,
) -> Dict[str, Any]:
input_tokens: int | None = None,
output_tokens: int | None = None,
cache_creation_input_tokens: int | None = None,
cache_read_input_tokens: int | None = None,
api_format: str | None = None,
) -> dict[str, Any]:
"""
计算并记录用量数据(遗留函数)
@@ -149,7 +148,7 @@ async def _calculate_and_record_usage(
"""
from src.services.usage.service import UsageService
from src.services.request.candidate import RequestCandidateService
from src.models.database import ApiKey, ProviderAPIKey, ProviderEndpoint
from src.models.database import ApiKey, ProviderAPIKey
# 获取Provider API Key对象不是用户API Key
provider_api_key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == api_key_id).first()
@@ -360,7 +359,7 @@ async def _calculate_and_record_usage(
}
def _extract_tokens_from_response(api_identifier: str, response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
def _extract_tokens_from_response(api_identifier: str, response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
"""
从响应中提取Token计数信息
@@ -446,7 +445,7 @@ def _extract_tokens_from_response(api_identifier: str, response_data: Optional[D
def _fallback_token_counting(request_data: Dict[str, Any], response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
def _fallback_token_counting(request_data: dict[str, Any], response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
"""
回退的Token计数方法简单估算
@@ -508,16 +507,16 @@ def _fallback_token_counting(request_data: Dict[str, Any], response_data: Option
class EndpointCheckRequest:
"""端点检查请求数据类"""
url: str
headers: Dict[str, str]
json_body: Dict[str, Any]
headers: dict[str, str]
json_body: dict[str, Any]
api_format: str
provider_name: Optional[str] = None
model_name: Optional[str] = None
api_key_id: Optional[str] = None
provider_id: Optional[str] = None
db: Optional[Any] = None
user: Optional[Any] = None
request_id: Optional[str] = None
provider_name: str | None = None
model_name: str | None = None
api_key_id: str | None = None
provider_id: str | None = None
db: Any | None = None
user: Any | None = None
request_id: str | None = None
timeout: float = 30.0
@@ -525,12 +524,12 @@ class EndpointCheckRequest:
class EndpointCheckResult:
"""端点检查结果数据类"""
status_code: int
headers: Dict[str, str]
headers: dict[str, str]
response_time_ms: int
request_id: str
response_data: Optional[Dict[str, Any]] = None
error_message: Optional[str] = None
usage_data: Optional[Dict[str, Any]] = None
response_data: dict[str, Any] | None = None
error_message: str | None = None
usage_data: dict[str, Any] | None = None
class HttpRequestExecutor:
@@ -613,7 +612,7 @@ class UsageCalculator:
return _extract_tokens_from_response(api_identifier, result.response_data)
@staticmethod
def _fallback_token_counting(request_data: Dict[str, Any], response_data: Optional[Dict[str, Any]]) -> tuple[int, int, int, int]:
def _fallback_token_counting(request_data: dict[str, Any], response_data: dict[str, Any] | None) -> tuple[int, int, int, int]:
"""回退的Token计数方法简单估算"""
# 估算输入Token
messages = request_data.get("messages", request_data.get("contents", []))
@@ -665,12 +664,12 @@ class AsyncBatchUsageRecorder:
def __init__(self, batch_size: int = 10, flush_interval: float = 2.0):
self.batch_size = batch_size
self.flush_interval = flush_interval
self.pending_records: List[Dict[str, Any]] = []
self._flush_task: Optional[asyncio.Task] = None
self.pending_records: list[dict[str, Any]] = []
self._flush_task: asyncio.Task | None = None
self._lock = asyncio.Lock()
self._running = True
async def add_record(self, usage_data: Dict[str, Any]) -> None:
async def add_record(self, usage_data: dict[str, Any]) -> None:
"""添加用量记录到批处理队列"""
async with self._lock:
self.pending_records.append(usage_data)
@@ -740,7 +739,7 @@ class AsyncBatchUsageRecorder:
# 全局批处理器实例(单例)
_global_batch_recorder: Optional[AsyncBatchUsageRecorder] = None
_global_batch_recorder: AsyncBatchUsageRecorder | None = None
def get_batch_recorder() -> AsyncBatchUsageRecorder:
"""获取全局批处理器实例"""
@@ -756,7 +755,7 @@ def get_batch_recorder() -> AsyncBatchUsageRecorder:
class EndpointCheckError(Exception):
"""端点检查错误基类"""
def __init__(self, message: str, error_type: str, status_code: int = 500, details: Optional[Dict[str, Any]] = None):
def __init__(self, message: str, error_type: str, status_code: int = 500, details: dict[str, Any] | None = None):
super().__init__(message)
self.message = message
self.error_type = error_type
@@ -765,22 +764,22 @@ class EndpointCheckError(Exception):
class NetworkError(EndpointCheckError):
"""网络请求错误"""
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
def __init__(self, message: str, details: dict[str, Any] | None = None):
super().__init__(message, "network_error", 0, details)
class AuthenticationError(EndpointCheckError):
"""认证错误"""
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
def __init__(self, message: str, details: dict[str, Any] | None = None):
super().__init__(message, "authentication_error", 401, details)
class RateLimitError(EndpointCheckError):
"""速率限制错误"""
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
def __init__(self, message: str, details: dict[str, Any] | None = None):
super().__init__(message, "rate_limit_error", 429, details)
class UpstreamError(EndpointCheckError):
"""上游服务错误"""
def __init__(self, message: str, status_code: int, details: Optional[Dict[str, Any]] = None):
def __init__(self, message: str, status_code: int, details: dict[str, Any] | None = None):
super().__init__(message, "upstream_error", status_code, details)
@@ -982,7 +981,7 @@ class EndpointCheckConfig:
retry_on_timeouts: bool = True
@classmethod
def from_env(cls) -> 'EndpointCheckConfig':
def from_env(cls) -> EndpointCheckConfig:
"""从环境变量创建配置"""
import os
@@ -1004,7 +1003,7 @@ class EndpointCheckConfig:
)
@classmethod
def from_dict(cls, config_dict: Dict[str, Any]) -> 'EndpointCheckConfig':
def from_dict(cls, config_dict: dict[str, Any]) -> EndpointCheckConfig:
"""从字典创建配置"""
return cls(**{k: v for k, v in config_dict.items() if hasattr(cls, k)})
@@ -1012,7 +1011,7 @@ class EndpointCheckConfig:
class ConfigurableEndpointChecker:
"""可配置的端点检查器"""
def __init__(self, config: Optional[EndpointCheckConfig] = None):
def __init__(self, config: EndpointCheckConfig | None = None):
self.config = config or EndpointCheckConfig()
self.executor = HttpRequestExecutor(timeout=self.config.timeout)
self.usage_calculator = UsageCalculator()
@@ -1171,9 +1170,9 @@ class ConfigurableEndpointChecker:
# 全局配置检查器实例
_global_configured_checker: Optional[ConfigurableEndpointChecker] = None
_global_configured_checker: ConfigurableEndpointChecker | None = None
def get_configured_checker(config: Optional[EndpointCheckConfig] = None) -> ConfigurableEndpointChecker:
def get_configured_checker(config: EndpointCheckConfig | None = None) -> ConfigurableEndpointChecker:
"""获取全局配置检查器实例"""
global _global_configured_checker
if _global_configured_checker is None or config is not None:
@@ -1186,8 +1185,8 @@ def get_configured_checker(config: Optional[EndpointCheckConfig] = None) -> Conf
class EndpointCheckOrchestrator:
"""端点检查协调器 - 协调整个流程"""
def __init__(self, executor: Optional[HttpRequestExecutor] = None,
usage_calculator: Optional[UsageCalculator] = None):
def __init__(self, executor: HttpRequestExecutor | None = None,
usage_calculator: UsageCalculator | None = None):
self.executor = executor or HttpRequestExecutor()
self.usage_calculator = usage_calculator or UsageCalculator()

View File

@@ -6,7 +6,7 @@
"""
import re
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
from src.api.handlers.base.response_parser import (
ParsedChunk,
@@ -20,7 +20,7 @@ from src.api.handlers.base.utils import extract_cache_creation_tokens
from src.core.api_format import is_cli_format
def _check_nested_error(response: Dict[str, Any]) -> Tuple[bool, Optional[Dict[str, Any]]]:
def _check_nested_error(response: dict[str, Any]) -> tuple[bool, dict[str, Any] | None]:
"""
检查响应中是否存在嵌套错误(某些代理服务返回 HTTP 200 但在响应体中包含错误)
@@ -62,7 +62,7 @@ def _check_nested_error(response: Dict[str, Any]) -> Tuple[bool, Optional[Dict[s
return False, None
def _extract_embedded_status_code(error_info: Optional[Dict[str, Any]]) -> Optional[int]:
def _extract_embedded_status_code(error_info: dict[str, Any] | None) -> int | None:
"""
从错误信息中提取嵌套的状态码
@@ -137,7 +137,7 @@ class OpenAIResponseParser(ResponseParser):
self.name = "OPENAI"
self.api_format = "OPENAI"
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
if not line or not line.strip():
return None
@@ -186,7 +186,7 @@ class OpenAIResponseParser(ResponseParser):
return chunk
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
result = ParsedResponse(
raw_response=response,
status_code=status_code,
@@ -217,7 +217,7 @@ class OpenAIResponseParser(ResponseParser):
return result
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
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),
@@ -226,7 +226,7 @@ class OpenAIResponseParser(ResponseParser):
"cache_read_tokens": 0,
}
def extract_text_content(self, response: Dict[str, Any]) -> str:
def extract_text_content(self, response: dict[str, Any]) -> str:
choices = response.get("choices", [])
if choices:
message = choices[0].get("message", {})
@@ -235,7 +235,7 @@ class OpenAIResponseParser(ResponseParser):
return content
return ""
def is_error_response(self, response: Dict[str, Any]) -> bool:
def is_error_response(self, response: dict[str, Any]) -> bool:
is_error, _ = _check_nested_error(response)
return is_error
@@ -259,7 +259,7 @@ class ClaudeResponseParser(ResponseParser):
self.name = "CLAUDE"
self.api_format = "CLAUDE"
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
if not line or not line.strip():
return None
@@ -324,7 +324,7 @@ class ClaudeResponseParser(ResponseParser):
return chunk
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
result = ParsedResponse(
raw_response=response,
status_code=status_code,
@@ -358,7 +358,7 @@ class ClaudeResponseParser(ResponseParser):
return result
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
# 对于 message_start 事件usage 在 message.usage 路径下
# 对于其他响应usage 在顶层
usage = response.get("usage") or {}
@@ -372,7 +372,7 @@ class ClaudeResponseParser(ResponseParser):
"cache_read_tokens": usage.get("cache_read_input_tokens", 0),
}
def extract_text_content(self, response: Dict[str, Any]) -> str:
def extract_text_content(self, response: dict[str, Any]) -> str:
content = response.get("content", [])
if isinstance(content, list):
text_parts = []
@@ -382,7 +382,7 @@ class ClaudeResponseParser(ResponseParser):
return "".join(text_parts)
return ""
def is_error_response(self, response: Dict[str, Any]) -> bool:
def is_error_response(self, response: dict[str, Any]) -> bool:
is_error, _ = _check_nested_error(response)
return is_error
@@ -406,7 +406,7 @@ class GeminiResponseParser(ResponseParser):
self.name = "GEMINI"
self.api_format = "GEMINI"
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
"""
解析 Gemini SSE 行
@@ -473,7 +473,7 @@ class GeminiResponseParser(ResponseParser):
return chunk
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
result = ParsedResponse(
raw_response=response,
status_code=status_code,
@@ -509,7 +509,7 @@ class GeminiResponseParser(ResponseParser):
return result
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
"""
从 Gemini 响应中提取 token 使用量
@@ -531,7 +531,7 @@ class GeminiResponseParser(ResponseParser):
"cache_read_tokens": usage.get("cached_tokens", 0),
}
def extract_text_content(self, response: Dict[str, Any]) -> str:
def extract_text_content(self, response: dict[str, Any]) -> str:
candidates = response.get("candidates", [])
if candidates:
content = candidates[0].get("content", {})
@@ -543,7 +543,7 @@ class GeminiResponseParser(ResponseParser):
return "".join(text_parts)
return ""
def is_error_response(self, response: Dict[str, Any]) -> bool:
def is_error_response(self, response: dict[str, Any]) -> bool:
"""
判断响应是否为错误响应
@@ -562,7 +562,7 @@ class GeminiCliResponseParser(GeminiResponseParser):
# 解析器注册表
_PARSERS: Dict[str, Type[ResponseParser]] = {
_PARSERS: dict[str, type[ResponseParser]] = {
"CLAUDE": ClaudeResponseParser,
"CLAUDE_CLI": ClaudeCliResponseParser,
"OPENAI": OpenAIResponseParser,

View File

@@ -11,12 +11,11 @@
payload, headers = builder.build(original_body, original_headers, endpoint, key)
"""
from __future__ import annotations
import json
from abc import ABC, abstractmethod
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Dict, FrozenSet, Optional, Tuple
from typing import TYPE_CHECKING, Any
from src.core.crypto import crypto_service
from src.core.api_format import HeaderBuilder, UPSTREAM_DROP_HEADERS
@@ -37,9 +36,9 @@ class ProviderAuthInfo:
auth_header: str
auth_value: str
# 解密后的认证配置(用于 URL 构建等场景,避免重复解密)
decrypted_auth_config: Optional[Dict[str, Any]] = None
decrypted_auth_config: dict[str, Any] | None = None
def as_tuple(self) -> Tuple[str, str]:
def as_tuple(self) -> tuple[str, str]:
"""返回 (auth_header, auth_value) 元组"""
return (self.auth_header, self.auth_value)
@@ -48,7 +47,7 @@ class ProviderAuthInfo:
# ==============================================================================
# 兼容别名:历史代码使用 SENSITIVE_HEADERS 命名
SENSITIVE_HEADERS: FrozenSet[str] = UPSTREAM_DROP_HEADERS
SENSITIVE_HEADERS: frozenset[str] = UPSTREAM_DROP_HEADERS
# ==============================================================================
@@ -57,14 +56,14 @@ SENSITIVE_HEADERS: FrozenSet[str] = UPSTREAM_DROP_HEADERS
# 标准测试请求体OpenAI 格式)
# 用于 check_endpoint 等测试场景,使用简单安全的消息内容避免触发安全过滤
DEFAULT_TEST_REQUEST: Dict[str, Any] = {
DEFAULT_TEST_REQUEST: dict[str, Any] = {
"messages": [{"role": "user", "content": "Hi"}],
"max_tokens": 5,
"temperature": 0,
}
def get_test_request_data(request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
def get_test_request_data(request_data: dict[str, Any] | None = None) -> dict[str, Any]:
"""获取测试请求数据
如果传入 request_data则合并到默认测试请求中
@@ -85,8 +84,8 @@ def get_test_request_data(request_data: Optional[Dict[str, Any]] = None) -> Dict
def build_test_request_body(
format_id: str,
request_data: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
request_data: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""构建测试请求体,自动处理格式转换
使用格式转换注册表将 OpenAI 格式的测试请求转换为目标格式。
@@ -127,39 +126,39 @@ class RequestBuilder(ABC):
@abstractmethod
def build_payload(
self,
original_body: Dict[str, Any],
original_body: dict[str, Any],
*,
mapped_model: Optional[str] = None,
mapped_model: str | None = None,
is_stream: bool = False,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""构建请求体"""
pass
@abstractmethod
def build_headers(
self,
original_headers: Dict[str, str],
original_headers: dict[str, str],
endpoint: Any,
key: Any,
*,
extra_headers: Optional[Dict[str, str]] = None,
pre_computed_auth: Optional[Tuple[str, str]] = None,
) -> Dict[str, str]:
extra_headers: dict[str, str] | None = None,
pre_computed_auth: tuple[str, str] | None = None,
) -> dict[str, str]:
"""构建请求头"""
pass
def build(
self,
original_body: Dict[str, Any],
original_headers: Dict[str, str],
original_body: dict[str, Any],
original_headers: dict[str, str],
endpoint: Any,
key: Any,
*,
mapped_model: Optional[str] = None,
mapped_model: str | None = None,
is_stream: bool = False,
extra_headers: Optional[Dict[str, str]] = None,
pre_computed_auth: Optional[Tuple[str, str]] = None,
) -> Tuple[Dict[str, Any], Dict[str, str]]:
extra_headers: dict[str, str] | None = None,
pre_computed_auth: tuple[str, str] | None = None,
) -> tuple[dict[str, Any], dict[str, str]]:
"""
构建完整的请求(请求体 + 请求头)
@@ -202,11 +201,11 @@ class PassthroughRequestBuilder(RequestBuilder):
def build_payload(
self,
original_body: Dict[str, Any],
original_body: dict[str, Any],
*,
mapped_model: Optional[str] = None, # noqa: ARG002 - 由 apply_mapped_model 处理
mapped_model: str | None = None, # noqa: ARG002 - 由 apply_mapped_model 处理
is_stream: bool = False, # noqa: ARG002 - 保留原始值,不自动添加
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
透传请求体 - 原样复制,不做任何修改
@@ -218,13 +217,13 @@ class PassthroughRequestBuilder(RequestBuilder):
def build_headers(
self,
original_headers: Dict[str, str],
original_headers: dict[str, str],
endpoint: Any,
key: Any,
*,
extra_headers: Optional[Dict[str, str]] = None,
pre_computed_auth: Optional[Tuple[str, str]] = None,
) -> Dict[str, str]:
extra_headers: dict[str, str] | None = None,
pre_computed_auth: tuple[str, str] | None = None,
) -> dict[str, str]:
"""
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
@@ -289,11 +288,11 @@ class PassthroughRequestBuilder(RequestBuilder):
def build_passthrough_request(
original_body: Dict[str, Any],
original_headers: Dict[str, str],
original_body: dict[str, Any],
original_headers: dict[str, str],
endpoint: Any,
key: Any,
) -> Tuple[Dict[str, Any], Dict[str, str]]:
) -> tuple[dict[str, Any], dict[str, str]]:
"""
构建透传模式的请求

View File

@@ -4,7 +4,7 @@
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Any, Dict, Optional
from typing import Any
@dataclass
@@ -13,14 +13,14 @@ class ParsedChunk:
# 原始数据
raw_line: str
event_type: Optional[str] = None
data: Optional[Dict[str, Any]] = None
event_type: str | None = None
data: dict[str, Any] | None = None
# 提取的内容
text_delta: str = ""
is_done: bool = False
is_error: bool = False
error_message: Optional[str] = None
error_message: str | None = None
# 使用量信息(通常在最后一个 chunk 中)
input_tokens: int = 0
@@ -29,7 +29,7 @@ class ParsedChunk:
cache_read_tokens: int = 0
# 响应 ID
response_id: Optional[str] = None
response_id: str | None = None
@dataclass
@@ -48,21 +48,21 @@ class StreamStats:
# 内容
collected_text: str = ""
response_id: Optional[str] = None
response_id: str | None = None
# 状态
has_completion: bool = False
status_code: int = 200
error_message: Optional[str] = None
error_message: str | None = None
# Provider 信息
provider_name: Optional[str] = None
endpoint_id: Optional[str] = None
key_id: Optional[str] = None
provider_name: str | None = None
endpoint_id: str | None = None
key_id: str | None = None
# 响应头和完整响应
response_headers: Dict[str, str] = field(default_factory=dict)
final_response: Optional[Dict[str, Any]] = None
response_headers: dict[str, str] = field(default_factory=dict)
final_response: dict[str, Any] | None = None
@dataclass
@@ -70,12 +70,12 @@ class ParsedResponse:
"""解析后的非流式响应"""
# 原始响应
raw_response: Dict[str, Any]
raw_response: dict[str, Any]
status_code: int
# 提取的内容
text_content: str = ""
response_id: Optional[str] = None
response_id: str | None = None
# 使用量
input_tokens: int = 0
@@ -85,10 +85,10 @@ class ParsedResponse:
# 错误信息
is_error: bool = False
error_type: Optional[str] = None
error_message: Optional[str] = None
error_type: str | None = None
error_message: str | None = None
# 从响应体解析出的嵌套状态码(当 HTTP 200 但响应体含错误时使用)
embedded_status_code: Optional[int] = None
embedded_status_code: int | None = None
class ResponseParser(ABC):
@@ -106,7 +106,7 @@ class ResponseParser(ABC):
api_format: str = "UNKNOWN"
@abstractmethod
def parse_sse_line(self, line: str, stats: StreamStats) -> Optional[ParsedChunk]:
def parse_sse_line(self, line: str, stats: StreamStats) -> ParsedChunk | None:
"""
解析单行 SSE 数据
@@ -120,7 +120,7 @@ class ResponseParser(ABC):
pass
@abstractmethod
def parse_response(self, response: Dict[str, Any], status_code: int) -> ParsedResponse:
def parse_response(self, response: dict[str, Any], status_code: int) -> ParsedResponse:
"""
解析非流式响应
@@ -134,7 +134,7 @@ class ResponseParser(ABC):
pass
@abstractmethod
def extract_usage_from_response(self, response: Dict[str, Any]) -> Dict[str, int]:
def extract_usage_from_response(self, response: dict[str, Any]) -> dict[str, int]:
"""
从响应中提取 token 使用量
@@ -147,7 +147,7 @@ class ResponseParser(ABC):
pass
@abstractmethod
def extract_text_content(self, response: Dict[str, Any]) -> str:
def extract_text_content(self, response: dict[str, Any]) -> str:
"""
从响应中提取文本内容
@@ -159,7 +159,7 @@ class ResponseParser(ABC):
"""
pass
def is_error_response(self, response: Dict[str, Any]) -> bool:
def is_error_response(self, response: dict[str, Any]) -> bool:
"""
判断响应是否为错误响应

View File

@@ -8,9 +8,11 @@
- 请求/响应数据
"""
from __future__ import annotations
import time
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any, Dict, List, Optional
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from src.core.api_format.conversion.stream_state import StreamState
@@ -35,16 +37,16 @@ class StreamContext:
api_key_id: int = 0
# Provider 信息(在请求执行时填充)
provider_name: Optional[str] = None
provider_id: Optional[str] = None
endpoint_id: Optional[str] = None
key_id: Optional[str] = None
attempt_id: Optional[str] = None
provider_name: str | None = None
provider_id: 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: Optional[str] = None # Provider 的响应格式
provider_api_format: str | None = None # Provider 的响应格式
# 模型映射
mapped_model: Optional[str] = None
mapped_model: str | None = None
# Token 统计
input_tokens: int = 0
@@ -53,33 +55,33 @@ class StreamContext:
cache_creation_tokens: int = 0
# 响应内容
_collected_text_parts: List[str] = field(default_factory=list, repr=False)
response_id: Optional[str] = None
final_usage: Optional[Dict[str, Any]] = None
final_response: Optional[Dict[str, Any]] = None
_collected_text_parts: list[str] = field(default_factory=list, 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: Optional[int] = None # 首字时间 (TTFB - Time To First Byte)
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: Optional[str] = None # 客户端友好的错误消息
upstream_response: Optional[str] = None # 原始 Provider 响应(用于请求链路追踪)
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: Optional[Dict[str, Any]] = None
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)
response_metadata: dict[str, Any] = field(default_factory=dict)
# 整流标记Thinking Rectifier
rectified: bool = False # 请求是否经过整流(移除 thinking 块后重试)
@@ -87,10 +89,10 @@ class StreamContext:
# 流式处理统计
data_count: int = 0
chunk_count: int = 0
parsed_chunks: List[Dict[str, Any]] = field(default_factory=list)
parsed_chunks: list[dict[str, Any]] = field(default_factory=list)
# 流式格式转换状态(跨 chunk 追踪)
stream_conversion_state: Optional["StreamState"] = None
stream_conversion_state: StreamState | None = None
def reset_for_retry(self) -> None:
"""
@@ -138,7 +140,7 @@ class StreamContext:
provider_id: str,
endpoint_id: str,
key_id: str,
provider_api_format: Optional[str] = None,
provider_api_format: str | None = None,
) -> None:
"""更新 Provider 信息"""
self.provider_name = provider_name
@@ -149,10 +151,10 @@ class StreamContext:
def update_usage(
self,
input_tokens: Optional[int] = None,
output_tokens: Optional[int] = None,
cached_tokens: Optional[int] = None,
cache_creation_tokens: Optional[int] = None,
input_tokens: int | None = None,
output_tokens: int | None = None,
cached_tokens: int | None = None,
cache_creation_tokens: int | None = None,
) -> None:
"""
更新 Token 使用统计
@@ -194,7 +196,7 @@ class StreamContext:
self,
status_code: int,
error_message: str,
upstream_response: Optional[str] = None,
upstream_response: str | None = None,
) -> None:
"""
标记请求失败
@@ -230,7 +232,7 @@ class StreamContext:
"""检查是否因客户端断开连接而结束"""
return self.status_code == 499
def build_response_body(self, response_time_ms: int) -> Dict[str, Any]:
def build_response_body(self, response_time_ms: int) -> dict[str, Any]:
"""
构建响应体元数据

View File

@@ -13,7 +13,10 @@ import asyncio
import codecs
import json
from dataclasses import dataclass
from typing import Any, AsyncGenerator, Callable, Optional
from typing import Any
from collections.abc import Callable
from collections.abc import AsyncGenerator
import httpx
@@ -65,10 +68,10 @@ class StreamProcessor:
self,
request_id: str,
default_parser: ResponseParser,
on_streaming_start: Optional[Callable[[], None]] = None,
on_streaming_start: Callable[[], None] | None = None,
*,
collect_text: bool = False,
smoothing_config: Optional[StreamSmoothingConfig] = None,
smoothing_config: StreamSmoothingConfig | None = None,
):
"""
初始化流处理器
@@ -105,7 +108,7 @@ class StreamProcessor:
def handle_sse_event(
self,
ctx: StreamContext,
event_name: Optional[str],
event_name: str | None,
data_str: str,
*,
skip_record: bool = False,
@@ -363,7 +366,7 @@ class StreamProcessor:
):
# 重新抛出可重试的 Provider 异常,触发故障转移
raise
except (OSError, IOError) as e:
except OSError as e:
# 网络 I/O 异常:记录警告,可能需要重试
logger.warning(f" [{self.request_id}] 预读流时发生网络异常: {type(e).__name__}: {e}")
except Exception as e:
@@ -382,10 +385,10 @@ class StreamProcessor:
byte_iterator: Any,
response_ctx: Any,
http_client: httpx.AsyncClient,
prefetched_chunks: Optional[list] = None,
prefetched_chunks: list | None = None,
*,
start_time: Optional[float] = None,
) -> AsyncGenerator[bytes, None]:
start_time: float | None = None,
) -> AsyncGenerator[bytes]:
"""
创建响应流生成器
@@ -547,9 +550,7 @@ class StreamProcessor:
logger.warning(f"[{self.request_id}] 流式格式转换失败: {conv_err}")
payload = _build_stream_error_payload("响应格式转换失败,请稍后重试")
error_bytes = (
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode(
"utf-8"
)
f"data: {json.dumps(payload, ensure_ascii=False)}\n\n".encode()
)
done_bytes = (
b"data: [DONE]\n\n" if client_format.startswith("OPENAI") else b""
@@ -570,7 +571,7 @@ class StreamProcessor:
# 统一使用 SSE 格式输出Gemini streamGenerateContent 也使用 SSE
# 参考: https://ai.google.dev/api/generate-content
out.append(
f"data: {json.dumps(evt, ensure_ascii=False)}\n\n".encode("utf-8")
f"data: {json.dumps(evt, ensure_ascii=False)}\n\n".encode()
)
return out
@@ -769,9 +770,9 @@ class StreamProcessor:
async def create_monitored_stream(
self,
ctx: StreamContext,
stream_generator: AsyncGenerator[bytes, None],
stream_generator: AsyncGenerator[bytes],
is_disconnected: Callable[[], Any],
) -> AsyncGenerator[bytes, None]:
) -> AsyncGenerator[bytes]:
"""
创建带监控的流生成器
@@ -833,8 +834,8 @@ class StreamProcessor:
async def create_smoothed_stream(
self,
stream_generator: AsyncGenerator[bytes, None],
) -> AsyncGenerator[bytes, None]:
stream_generator: AsyncGenerator[bytes],
) -> AsyncGenerator[bytes]:
"""
创建平滑输出的流生成器
@@ -933,7 +934,7 @@ class StreamProcessor:
if buffer:
yield buffer
def _get_extractor(self, format_name: str) -> Optional[ContentExtractor]:
def _get_extractor(self, format_name: str) -> ContentExtractor | None:
"""获取或创建格式对应的提取器(带缓存)"""
if format_name not in self._extractors:
extractor = get_extractor(format_name)
@@ -943,7 +944,7 @@ class StreamProcessor:
def _detect_format_and_extract(
self, data: dict
) -> tuple[Optional[str], Optional[ContentExtractor]]:
) -> tuple[str | None, ContentExtractor | None]:
"""
检测数据格式并提取内容
@@ -998,10 +999,10 @@ class StreamProcessor:
async def create_smoothed_stream(
stream_generator: AsyncGenerator[bytes, None],
stream_generator: AsyncGenerator[bytes],
chunk_size: int = 20,
delay_ms: int = 8,
) -> AsyncGenerator[bytes, None]:
) -> AsyncGenerator[bytes]:
"""
独立的平滑流生成函数
@@ -1032,7 +1033,7 @@ class _LightweightSmoother:
self.delay_ms = delay_ms
self._extractors: dict[str, ContentExtractor] = {}
def _get_extractor(self, format_name: str) -> Optional[ContentExtractor]:
def _get_extractor(self, format_name: str) -> ContentExtractor | None:
if format_name not in self._extractors:
extractor = get_extractor(format_name)
if extractor:
@@ -1041,7 +1042,7 @@ class _LightweightSmoother:
def _detect_format_and_extract(
self, data: dict
) -> tuple[Optional[str], Optional[ContentExtractor]]:
) -> tuple[str | None, ContentExtractor | None]:
for format_name in get_extractor_formats():
extractor = self._get_extractor(format_name)
if extractor:
@@ -1060,8 +1061,8 @@ class _LightweightSmoother:
return [content[i : i + self.chunk_size] for i in range(0, text_length, self.chunk_size)]
async def smooth(
self, stream_generator: AsyncGenerator[bytes, None]
) -> AsyncGenerator[bytes, None]:
self, stream_generator: AsyncGenerator[bytes]
) -> AsyncGenerator[bytes]:
buffer = b""
is_first_content = True

View File

@@ -9,7 +9,7 @@
import asyncio
import time
from typing import Any, Dict, Optional
from typing import Any
from sqlalchemy.orm import Session
@@ -58,8 +58,8 @@ class StreamTelemetryRecorder:
async def record_stream_stats(
self,
ctx: StreamContext,
original_headers: Dict[str, str],
original_request_body: Dict[str, Any],
original_headers: dict[str, str],
original_request_body: dict[str, Any],
start_time: float,
) -> None:
"""
@@ -144,9 +144,9 @@ class StreamTelemetryRecorder:
self,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: Dict[str, str],
actual_request_body: Dict[str, Any],
response_body: Optional[Dict[str, Any]],
original_headers: dict[str, str],
actual_request_body: dict[str, Any],
response_body: dict[str, Any] | None,
response_time_ms: int,
) -> None:
"""记录成功的请求"""
@@ -193,9 +193,9 @@ class StreamTelemetryRecorder:
self,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: Dict[str, str],
actual_request_body: Dict[str, Any],
response_body: Optional[Dict[str, Any]],
original_headers: dict[str, str],
actual_request_body: dict[str, Any],
response_body: dict[str, Any] | None,
response_time_ms: int,
) -> None:
"""记录失败的请求"""
@@ -236,9 +236,9 @@ class StreamTelemetryRecorder:
self,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: Dict[str, str],
actual_request_body: Dict[str, Any],
response_body: Optional[Dict[str, Any]],
original_headers: dict[str, str],
actual_request_body: dict[str, Any],
response_body: dict[str, Any] | None,
response_time_ms: int,
) -> None:
"""记录客户端取消的请求"""
@@ -285,7 +285,7 @@ class StreamTelemetryRecorder:
from src.services.request.candidate import RequestCandidateService
extra_data: Dict[str, Any] = {
extra_data: dict[str, Any] = {
"stream_completed": ctx.is_success(),
"data_count": ctx.data_count,
}
@@ -358,7 +358,7 @@ class StreamTelemetryRecorder:
status: str,
response_time_ms: int,
status_code: int = 200,
error_message: Optional[str] = None,
error_message: str | None = None,
) -> None:
"""直接更新 Usage 表的状态字段"""
try:
@@ -378,7 +378,7 @@ class StreamTelemetryRecorder:
async def _get_telemetry_writer(
self, bg_db: Session, ctx: StreamContext, response_time_ms: int
) -> Optional[TelemetryWriter]:
) -> TelemetryWriter | None:
if config.usage_queue_enabled and self.user_id and self.api_key_id:
return QueueTelemetryWriter(
request_id=self.request_id,
@@ -400,9 +400,9 @@ class StreamTelemetryRecorder:
self,
writer: TelemetryWriter,
ctx: StreamContext,
original_headers: Dict[str, str],
actual_request_body: Dict[str, Any],
response_body: Optional[Dict[str, Any]],
original_headers: dict[str, str],
actual_request_body: dict[str, Any],
response_body: dict[str, Any] | None,
response_time_ms: int,
) -> None:
"""根据上下文状态分发到对应的记录方法"""
@@ -430,7 +430,7 @@ class StreamTelemetryRecorder:
return "cancelled"
return "failed"
def _build_db_writer(self, bg_db: Session) -> Optional[DbTelemetryWriter]:
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()

View File

@@ -4,8 +4,9 @@ Handler 基础工具函数
from __future__ import annotations
import json
from typing import TYPE_CHECKING, Any, Dict, Optional
from typing import TYPE_CHECKING, Any
from src.core.exceptions import EmbeddedErrorException, ProviderNotAvailableException
from src.core.api_format import filter_response_headers
@@ -15,7 +16,7 @@ if TYPE_CHECKING:
from src.core.api_format.conversion.registry import FormatConversionRegistry
def get_format_converter_registry() -> "FormatConversionRegistry":
def get_format_converter_registry() -> FormatConversionRegistry:
"""
获取格式转换注册表(线程安全)
@@ -31,7 +32,7 @@ def get_format_converter_registry() -> "FormatConversionRegistry":
return format_conversion_registry
def extract_cache_creation_tokens(usage: Dict[str, Any]) -> int:
def extract_cache_creation_tokens(usage: dict[str, Any]) -> int:
"""
提取缓存创建 tokens兼容三种格式
@@ -99,7 +100,7 @@ def extract_cache_creation_tokens(usage: Dict[str, Any]) -> int:
return old_format
def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[str, str]:
def build_sse_headers(extra_headers: dict[str, str] | None = None) -> dict[str, str]:
"""
构建 SSEtext/event-stream推荐响应头用于减少代理缓冲带来的卡顿/成段输出。
@@ -107,7 +108,7 @@ def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[st
- Cache-Control: no-transform 可避免部分代理对流做压缩/改写导致缓冲
- X-Accel-Buffering: no 可显式提示 Nginx 关闭缓冲(即使全局已关闭也无害)
"""
headers: Dict[str, str] = {
headers: dict[str, str] = {
"Cache-Control": "no-cache, no-transform",
"X-Accel-Buffering": "no",
}
@@ -116,7 +117,7 @@ def build_sse_headers(extra_headers: Optional[Dict[str, str]] = None) -> Dict[st
return headers
def filter_proxy_response_headers(headers: Optional[Dict[str, str]]) -> Dict[str, str]:
def filter_proxy_response_headers(headers: dict[str, str] | None) -> dict[str, str]:
"""
过滤上游响应头中不应透传给客户端的字段。
@@ -148,8 +149,8 @@ def check_prefetched_response_error(
parser: Any,
request_id: str,
provider_name: str,
endpoint_id: Optional[str],
base_url: Optional[str],
endpoint_id: str | None,
base_url: str | None,
) -> None:
"""
检查预读的响应是否为非 SSE 格式的错误响应HTML 或纯 JSON 错误)

View File

@@ -4,7 +4,7 @@ Claude Chat Adapter - 基于 ChatAdapterBase 的 Claude Chat API 适配器
处理 /v1/messages 端点的 Claude Chat 格式请求。
"""
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
from fastapi import HTTPException, Request
@@ -25,9 +25,9 @@ class ClaudeCapabilityDetector:
@staticmethod
def detect_from_headers(
headers: Dict[str, str],
request_body: Optional[Dict[str, Any]] = None,
) -> Dict[str, bool]:
headers: dict[str, str],
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""
从 Claude 请求头检测能力需求
@@ -38,7 +38,7 @@ class ClaudeCapabilityDetector:
headers: 请求头字典
request_body: 请求体Claude 不使用,保留用于接口统一)
"""
requirements: Dict[str, bool] = {}
requirements: dict[str, bool] = {}
# 使用统一的大小写不敏感获取
beta_header = get_header_value(headers, "anthropic-beta")
@@ -61,21 +61,21 @@ class ClaudeChatAdapter(ChatAdapterBase):
name = "claude.chat"
@property
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.claude.handler import ClaudeChatHandler
return ClaudeChatHandler
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats or ["CLAUDE"])
logger.info(f"[{self.name}] 初始化Chat模式适配器 | API格式: {self.allowed_api_formats}")
def detect_capability_requirements(
self,
headers: Dict[str, str],
request_body: Optional[Dict[str, Any]] = None,
) -> Dict[str, bool]:
headers: dict[str, str],
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""检测 Claude 请求中隐含的能力需求"""
return ClaudeCapabilityDetector.detect_from_headers(headers)
@@ -124,7 +124,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
)
return request
def _build_audit_metadata(self, _payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
def _build_audit_metadata(self, _payload: dict[str, Any], request_obj) -> dict[str, Any]:
"""构建 Claude Chat 特定的审计元数据"""
role_counts: dict[str, int] = {}
for message in request_obj.messages:
@@ -153,8 +153,8 @@ class ClaudeChatAdapter(ChatAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 Claude API 支持的模型列表"""
headers = cls.build_headers_with_extra(api_key, extra_headers)
@@ -201,7 +201,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE
def build_claude_adapter(x_app_header: Optional[str]):
def build_claude_adapter(x_app_header: str | None):
"""根据 x-app 头部构造 Chat 或 Claude Code 适配器。"""
if x_app_header and x_app_header.lower() == "cli":
from src.api.handlers.claude_cli.adapter import ClaudeCliAdapter
@@ -216,7 +216,7 @@ class ClaudeTokenCountAdapter(ApiAdapter):
name = "claude.token_count"
mode = ApiMode.STANDARD
def extract_api_key(self, request: Request) -> Optional[str]:
def extract_api_key(self, request: Request) -> str | None:
"""从请求中提取 API 密钥 (x-api-key 或 Authorization: Bearer)"""
# 优先检查 x-api-key
api_key = request.headers.get("x-api-key")

View File

@@ -5,7 +5,7 @@ Claude Chat Handler - 基于通用 Chat Handler 基类的简化实现
代码量从原来的 ~1470 行减少到 ~120 行。
"""
from typing import Any, Dict, Optional
from typing import Any
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.api.handlers.base.utils import extract_cache_creation_tokens
@@ -25,8 +25,8 @@ class ClaudeChatHandler(ChatHandlerBase):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> str:
"""
从请求中提取模型名 - Claude 格式实现
@@ -45,9 +45,9 @@ class ClaudeChatHandler(ChatHandlerBase):
def apply_mapped_model(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
mapped_model: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
将映射后的模型名应用到请求体
@@ -90,7 +90,7 @@ class ClaudeChatHandler(ChatHandlerBase):
return request
def _extract_usage(self, response: Dict) -> Dict[str, int]:
def _extract_usage(self, response: dict) -> dict[str, int]:
"""
从 Claude 响应中提取 token 使用情况
@@ -108,7 +108,7 @@ class ClaudeChatHandler(ChatHandlerBase):
"cache_read_input_tokens": usage.get("cache_read_input_tokens", 0),
}
def _normalize_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
def _normalize_response(self, response: dict[str, Any]) -> dict[str, Any]:
"""
规范化 Claude 响应

View File

@@ -4,10 +4,9 @@ Claude SSE 流解析器
解析 Claude Messages API 的 Server-Sent Events 流。
"""
from __future__ import annotations
import json
from typing import Any, Dict, List, Optional
from typing import Any
from src.api.handlers.base.utils import extract_cache_creation_tokens
@@ -43,7 +42,7 @@ class ClaudeStreamParser:
DELTA_TEXT = "text_delta"
DELTA_INPUT_JSON = "input_json_delta"
def parse_chunk(self, chunk: bytes | str) -> List[Dict[str, Any]]:
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
"""
解析 SSE 数据块
@@ -58,10 +57,10 @@ class ClaudeStreamParser:
else:
text = chunk
events: List[Dict[str, Any]] = []
events: list[dict[str, Any]] = []
lines = text.strip().split("\n")
current_event_type: Optional[str] = None
current_event_type: str | None = None
for line in lines:
line = line.strip()
@@ -96,7 +95,7 @@ class ClaudeStreamParser:
return events
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
def parse_line(self, line: str) -> dict[str, Any] | None:
"""
解析单行 SSE 数据
@@ -117,7 +116,7 @@ class ClaudeStreamParser:
except json.JSONDecodeError:
return None
def is_done_event(self, event: Dict[str, Any]) -> bool:
def is_done_event(self, event: dict[str, Any]) -> bool:
"""
判断是否为结束事件
@@ -130,7 +129,7 @@ class ClaudeStreamParser:
event_type = event.get("type")
return event_type in (self.EVENT_MESSAGE_STOP, "__done__")
def is_error_event(self, event: Dict[str, Any]) -> bool:
def is_error_event(self, event: dict[str, Any]) -> bool:
"""
判断是否为错误事件
@@ -142,7 +141,7 @@ class ClaudeStreamParser:
"""
return event.get("type") == self.EVENT_ERROR
def get_event_type(self, event: Dict[str, Any]) -> Optional[str]:
def get_event_type(self, event: dict[str, Any]) -> str | None:
"""
获取事件类型
@@ -155,7 +154,7 @@ class ClaudeStreamParser:
event_type = event.get("type")
return str(event_type) if event_type is not None else None
def extract_text_delta(self, event: Dict[str, Any]) -> Optional[str]:
def extract_text_delta(self, event: dict[str, Any]) -> str | None:
"""
从 content_block_delta 事件中提取文本增量
@@ -175,7 +174,7 @@ class ClaudeStreamParser:
return None
def extract_usage(self, event: Dict[str, Any]) -> Optional[Dict[str, int]]:
def extract_usage(self, event: dict[str, Any]) -> dict[str, int] | None:
"""
从事件中提取 token 使用量
@@ -212,7 +211,7 @@ class ClaudeStreamParser:
return None
def extract_message_id(self, event: Dict[str, Any]) -> Optional[str]:
def extract_message_id(self, event: dict[str, Any]) -> str | None:
"""
从 message_start 事件中提取消息 ID
@@ -229,7 +228,7 @@ class ClaudeStreamParser:
msg_id = message.get("id")
return str(msg_id) if msg_id is not None else None
def extract_stop_reason(self, event: Dict[str, Any]) -> Optional[str]:
def extract_stop_reason(self, event: dict[str, Any]) -> str | None:
"""
从 message_delta 事件中提取停止原因

View File

@@ -4,7 +4,7 @@ Claude CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
继承 CliAdapterBase只需配置 FORMAT_ID 和 HANDLER_CLASS。
"""
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
@@ -27,20 +27,20 @@ class ClaudeCliAdapter(CliAdapterBase):
name = "claude.cli"
@property
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.claude_cli.handler import ClaudeCliMessageHandler
return ClaudeCliMessageHandler
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats or ["CLAUDE_CLI"])
def detect_capability_requirements(
self,
headers: Dict[str, str],
request_body: Optional[Dict[str, Any]] = None,
) -> Dict[str, bool]:
headers: dict[str, str],
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""检测 Claude CLI 请求中隐含的能力需求"""
return ClaudeCapabilityDetector.detect_from_headers(headers)
@@ -61,16 +61,16 @@ class ClaudeCliAdapter(CliAdapterBase):
"""
return input_tokens + cache_creation_input_tokens + cache_read_input_tokens
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
def _extract_message_count(self, payload: dict[str, Any]) -> int:
"""Claude CLI 使用 messages 字段"""
messages = payload.get("messages", [])
return len(messages) if isinstance(messages, list) else 0
def _build_audit_metadata(
self,
payload: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
) -> Dict[str, Any]:
payload: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> dict[str, Any]:
"""Claude CLI 特定的审计元数据"""
model = payload.get("model", "unknown")
stream = payload.get("stream", False)
@@ -104,8 +104,8 @@ class ClaudeCliAdapter(CliAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 Claude API 支持的模型列表(带 CLI User-Agent"""
# 复用 ClaudeChatAdapter 的实现,添加 CLI User-Agent
cli_headers = {"User-Agent": config.internal_user_agent_claude_cli}
@@ -120,7 +120,7 @@ class ClaudeCliAdapter(CliAdapterBase):
return models, error
@classmethod
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
"""构建Claude CLI API端点URL"""
base_url = base_url.rstrip("/")
if base_url.endswith("/v1"):
@@ -131,12 +131,12 @@ class ClaudeCliAdapter(CliAdapterBase):
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE_CLI
@classmethod
def get_cli_user_agent(cls) -> Optional[str]:
def get_cli_user_agent(cls) -> str | None:
"""获取Claude CLI User-Agent"""
return config.internal_user_agent_claude_cli
@classmethod
def get_cli_extra_headers(cls) -> Dict[str, str]:
def get_cli_extra_headers(cls) -> dict[str, str]:
"""获取Claude CLI额外请求头包含 x-app: cli 标识"""
headers = super().get_cli_extra_headers()
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的认证方式

View File

@@ -4,7 +4,7 @@ Claude CLI Message Handler - 基于通用 CLI Handler 基类的简化实现
继承 CliMessageHandlerBase只需覆盖格式特定的配置和事件处理逻辑。
"""
from typing import Any, Dict, Optional
from typing import Any
from src.api.handlers.base.cli_handler_base import (
CliMessageHandlerBase,
@@ -33,8 +33,8 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> str:
"""
从请求中提取模型名 - Claude 格式实现
@@ -53,9 +53,9 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
def apply_mapped_model(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
mapped_model: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
Claude API 的 model 在请求体顶级
@@ -74,7 +74,7 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
self,
ctx: StreamContext,
event_type: str,
data: Dict[str, Any],
data: dict[str, Any],
) -> None:
"""
处理 Claude CLI 格式的 SSE 事件
@@ -142,8 +142,8 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
def _extract_response_metadata(
self,
response: Dict[str, Any],
) -> Dict[str, Any]:
response: dict[str, Any],
) -> dict[str, Any]:
"""
从 Claude 响应中提取元数据
@@ -155,7 +155,7 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
Returns:
提取的元数据字典
"""
metadata: Dict[str, Any] = {}
metadata: dict[str, Any] = {}
# 提取模型名称(实际使用的模型)
if "model" in response:

View File

@@ -4,7 +4,7 @@ Gemini Chat Adapter
处理 Gemini API 格式的请求适配
"""
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
from fastapi import HTTPException, Request
@@ -33,17 +33,17 @@ class GeminiChatAdapter(ChatAdapterBase):
name = "gemini.chat"
@property
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.gemini.handler import GeminiChatHandler
return GeminiChatHandler
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats or ["GEMINI"])
logger.info(f"[{self.name}] 初始化 Gemini Chat 适配器 | API格式: {self.allowed_api_formats}")
def extract_api_key(self, request: Request) -> Optional[str]:
def extract_api_key(self, request: Request) -> str | None:
"""
从请求中提取 API 密钥 - Gemini 支持 header 和 query 两种方式
@@ -68,8 +68,8 @@ class GeminiChatAdapter(ChatAdapterBase):
return {}
def _merge_path_params(
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
) -> Dict[str, Any]:
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
) -> dict[str, Any]:
"""
合并 URL 路径参数到请求体 - Gemini 特化版本
@@ -122,14 +122,14 @@ class GeminiChatAdapter(ChatAdapterBase):
request.stream = is_stream
return request
def _extract_message_count(self, payload: Dict[str, Any], request_obj) -> int:
def _extract_message_count(self, payload: dict[str, Any], request_obj) -> int:
"""提取消息数量"""
contents = payload.get("contents", [])
if hasattr(request_obj, "contents"):
contents = request_obj.contents
return len(contents) if isinstance(contents, list) else 0
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
"""构建 Gemini Chat 特定的审计元数据"""
role_counts: dict[str, int] = {}
@@ -182,8 +182,8 @@ class GeminiChatAdapter(ChatAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 Gemini API 支持的模型列表"""
# Gemini 使用 URL 参数传递 key不需要 headers 中的认证
base_url_clean = base_url.rstrip("/")
@@ -192,7 +192,7 @@ class GeminiChatAdapter(ChatAdapterBase):
else:
models_url = f"{base_url_clean}/v1beta/models?key={api_key}"
headers: Dict[str, str] = {}
headers: dict[str, str] = {}
if extra_headers:
headers.update(extra_headers)
@@ -242,16 +242,16 @@ class GeminiChatAdapter(ChatAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
request_data: Dict[str, Any],
extra_headers: Optional[Dict[str, str]] = None,
request_data: dict[str, Any],
extra_headers: dict[str, str] | None = None,
# 用量计算参数
db: Optional[Any] = None,
user: Optional[Any] = None,
provider_name: Optional[str] = None,
provider_id: Optional[str] = None,
api_key_id: Optional[str] = None,
model_name: Optional[str] = None,
) -> Dict[str, Any]:
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,
) -> dict[str, Any]:
"""测试 Gemini API 模型连接性(非流式)"""
# Gemini需要从request_data或model_name参数获取model名称
effective_model_name = model_name or request_data.get("model", "")

View File

@@ -4,7 +4,7 @@ Gemini Chat Handler
处理 Gemini API 格式的请求
"""
from typing import Any, Dict, Optional
from typing import Any
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
@@ -76,8 +76,8 @@ class GeminiChatHandler(ChatHandlerBase):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None,
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None,
) -> str:
"""
从请求中提取模型名 - Gemini Chat 格式实现
@@ -126,7 +126,7 @@ class GeminiChatHandler(ChatHandlerBase):
return request
def _extract_usage(self, response: Dict) -> Dict[str, int]:
def _extract_usage(self, response: dict) -> dict[str, int]:
"""
从 Gemini 响应中提取 token 使用情况
@@ -151,7 +151,7 @@ class GeminiChatHandler(ChatHandlerBase):
"cache_read_input_tokens": usage.get("cached_tokens", 0),
}
def _normalize_response(self, response: Dict) -> Dict:
def _normalize_response(self, response: dict) -> dict:
"""
规范化 Gemini 响应

View File

@@ -15,7 +15,7 @@ Gemini streamGenerateContent 常见两种返回(与上游/代理实现有关
"""
import json
from typing import Any, Dict, List, Optional, Union
from typing import Any
class GeminiStreamParser:
@@ -43,7 +43,7 @@ class GeminiStreamParser:
self._in_array = False
self._brace_depth = 0
def parse_chunk(self, chunk: Union[bytes, str]) -> List[Dict[str, Any]]:
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
"""
解析流式数据块
@@ -58,7 +58,7 @@ class GeminiStreamParser:
else:
text = chunk
events: List[Dict[str, Any]] = []
events: list[dict[str, Any]] = []
for char in text:
if char == "[" and not self._in_array:
@@ -97,7 +97,7 @@ class GeminiStreamParser:
return events
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
def parse_line(self, line: str) -> dict[str, Any] | None:
"""
解析单行 JSON 数据
@@ -118,7 +118,7 @@ class GeminiStreamParser:
except json.JSONDecodeError:
return None
def is_done_event(self, event: Dict[str, Any]) -> bool:
def is_done_event(self, event: dict[str, Any]) -> bool:
"""
判断是否为结束事件
@@ -143,7 +143,7 @@ class GeminiStreamParser:
return False
def is_error_event(self, event: Dict[str, Any]) -> bool:
def is_error_event(self, event: dict[str, Any]) -> bool:
"""
判断是否为错误事件
@@ -171,7 +171,7 @@ class GeminiStreamParser:
return False
def extract_error_info(self, event: Dict[str, Any]) -> Optional[Dict[str, Any]]:
def extract_error_info(self, event: dict[str, Any]) -> dict[str, Any] | None:
"""
从事件中提取错误信息
@@ -208,7 +208,7 @@ class GeminiStreamParser:
return None
def get_finish_reason(self, event: Dict[str, Any]) -> Optional[str]:
def get_finish_reason(self, event: dict[str, Any]) -> str | None:
"""
获取结束原因
@@ -224,7 +224,7 @@ class GeminiStreamParser:
return str(reason) if reason is not None else None
return None
def extract_text_delta(self, event: Dict[str, Any]) -> Optional[str]:
def extract_text_delta(self, event: dict[str, Any]) -> str | None:
"""
从响应中提取文本内容
@@ -248,7 +248,7 @@ class GeminiStreamParser:
return "".join(text_parts) if text_parts else None
def extract_usage(self, event: Dict[str, Any]) -> Optional[Dict[str, int]]:
def extract_usage(self, event: dict[str, Any]) -> dict[str, int] | None:
"""
从事件中提取 token 使用量
@@ -280,7 +280,7 @@ class GeminiStreamParser:
"cached_tokens": usage_metadata.get("cachedContentTokenCount", 0),
}
def extract_model_version(self, event: Dict[str, Any]) -> Optional[str]:
def extract_model_version(self, event: dict[str, Any]) -> str | None:
"""
从响应中提取模型版本
@@ -293,7 +293,7 @@ class GeminiStreamParser:
version = event.get("modelVersion")
return str(version) if version is not None else None
def extract_safety_ratings(self, event: Dict[str, Any]) -> Optional[List[Dict[str, Any]]]:
def extract_safety_ratings(self, event: dict[str, Any]) -> list[dict[str, Any]] | None:
"""
从响应中提取安全评级

View File

@@ -4,7 +4,7 @@ Gemini CLI Adapter - 基于通用 CLI Adapter 基类的实现
继承 CliAdapterBase处理 Gemini CLI 格式的请求。
"""
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
from fastapi import Request
@@ -29,16 +29,16 @@ class GeminiCliAdapter(CliAdapterBase):
name = "gemini.cli"
@property
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.gemini_cli.handler import GeminiCliMessageHandler
return GeminiCliMessageHandler
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats or ["GEMINI_CLI"])
def extract_api_key(self, request: Request) -> Optional[str]:
def extract_api_key(self, request: Request) -> str | None:
"""
从请求中提取 API 密钥 - Gemini CLI 支持 header 和 query 两种方式
@@ -53,8 +53,8 @@ class GeminiCliAdapter(CliAdapterBase):
)
def _merge_path_params(
self, original_request_body: Dict[str, Any], path_params: Dict[str, Any] # noqa: ARG002
) -> Dict[str, Any]:
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
) -> dict[str, Any]:
"""
合并 URL 路径参数到请求体 - Gemini CLI 特化版本
@@ -74,23 +74,23 @@ class GeminiCliAdapter(CliAdapterBase):
# Gemini: 不合并任何 path_params 到请求体
return original_request_body.copy()
def _extract_message_count(self, payload: Dict[str, Any]) -> int:
def _extract_message_count(self, payload: dict[str, Any]) -> int:
"""Gemini CLI 使用 contents 字段"""
contents = payload.get("contents", [])
return len(contents) if isinstance(contents, list) else 0
def _build_audit_metadata(
self,
payload: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
payload: dict[str, Any],
path_params: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""Gemini CLI 特定的审计元数据"""
# 从 path_params 获取 modelGemini 请求体不含 model
model = path_params.get("model", "unknown") if path_params else "unknown"
contents = payload.get("contents", [])
generation_config = payload.get("generation_config", {}) or {}
role_counts: Dict[str, int] = {}
role_counts: dict[str, int] = {}
for content in contents:
role = content.get("role", "unknown") if isinstance(content, dict) else "unknown"
role_counts[role] = role_counts.get(role, 0) + 1
@@ -120,8 +120,8 @@ class GeminiCliAdapter(CliAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 Gemini API 支持的模型列表(带 CLI User-Agent"""
# 复用 GeminiChatAdapter 的实现,添加 CLI User-Agent
cli_headers = {"User-Agent": config.internal_user_agent_gemini_cli}
@@ -136,7 +136,7 @@ class GeminiCliAdapter(CliAdapterBase):
return models, error
@classmethod
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
"""构建Gemini CLI API端点URL"""
effective_model_name = model_name or request_data.get("model", "")
if not effective_model_name:
@@ -152,12 +152,12 @@ class GeminiCliAdapter(CliAdapterBase):
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> GEMINI_CLI
@classmethod
def get_cli_user_agent(cls) -> Optional[str]:
def get_cli_user_agent(cls) -> str | None:
"""获取Gemini CLI User-Agent"""
return config.internal_user_agent_gemini_cli
@classmethod
def get_cli_extra_headers(cls) -> Dict[str, str]:
def get_cli_extra_headers(cls) -> dict[str, str]:
"""获取Gemini CLI额外请求头包含 x-app: cli 标识"""
headers = super().get_cli_extra_headers()
headers["x-app"] = "cli" # 标识 CLI 模式,让上游使用正确的 adapter

View File

@@ -4,7 +4,7 @@ Gemini CLI Message Handler - 基于通用 CLI Handler 基类的实现
继承 CliMessageHandlerBase处理 Gemini CLI API 格式的请求。
"""
from typing import Any, Dict, Optional
from typing import Any
from src.api.handlers.base.cli_handler_base import (
CliMessageHandlerBase,
@@ -34,8 +34,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
def extract_model_from_request(
self,
request_body: Dict[str, Any], # noqa: ARG002 - 基类签名要求
path_params: Optional[Dict[str, Any]] = None,
request_body: dict[str, Any], # noqa: ARG002 - 基类签名要求
path_params: dict[str, Any] | None = None,
) -> str:
"""
从请求中提取模型名 - Gemini 格式实现
@@ -57,8 +57,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
def prepare_provider_request_body(
self,
request_body: Dict[str, Any],
) -> Dict[str, Any]:
request_body: dict[str, Any],
) -> dict[str, Any]:
"""
准备发送给 Gemini API 的请求体 - 移除 model 字段
@@ -77,9 +77,9 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
def get_model_for_url(
self,
request_body: Dict[str, Any],
mapped_model: Optional[str],
) -> Optional[str]:
request_body: dict[str, Any],
mapped_model: str | None,
) -> str | None:
"""
Gemini 需要将 model 放入 URL 路径中
@@ -93,7 +93,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
# 优先使用映射后的模型名,否则使用请求体中的
return mapped_model or request_body.get("model")
def _extract_usage_from_event(self, event: Dict[str, Any]) -> Dict[str, int]:
def _extract_usage_from_event(self, event: dict[str, Any]) -> dict[str, int]:
"""
从 Gemini 事件中提取 token 使用情况
@@ -126,7 +126,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
self,
ctx: StreamContext,
_event_type: str,
data: Dict[str, Any],
data: dict[str, Any],
) -> None:
"""
处理 Gemini CLI 格式的流式事件
@@ -190,8 +190,8 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
def _extract_response_metadata(
self,
response: Dict[str, Any],
) -> Dict[str, Any]:
response: dict[str, Any],
) -> dict[str, Any]:
"""
从 Gemini 响应中提取元数据
@@ -203,7 +203,7 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
Returns:
包含 model_version 的元数据字典
"""
metadata: Dict[str, Any] = {}
metadata: dict[str, Any] = {}
model_version = response.get("modelVersion")
if model_version:
metadata["model_version"] = model_version

View File

@@ -4,7 +4,7 @@ OpenAI Chat Adapter - 基于 ChatAdapterBase 的 OpenAI Chat API 适配器
处理 /v1/chat/completions 端点的 OpenAI Chat 格式请求。
"""
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
from fastapi.responses import JSONResponse
@@ -28,13 +28,13 @@ class OpenAIChatAdapter(ChatAdapterBase):
name = "openai.chat"
@property
def HANDLER_CLASS(self) -> Type[ChatHandlerBase]:
def HANDLER_CLASS(self) -> type[ChatHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.openai.handler import OpenAIChatHandler
return OpenAIChatHandler
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats or ["OPENAI"])
def _validate_request_body(self, original_request_body: dict, path_params: dict = None):
@@ -66,7 +66,7 @@ class OpenAIChatAdapter(ChatAdapterBase):
max_tokens=original_request_body.get("max_tokens"),
)
def _build_audit_metadata(self, payload: Dict[str, Any], request_obj) -> Dict[str, Any]:
def _build_audit_metadata(self, payload: dict[str, Any], request_obj) -> dict[str, Any]:
"""构建 OpenAI Chat 特定的审计元数据"""
role_counts = {}
for message in request_obj.messages:
@@ -105,8 +105,8 @@ class OpenAIChatAdapter(ChatAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 OpenAI 兼容 API 支持的模型列表"""
headers = cls.build_headers_with_extra(api_key, extra_headers)

View File

@@ -5,7 +5,7 @@ OpenAI Chat Handler - 基于通用 Chat Handler 基类的简化实现
代码量从原来的 ~1315 行减少到 ~100 行。
"""
from typing import Any, Dict, Optional
from typing import Any
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
@@ -24,8 +24,8 @@ class OpenAIChatHandler(ChatHandlerBase):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> str:
"""
从请求中提取模型名 - OpenAI 格式实现
@@ -44,9 +44,9 @@ class OpenAIChatHandler(ChatHandlerBase):
def apply_mapped_model(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
mapped_model: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
将映射后的模型名应用到请求体
@@ -89,7 +89,7 @@ class OpenAIChatHandler(ChatHandlerBase):
return request
def _extract_usage(self, response: Dict) -> Dict[str, int]:
def _extract_usage(self, response: dict) -> dict[str, int]:
"""
从 OpenAI 响应中提取 token 使用情况
@@ -106,7 +106,7 @@ class OpenAIChatHandler(ChatHandlerBase):
"cache_read_input_tokens": 0,
}
def _normalize_response(self, response: Dict) -> Dict:
def _normalize_response(self, response: dict) -> dict:
"""
规范化 OpenAI 响应

View File

@@ -4,10 +4,9 @@ OpenAI SSE 流解析器
解析 OpenAI Chat Completions API 的 Server-Sent Events 流。
"""
from __future__ import annotations
import json
from typing import Any, Dict, List, Optional
from typing import Any
class OpenAIStreamParser:
@@ -23,7 +22,7 @@ class OpenAIStreamParser:
- 流结束时发送 data: [DONE]
"""
def parse_chunk(self, chunk: bytes | str) -> List[Dict[str, Any]]:
def parse_chunk(self, chunk: bytes | str) -> list[dict[str, Any]]:
"""
解析 SSE 数据块
@@ -38,7 +37,7 @@ class OpenAIStreamParser:
else:
text = chunk
chunks: List[Dict[str, Any]] = []
chunks: list[dict[str, Any]] = []
lines = text.strip().split("\n")
for line in lines:
@@ -64,7 +63,7 @@ class OpenAIStreamParser:
return chunks
def parse_line(self, line: str) -> Optional[Dict[str, Any]]:
def parse_line(self, line: str) -> dict[str, Any] | None:
"""
解析单行 SSE 数据
@@ -85,7 +84,7 @@ class OpenAIStreamParser:
except json.JSONDecodeError:
return None
def is_done_chunk(self, chunk: Dict[str, Any]) -> bool:
def is_done_chunk(self, chunk: dict[str, Any]) -> bool:
"""
判断是否为结束 chunk
@@ -107,7 +106,7 @@ class OpenAIStreamParser:
return False
def get_finish_reason(self, chunk: Dict[str, Any]) -> Optional[str]:
def get_finish_reason(self, chunk: dict[str, Any]) -> str | None:
"""
获取结束原因
@@ -123,7 +122,7 @@ class OpenAIStreamParser:
return str(reason) if reason is not None else None
return None
def extract_text_delta(self, chunk: Dict[str, Any]) -> Optional[str]:
def extract_text_delta(self, chunk: dict[str, Any]) -> str | None:
"""
从 chunk 中提取文本增量
@@ -145,7 +144,7 @@ class OpenAIStreamParser:
return None
def extract_tool_calls_delta(self, chunk: Dict[str, Any]) -> Optional[List[Dict[str, Any]]]:
def extract_tool_calls_delta(self, chunk: dict[str, Any]) -> list[dict[str, Any]] | None:
"""
从 chunk 中提取工具调用增量
@@ -165,7 +164,7 @@ class OpenAIStreamParser:
return tool_calls
return None
def extract_role(self, chunk: Dict[str, Any]) -> Optional[str]:
def extract_role(self, chunk: dict[str, Any]) -> str | None:
"""
从 chunk 中提取角色

View File

@@ -4,7 +4,7 @@ OpenAI CLI Adapter - 基于通用 CLI Adapter 基类的简化实现
继承 CliAdapterBase只需配置 FORMAT_ID 和 HANDLER_CLASS。
"""
from typing import Any, Dict, Optional, Tuple, Type
from typing import Any
import httpx
@@ -27,13 +27,13 @@ class OpenAICliAdapter(CliAdapterBase):
name = "openai.cli"
@property
def HANDLER_CLASS(self) -> Type[CliMessageHandlerBase]:
def HANDLER_CLASS(self) -> type[CliMessageHandlerBase]:
"""延迟导入 Handler 类避免循环依赖"""
from src.api.handlers.openai_cli.handler import OpenAICliMessageHandler
return OpenAICliMessageHandler
def __init__(self, allowed_api_formats: Optional[list[str]] = None):
def __init__(self, allowed_api_formats: list[str] | None = None):
super().__init__(allowed_api_formats or ["OPENAI_CLI"])
# =========================================================================
@@ -46,8 +46,8 @@ class OpenAICliAdapter(CliAdapterBase):
client: httpx.AsyncClient,
base_url: str,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Tuple[list, Optional[str]]:
extra_headers: dict[str, str] | None = None,
) -> tuple[list, str | None]:
"""查询 OpenAI 兼容 API 支持的模型列表(带 CLI User-Agent"""
# 复用 OpenAIChatAdapter 的实现,添加 CLI User-Agent
cli_headers = {"User-Agent": config.internal_user_agent_openai_cli}
@@ -62,7 +62,7 @@ class OpenAICliAdapter(CliAdapterBase):
return models, error
@classmethod
def build_endpoint_url(cls, base_url: str, request_data: Dict[str, Any], model_name: Optional[str] = None) -> str:
def build_endpoint_url(cls, base_url: str, request_data: dict[str, Any], model_name: str | None = None) -> str:
"""构建OpenAI CLI API端点URL"""
base_url = base_url.rstrip("/")
if base_url.endswith("/v1"):
@@ -74,7 +74,7 @@ class OpenAICliAdapter(CliAdapterBase):
# OPENAI -> OPENAI_CLI 无转换器,会直接透传原始请求
@classmethod
def get_cli_user_agent(cls) -> Optional[str]:
def get_cli_user_agent(cls) -> str | None:
"""获取OpenAI CLI User-Agent"""
return config.internal_user_agent_openai_cli

View File

@@ -5,7 +5,7 @@ OpenAI CLI Message Handler - 基于通用 CLI Handler 基类的简化实现
代码量从原来的 900+ 行减少到 ~100 行。
"""
from typing import Any, Dict, Optional
from typing import Any
from src.api.handlers.base.cli_handler_base import (
CliMessageHandlerBase,
@@ -32,8 +32,8 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
def extract_model_from_request(
self,
request_body: Dict[str, Any],
path_params: Optional[Dict[str, Any]] = None, # noqa: ARG002
request_body: dict[str, Any],
path_params: dict[str, Any] | None = None, # noqa: ARG002
) -> str:
"""
从请求中提取模型名 - OpenAI 格式实现
@@ -52,9 +52,9 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
def apply_mapped_model(
self,
request_body: Dict[str, Any],
request_body: dict[str, Any],
mapped_model: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
OpenAI CLI (Responses API) 的 model 在请求体顶级字段。
@@ -73,7 +73,7 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
self,
ctx: StreamContext,
event_type: str,
data: Dict[str, Any],
data: dict[str, Any],
) -> None:
"""
处理 OpenAI CLI 格式的 SSE 事件
@@ -144,8 +144,8 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
def _extract_response_metadata(
self,
response: Dict[str, Any],
) -> Dict[str, Any]:
response: dict[str, Any],
) -> dict[str, Any]:
"""
从 OpenAI 响应中提取元数据
@@ -157,7 +157,7 @@ class OpenAICliMessageHandler(CliMessageHandlerBase):
Returns:
提取的元数据字典
"""
metadata: Dict[str, Any] = {}
metadata: dict[str, Any] = {}
# 提取模型名称(实际使用的模型)
if "model" in response: