chore: 升级到 Python 3.14 并现代化代码

- 升级 Docker 基础镜像从 Python 3.12 到 3.14
- 更新 pyproject.toml 支持 Python 3.13/3.14
- 移除 Python 3.8/3.9/3.10/3.11 分类器
- 更新 black 和 mypy 配置目标版本
- 将 get_event_loop() 替换为 get_running_loop() 加上 RuntimeError 处理
- 简化 compute_cost_sync 中的 asyncio.run 使用
- Dict/List/Tuple/Set → dict/list/tuple/set (PEP 585)
- Optional[T] → T | None (PEP 604)
- Union[A, B] → A | B (PEP 604)
- 移除废弃的 typing 导入
- 移除不必要的字符串引号注解
This commit is contained in:
AAEE86
2026-01-30 03:10:21 +08:00
parent 3e75bc8964
commit 24d24f6829
255 changed files with 4062 additions and 4173 deletions

View File

@@ -10,10 +10,9 @@
- data_format_id 不同 -> 需要转换,检查全局开关 + 端点配置 + 转换器能力
"""
from __future__ import annotations
import logging
from typing import TYPE_CHECKING, Optional, Tuple
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from src.core.api_format.conversion.registry import FormatConversionRegistry
@@ -26,11 +25,11 @@ logger = logging.getLogger(__name__)
def is_format_compatible(
client_format: str,
endpoint_api_format: str,
endpoint_format_acceptance_config: Optional[dict],
endpoint_format_acceptance_config: dict | None,
is_stream: bool,
global_conversion_enabled: bool,
registry: Optional["FormatConversionRegistry"] = None,
) -> Tuple[bool, bool, Optional[str]]:
registry: FormatConversionRegistry | None = None,
) -> tuple[bool, bool, str | None]:
"""
检查端点是否兼容客户端格式

View File

@@ -4,7 +4,6 @@
用于严格模式转换失败时抛出,让编排器可以尝试下一个候选。
"""
from __future__ import annotations
class FormatConversionError(Exception):

View File

@@ -9,13 +9,11 @@
这些应复用 `src/core/api_format/metadata.py`API_FORMAT_DEFINITIONS作为单一事实来源。
"""
from __future__ import annotations
from typing import Dict, Set
# 角色映射仅作为辅助system/tool 的具体落点以 Normalizer 规则为准)
ROLE_MAPPINGS: Dict[str, Dict[str, str]] = {
ROLE_MAPPINGS: dict[str, dict[str, str]] = {
"OPENAI": {
"user": "user",
"assistant": "assistant",
@@ -29,7 +27,7 @@ ROLE_MAPPINGS: Dict[str, Dict[str, str]] = {
# 停止原因映射internal -> provider未知值使用 UNKNOWN 并写入 extra/raw
STOP_REASON_MAPPINGS: Dict[str, Dict[str, str]] = {
STOP_REASON_MAPPINGS: dict[str, dict[str, str]] = {
"CLAUDE": {
"end_turn": "end_turn",
"max_tokens": "max_tokens",
@@ -60,7 +58,7 @@ STOP_REASON_MAPPINGS: Dict[str, Dict[str, str]] = {
# 使用量字段映射provider usage field -> internal UsageInfo field
USAGE_FIELD_MAPPINGS: Dict[str, Dict[str, str]] = {
USAGE_FIELD_MAPPINGS: dict[str, dict[str, str]] = {
"CLAUDE": {
"input_tokens": "input_tokens",
"output_tokens": "output_tokens",
@@ -82,7 +80,7 @@ USAGE_FIELD_MAPPINGS: Dict[str, Dict[str, str]] = {
# 错误类型映射provider -> internal ErrorType.value
ERROR_TYPE_MAPPINGS: Dict[str, Dict[str, str]] = {
ERROR_TYPE_MAPPINGS: dict[str, dict[str, str]] = {
"CLAUDE": {
"invalid_request_error": "invalid_request",
"authentication_error": "authentication",
@@ -116,7 +114,7 @@ ERROR_TYPE_MAPPINGS: Dict[str, Dict[str, str]] = {
# 可重试的错误类型internal ErrorType.value
RETRYABLE_ERROR_TYPES: Set[str] = {
RETRYABLE_ERROR_TYPES: set[str] = {
"rate_limit",
"overloaded",
"server_error",

View File

@@ -10,11 +10,10 @@
- 兼容优先UnknownBlock 在内部保留,但默认在输出阶段丢弃(可观测、可随时调整策略)
"""
from __future__ import annotations
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Dict, FrozenSet, List, Optional, Union
from typing import Any
class Role(str, Enum):
@@ -65,7 +64,7 @@ class TextBlock:
type: ContentType = field(default=ContentType.TEXT, init=False)
text: str = ""
extra: Dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -74,11 +73,11 @@ class ImageBlock:
type: ContentType = field(default=ContentType.IMAGE, init=False)
# base64 编码的图片数据(二选一)
data: Optional[str] = None
media_type: Optional[str] = None
data: str | None = None
media_type: str | None = None
# 或者 URL 引用
url: Optional[str] = None
extra: Dict[str, Any] = field(default_factory=dict)
url: str | None = None
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -88,8 +87,8 @@ class ToolUseBlock:
type: ContentType = field(default=ContentType.TOOL_USE, init=False)
tool_id: str = ""
tool_name: str = ""
tool_input: Dict[str, Any] = field(default_factory=dict)
extra: Dict[str, Any] = field(default_factory=dict)
tool_input: dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -100,9 +99,9 @@ class ToolResultBlock:
tool_use_id: str = "" # 对应的 ToolUseBlock.tool_id
# 工具输出可能是纯文本,也可能是结构化 JSONGemini functionResponse 等)
output: Any = None
content_text: Optional[str] = None
content_text: str | None = None
is_error: bool = False
extra: Dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -111,11 +110,11 @@ class UnknownBlock:
type: ContentType = field(default=ContentType.UNKNOWN, init=False)
raw_type: str = "" # 原始的类型字符串(各格式不一致)
payload: Dict[str, Any] = field(default_factory=dict) # 原始结构(尽量保持)
extra: Dict[str, Any] = field(default_factory=dict)
payload: dict[str, Any] = field(default_factory=dict) # 原始结构(尽量保持)
extra: dict[str, Any] = field(default_factory=dict)
ContentBlock = Union[TextBlock, ImageBlock, ToolUseBlock, ToolResultBlock, UnknownBlock]
ContentBlock = TextBlock | ImageBlock | ToolUseBlock | ToolResultBlock | UnknownBlock
@dataclass
@@ -123,8 +122,8 @@ class InternalMessage:
"""统一的消息表示"""
role: Role
content: List[ContentBlock] # 统一使用列表,纯文本用单个 TextBlock
extra: Dict[str, Any] = field(default_factory=dict)
content: list[ContentBlock] # 统一使用列表,纯文本用单个 TextBlock
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -132,9 +131,9 @@ class ToolDefinition:
"""统一的工具定义"""
name: str
description: Optional[str] = None
parameters: Optional[Dict[str, Any]] = None # JSON Schema
extra: Dict[str, Any] = field(default_factory=dict)
description: str | None = None
parameters: dict[str, Any] | None = None # JSON Schema
extra: dict[str, Any] = field(default_factory=dict)
class ToolChoiceType(str, Enum):
@@ -149,8 +148,8 @@ class ToolChoice:
"""统一的工具选择"""
type: ToolChoiceType
tool_name: Optional[str] = None
extra: Dict[str, Any] = field(default_factory=dict)
tool_name: str | None = None
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -159,7 +158,7 @@ class InstructionSegment:
role: Role # 仅允许 Role.SYSTEM / Role.DEVELOPER
text: str = ""
extra: Dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -167,25 +166,25 @@ class InternalRequest:
"""统一的请求表示"""
model: str
messages: List[InternalMessage]
messages: list[InternalMessage]
# 指令层:保留 system/developer 结构与顺序
instructions: List[InstructionSegment] = field(default_factory=list)
instructions: list[InstructionSegment] = field(default_factory=list)
# 兼容字段instructions 的 join 文本(无 role 标签),用于 Claude/Gemini 这类仅接受字符串 system 的格式
system: Optional[str] = None
system: str | None = None
max_tokens: Optional[int] = None
temperature: Optional[float] = None
top_p: Optional[float] = None
top_k: Optional[int] = None
stop_sequences: Optional[List[str]] = None
max_tokens: int | None = None
temperature: float | None = None
top_p: float | None = None
top_k: int | None = None
stop_sequences: list[str] | None = None
stream: bool = False
tools: Optional[List[ToolDefinition]] = None
tool_choice: Optional[ToolChoice] = None # auto/none/required 或指定 tool_name
extra: Dict[str, Any] = field(default_factory=dict) # 未识别字段透传
tools: list[ToolDefinition] | None = None
tool_choice: ToolChoice | None = None # auto/none/required 或指定 tool_name
extra: dict[str, Any] = field(default_factory=dict) # 未识别字段透传
def to_debug_dict(self) -> Dict[str, Any]:
def to_debug_dict(self) -> dict[str, Any]:
"""用于日志和调试的简化表示"""
return {
"model": self.model,
@@ -208,7 +207,7 @@ class UsageInfo:
total_tokens: int = 0
cache_read_tokens: int = 0
cache_write_tokens: int = 0
extra: Dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -217,12 +216,12 @@ class InternalResponse:
id: str
model: str
content: List[ContentBlock]
stop_reason: Optional[StopReason] = None
usage: Optional[UsageInfo] = None
extra: Dict[str, Any] = field(default_factory=dict)
content: list[ContentBlock]
stop_reason: StopReason | None = None
usage: UsageInfo | None = None
extra: dict[str, Any] = field(default_factory=dict)
def to_debug_dict(self) -> Dict[str, Any]:
def to_debug_dict(self) -> dict[str, Any]:
"""用于日志和调试的简化表示"""
usage = None
if self.usage:
@@ -246,12 +245,12 @@ class InternalError:
type: ErrorType
message: str
code: Optional[str] = None # 原始错误码
param: Optional[str] = None # 导致错误的参数
code: str | None = None # 原始错误码
param: str | None = None # 导致错误的参数
retryable: bool = False # 是否可重试
extra: Dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
def to_debug_dict(self) -> Dict[str, Any]:
def to_debug_dict(self) -> dict[str, Any]:
"""用于日志和调试"""
return {
"type": self.type.value,
@@ -269,7 +268,7 @@ class FormatCapabilities:
supports_error_conversion: bool = True
supports_tools: bool = True
supports_images: bool = False
supported_features: FrozenSet[str] = field(default_factory=frozenset)
supported_features: frozenset[str] = field(default_factory=frozenset)
__all__ = [

View File

@@ -5,10 +5,9 @@
再从 internal 输出到目标格式。
"""
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import Any, Dict, List, Optional
from typing import Any
from .internal import FormatCapabilities, InternalError, InternalRequest, InternalResponse
from .stream_events import InternalStreamEvent
@@ -24,19 +23,19 @@ class FormatNormalizer(ABC):
# ============ 请求转换 ============
@abstractmethod
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
"""将格式特定请求转换为内部表示"""
raise NotImplementedError
@abstractmethod
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
"""将内部表示转换为格式特定请求"""
raise NotImplementedError
# ============ 响应转换 ============
@abstractmethod
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse:
def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
"""将格式特定响应转换为内部表示"""
raise NotImplementedError
@@ -45,8 +44,8 @@ class FormatNormalizer(ABC):
self,
internal: InternalResponse,
*,
requested_model: Optional[str] = None,
) -> Dict[str, Any]:
requested_model: str | None = None,
) -> dict[str, Any]:
"""将内部表示转换为格式特定响应
Args:
@@ -61,9 +60,9 @@ class FormatNormalizer(ABC):
def stream_chunk_to_internal(
self,
chunk: Dict[str, Any],
chunk: dict[str, Any],
state: StreamState,
) -> List[InternalStreamEvent]:
) -> list[InternalStreamEvent]:
"""将格式特定流式块转换为内部事件"""
raise NotImplementedError
@@ -71,21 +70,21 @@ class FormatNormalizer(ABC):
self,
event: InternalStreamEvent,
state: StreamState,
) -> List[Dict[str, Any]]:
) -> list[dict[str, Any]]:
"""将内部事件转换为格式特定流式块"""
raise NotImplementedError
# ============ 错误转换(可选) ============
def is_error_response(self, response: Dict[str, Any]) -> bool:
def is_error_response(self, response: dict[str, Any]) -> bool:
"""基于 body 的兜底判断(不可靠),子类可覆盖"""
return False
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError:
def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
"""将格式特定错误转换为内部表示"""
raise NotImplementedError
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]:
def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
"""将内部错误表示转换为格式特定错误"""
raise NotImplementedError

View File

@@ -6,7 +6,6 @@ Normalizers
本目录在 Phase 1 仅创建结构;具体实现将在 Phase 2+ 补齐。
"""
from __future__ import annotations
__all__: list[str] = []

View File

@@ -7,10 +7,9 @@ Claude Messages API Normalizer
- 可选Claude error <-> InternalError
"""
from __future__ import annotations
import json
from typing import Any, Dict, List, Optional, Tuple
from typing import Any
from src.core.api_format.conversion.field_mappings import (
ERROR_TYPE_MAPPINGS,
@@ -63,7 +62,7 @@ class ClaudeNormalizer(FormatNormalizer):
supports_images=True,
)
_CLAUDE_STOP_TO_INTERNAL: Dict[str, StopReason] = {
_CLAUDE_STOP_TO_INTERNAL: dict[str, StopReason] = {
"end_turn": StopReason.END_TURN,
"max_tokens": StopReason.MAX_TOKENS,
"stop_sequence": StopReason.STOP_SEQUENCE,
@@ -73,7 +72,7 @@ class ClaudeNormalizer(FormatNormalizer):
"content_filtered": StopReason.CONTENT_FILTERED,
}
_ERROR_TYPE_TO_CLAUDE: Dict[ErrorType, str] = {
_ERROR_TYPE_TO_CLAUDE: dict[ErrorType, str] = {
ErrorType.INVALID_REQUEST: "invalid_request_error",
ErrorType.AUTHENTICATION: "authentication_error",
ErrorType.PERMISSION_DENIED: "permission_error",
@@ -90,11 +89,11 @@ class ClaudeNormalizer(FormatNormalizer):
# Requests
# =========================
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
model = str(request.get("model") or "")
dropped: Dict[str, int] = {}
dropped: dict[str, int] = {}
instructions: List[InstructionSegment] = []
instructions: list[InstructionSegment] = []
# 顶层 system 先进入 instructions保持确定性优先级
sys_value = request.get("system")
@@ -103,7 +102,7 @@ class ClaudeNormalizer(FormatNormalizer):
if sys_text:
instructions.append(InstructionSegment(role=Role.SYSTEM, text=sys_text))
messages: List[InternalMessage] = []
messages: list[InternalMessage] = []
for msg in request.get("messages") or []:
if not isinstance(msg, dict):
dropped["claude_message_non_dict"] = dropped.get("claude_message_non_dict", 0) + 1
@@ -156,15 +155,15 @@ class ClaudeNormalizer(FormatNormalizer):
return internal
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
system_text = internal.system or self._join_instructions(internal.instructions)
# Claude Messages API: messages[] 仅允许 user/assistant且需要交替这里做最小修复
fixed_messages = self._coerce_claude_message_sequence(internal.messages)
out_messages: List[Dict[str, Any]] = [self._internal_message_to_claude(m) for m in fixed_messages]
out_messages: list[dict[str, Any]] = [self._internal_message_to_claude(m) for m in fixed_messages]
result: Dict[str, Any] = {
result: dict[str, Any] = {
"model": internal.model,
"messages": out_messages,
"max_tokens": internal.max_tokens if internal.max_tokens is not None else 4096,
@@ -210,20 +209,20 @@ class ClaudeNormalizer(FormatNormalizer):
# Responses
# =========================
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse:
def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
rid = str(response.get("id") or "")
model = str(response.get("model") or "")
blocks, dropped = self._claude_content_to_blocks(response.get("content"))
raw_stop = response.get("stop_reason")
stop_reason: Optional[StopReason] = None
stop_reason: StopReason | None = None
if raw_stop is not None:
stop_reason = self._CLAUDE_STOP_TO_INTERNAL.get(str(raw_stop), StopReason.UNKNOWN)
usage_info = self._claude_usage_to_internal(response.get("usage"))
extra: Dict[str, Any] = {}
extra: dict[str, Any] = {}
if raw_stop is not None:
extra.setdefault("raw", {})["stop_reason"] = raw_stop
@@ -245,13 +244,13 @@ class ClaudeNormalizer(FormatNormalizer):
self,
internal: InternalResponse,
*,
requested_model: Optional[str] = None,
) -> Dict[str, Any]:
requested_model: str | None = None,
) -> dict[str, Any]:
cid = internal.id or "unknown"
if not cid.startswith("msg_"):
cid = f"msg_{cid}"
content: List[Dict[str, Any]] = []
content: list[dict[str, Any]] = []
for b in internal.content:
if isinstance(b, TextBlock):
if b.text:
@@ -288,7 +287,7 @@ class ClaudeNormalizer(FormatNormalizer):
if internal.stop_reason is not None:
stop_reason = STOP_REASON_MAPPINGS.get("CLAUDE", {}).get(internal.stop_reason.value, "end_turn")
usage: Dict[str, Any] = {"input_tokens": 0, "output_tokens": 0}
usage: dict[str, Any] = {"input_tokens": 0, "output_tokens": 0}
if internal.usage:
usage = {
"input_tokens": int(internal.usage.input_tokens),
@@ -319,11 +318,11 @@ class ClaudeNormalizer(FormatNormalizer):
def stream_chunk_to_internal(
self,
chunk: Dict[str, Any],
chunk: dict[str, Any],
state: StreamState,
) -> List[InternalStreamEvent]:
) -> list[InternalStreamEvent]:
ss = state.substate(self.FORMAT_ID)
events: List[InternalStreamEvent] = []
events: list[InternalStreamEvent] = []
event_type = chunk.get("type")
if event_type is None:
@@ -335,7 +334,7 @@ class ClaudeNormalizer(FormatNormalizer):
if event_type == "message_start":
message_raw = chunk.get("message")
message: Dict[str, Any] = message_raw if isinstance(message_raw, dict) else {}
message: dict[str, Any] = message_raw if isinstance(message_raw, dict) else {}
msg_id = str(message.get("id") or "")
# 保留初始化时设置的 model客户端请求的模型仅在空时用上游值
model = state.model or str(message.get("model") or "")
@@ -350,7 +349,7 @@ class ClaudeNormalizer(FormatNormalizer):
if event_type == "content_block_start":
index = int(chunk.get("index") or 0)
block_raw = chunk.get("content_block")
block: Dict[str, Any] = block_raw if isinstance(block_raw, dict) else {}
block: dict[str, Any] = block_raw if isinstance(block_raw, dict) else {}
btype = str(block.get("type") or "unknown")
if btype == "text":
@@ -385,7 +384,7 @@ class ClaudeNormalizer(FormatNormalizer):
if event_type == "content_block_delta":
index = int(chunk.get("index") or 0)
delta_raw = chunk.get("delta")
delta: Dict[str, Any] = delta_raw if isinstance(delta_raw, dict) else {}
delta: dict[str, Any] = delta_raw if isinstance(delta_raw, dict) else {}
dtype = str(delta.get("type") or "unknown")
if dtype == "text_delta":
@@ -415,7 +414,7 @@ class ClaudeNormalizer(FormatNormalizer):
if event_type == "message_delta":
delta_raw2 = chunk.get("delta")
delta2: Dict[str, Any] = delta_raw2 if isinstance(delta_raw2, dict) else {}
delta2: dict[str, Any] = delta_raw2 if isinstance(delta_raw2, dict) else {}
raw_stop = delta2.get("stop_reason")
if raw_stop is not None:
ss["stop_reason"] = str(raw_stop)
@@ -426,7 +425,7 @@ class ClaudeNormalizer(FormatNormalizer):
if event_type == "message_stop":
raw_stop = ss.get("stop_reason")
stop_reason: Optional[StopReason] = None
stop_reason: StopReason | None = None
if raw_stop is not None:
stop_reason = self._CLAUDE_STOP_TO_INTERNAL.get(str(raw_stop), StopReason.UNKNOWN)
usage_info = self._claude_usage_to_internal(ss.get("usage"))
@@ -444,9 +443,9 @@ class ClaudeNormalizer(FormatNormalizer):
self,
event: InternalStreamEvent,
state: StreamState,
) -> List[Dict[str, Any]]:
) -> list[dict[str, Any]]:
ss = state.substate(self.FORMAT_ID)
out: List[Dict[str, Any]] = []
out: list[dict[str, Any]] = []
if isinstance(event, MessageStartEvent):
state.message_id = event.message_id or state.message_id
@@ -454,7 +453,7 @@ class ClaudeNormalizer(FormatNormalizer):
if not state.model:
state.model = event.model or ""
ss.setdefault("block_index_to_tool_id", {})
message_obj: Dict[str, Any] = {
message_obj: dict[str, Any] = {
"id": state.message_id or "msg_stream",
"type": "message",
"role": "assistant",
@@ -531,7 +530,7 @@ class ClaudeNormalizer(FormatNormalizer):
if event.stop_reason is not None:
stop_reason = STOP_REASON_MAPPINGS.get("CLAUDE", {}).get(event.stop_reason.value, "end_turn")
msg_delta: Dict[str, Any] = {
msg_delta: dict[str, Any] = {
"type": "message_delta",
"delta": {"stop_reason": stop_reason},
}
@@ -553,15 +552,15 @@ class ClaudeNormalizer(FormatNormalizer):
# Error conversion
# =========================
def is_error_response(self, response: Dict[str, Any]) -> bool:
def is_error_response(self, response: dict[str, Any]) -> bool:
if not isinstance(response, dict):
return False
if response.get("type") == "error":
return True
return "error" in response
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError:
err: Dict[str, Any] = {}
def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
err: dict[str, Any] = {}
if isinstance(error_response, dict):
err_raw = error_response.get("error")
err = err_raw if isinstance(err_raw, dict) else {}
@@ -580,9 +579,9 @@ class ClaudeNormalizer(FormatNormalizer):
extra={"claude": {"error": err}, "raw": {"type": raw_type}},
)
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]:
def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
type_str = self._ERROR_TYPE_TO_CLAUDE.get(internal.type, "api_error")
payload: Dict[str, Any] = {"type": type_str, "message": internal.message}
payload: dict[str, Any] = {"type": type_str, "message": internal.message}
if internal.param is not None:
payload["param"] = internal.param
if internal.code is not None:
@@ -593,8 +592,8 @@ class ClaudeNormalizer(FormatNormalizer):
# Helpers
# =========================
def _claude_message_to_internal(self, msg: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]:
dropped: Dict[str, int] = {}
def _claude_message_to_internal(self, msg: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
dropped: dict[str, int] = {}
role_raw = str(msg.get("role") or "unknown")
if role_raw == "user":
@@ -616,8 +615,8 @@ class ClaudeNormalizer(FormatNormalizer):
dropped,
)
def _claude_content_to_blocks(self, content: Any) -> Tuple[List[ContentBlock], Dict[str, int]]:
dropped: Dict[str, int] = {}
def _claude_content_to_blocks(self, content: Any) -> tuple[list[ContentBlock], dict[str, int]]:
dropped: dict[str, int] = {}
if content is None:
return [], dropped
if isinstance(content, str):
@@ -626,7 +625,7 @@ class ClaudeNormalizer(FormatNormalizer):
dropped["claude_content_non_list"] = dropped.get("claude_content_non_list", 0) + 1
return [], dropped
blocks: List[ContentBlock] = []
blocks: list[ContentBlock] = []
for block in content:
if not isinstance(block, dict):
dropped["claude_block_non_dict"] = dropped.get("claude_block_non_dict", 0) + 1
@@ -641,7 +640,7 @@ class ClaudeNormalizer(FormatNormalizer):
if btype == "image":
src_raw = block.get("source")
src: Dict[str, Any] = src_raw if isinstance(src_raw, dict) else {}
src: dict[str, Any] = src_raw if isinstance(src_raw, dict) else {}
stype = src.get("type")
if stype == "base64":
data = src.get("data")
@@ -687,7 +686,7 @@ class ClaudeNormalizer(FormatNormalizer):
tool_use_id: str,
raw_content: Any,
is_error: bool,
raw_block: Dict[str, Any],
raw_block: dict[str, Any],
) -> ToolResultBlock:
if raw_content is None:
return ToolResultBlock(
@@ -723,7 +722,7 @@ class ClaudeNormalizer(FormatNormalizer):
)
if isinstance(raw_content, list):
text_parts: List[str] = []
text_parts: list[str] = []
for part in raw_content:
if isinstance(part, dict) and part.get("type") == "text":
text = part.get("text")
@@ -747,15 +746,15 @@ class ClaudeNormalizer(FormatNormalizer):
extra={"claude": raw_block},
)
def _collapse_claude_system(self, system_value: Any) -> Tuple[Optional[str], Dict[str, int]]:
dropped: Dict[str, int] = {}
def _collapse_claude_system(self, system_value: Any) -> tuple[str | None, dict[str, int]]:
dropped: dict[str, int] = {}
if system_value is None:
return None, dropped
if isinstance(system_value, str):
return (system_value or None), dropped
if isinstance(system_value, list):
texts: List[str] = []
texts: list[str] = []
for item in system_value:
if not isinstance(item, dict):
dropped["claude_system_item_non_dict"] = dropped.get("claude_system_item_non_dict", 0) + 1
@@ -773,16 +772,16 @@ class ClaudeNormalizer(FormatNormalizer):
dropped["claude_system_unsupported"] = dropped.get("claude_system_unsupported", 0) + 1
return None, dropped
def _join_instructions(self, instructions: List[InstructionSegment]) -> Optional[str]:
def _join_instructions(self, instructions: list[InstructionSegment]) -> str | None:
parts = [seg.text for seg in instructions if seg.text]
joined = "\n\n".join(parts)
return joined or None
def _claude_tools_to_internal(self, tools: Any) -> Optional[List[ToolDefinition]]:
def _claude_tools_to_internal(self, tools: Any) -> list[ToolDefinition] | None:
if not tools or not isinstance(tools, list):
return None
out: List[ToolDefinition] = []
out: list[ToolDefinition] = []
for tool in tools:
if not isinstance(tool, dict):
continue
@@ -799,7 +798,7 @@ class ClaudeNormalizer(FormatNormalizer):
)
return out or None
def _claude_tool_choice_to_internal(self, tool_choice: Any) -> Optional[ToolChoice]:
def _claude_tool_choice_to_internal(self, tool_choice: Any) -> ToolChoice | None:
if tool_choice is None:
return None
if not isinstance(tool_choice, dict):
@@ -818,7 +817,7 @@ class ClaudeNormalizer(FormatNormalizer):
return ToolChoice(type=ToolChoiceType.AUTO, extra={"claude": tool_choice})
def _tool_choice_to_claude(self, tool_choice: ToolChoice) -> Dict[str, Any]:
def _tool_choice_to_claude(self, tool_choice: ToolChoice) -> dict[str, Any]:
if tool_choice.type == ToolChoiceType.NONE:
return {"type": "none"}
if tool_choice.type == ToolChoiceType.AUTO:
@@ -829,11 +828,11 @@ class ClaudeNormalizer(FormatNormalizer):
return {"type": "tool_use", "name": tool_choice.tool_name or ""}
return {"type": "auto"}
def _internal_message_to_claude(self, msg: InternalMessage) -> Dict[str, Any]:
def _internal_message_to_claude(self, msg: InternalMessage) -> dict[str, Any]:
role = "user" if msg.role == Role.USER else "assistant"
blocks: List[Dict[str, Any]] = []
text_parts: List[str] = []
blocks: list[dict[str, Any]] = []
text_parts: list[str] = []
for b in msg.content:
if isinstance(b, UnknownBlock):
@@ -902,8 +901,8 @@ class ClaudeNormalizer(FormatNormalizer):
return {"role": role, "content": blocks}
def _coerce_claude_message_sequence(self, messages: List[InternalMessage]) -> List[InternalMessage]:
normalized: List[InternalMessage] = []
def _coerce_claude_message_sequence(self, messages: list[InternalMessage]) -> list[InternalMessage]:
normalized: list[InternalMessage] = []
for m in messages:
role = m.role
if role not in (Role.USER, Role.ASSISTANT):
@@ -916,7 +915,7 @@ class ClaudeNormalizer(FormatNormalizer):
if normalized[0].role != Role.USER:
normalized = [InternalMessage(role=Role.USER, content=[])] + normalized
merged: List[InternalMessage] = []
merged: list[InternalMessage] = []
for m in normalized:
if merged and merged[-1].role == m.role:
merged[-1].content.extend(m.content)
@@ -925,12 +924,12 @@ class ClaudeNormalizer(FormatNormalizer):
return merged
def _claude_usage_to_internal(self, usage: Any) -> Optional[UsageInfo]:
def _claude_usage_to_internal(self, usage: Any) -> UsageInfo | None:
if not isinstance(usage, dict):
return None
mapping = USAGE_FIELD_MAPPINGS.get("CLAUDE", {})
fields: Dict[str, int] = {}
fields: dict[str, int] = {}
extra = self._extract_extra(usage, set(mapping.keys()))
for provider_key, internal_key in mapping.items():
@@ -952,8 +951,8 @@ class ClaudeNormalizer(FormatNormalizer):
extra={"claude": extra} if extra else {},
)
def _usage_to_claude(self, usage: UsageInfo) -> Dict[str, Any]:
result: Dict[str, Any] = {
def _usage_to_claude(self, usage: UsageInfo) -> dict[str, Any]:
result: dict[str, Any] = {
"input_tokens": int(usage.input_tokens),
"output_tokens": int(usage.output_tokens),
}
@@ -969,7 +968,7 @@ class ClaudeNormalizer(FormatNormalizer):
except ValueError:
return ErrorType.UNKNOWN
def _optional_int(self, value: Any) -> Optional[int]:
def _optional_int(self, value: Any) -> int | None:
if value is None:
return None
try:
@@ -977,7 +976,7 @@ class ClaudeNormalizer(FormatNormalizer):
except (TypeError, ValueError):
return None
def _optional_float(self, value: Any) -> Optional[float]:
def _optional_float(self, value: Any) -> float | None:
if value is None:
return None
try:
@@ -985,7 +984,7 @@ class ClaudeNormalizer(FormatNormalizer):
except (TypeError, ValueError):
return None
def _coerce_str_list(self, value: Any) -> Optional[List[str]]:
def _coerce_str_list(self, value: Any) -> list[str] | None:
if value is None:
return None
if isinstance(value, str):
@@ -994,10 +993,10 @@ class ClaudeNormalizer(FormatNormalizer):
return [str(x) for x in value if x is not None]
return None
def _extract_extra(self, payload: Dict[str, Any], known_keys: set[str]) -> Dict[str, Any]:
def _extract_extra(self, payload: dict[str, Any], known_keys: set[str]) -> dict[str, Any]:
return {k: v for k, v in payload.items() if k not in known_keys}
def _merge_dropped(self, target: Dict[str, int], source: Dict[str, int]) -> None:
def _merge_dropped(self, target: dict[str, int], source: dict[str, int]) -> None:
for k, v in source.items():
target[k] = target.get(k, 0) + int(v)

View File

@@ -7,7 +7,6 @@ CLAUDE_CLI 的请求/响应 body 与 CLAUDE 一致Anthropic Messages API
如需 CLI 特殊处理,可覆盖 request_from_internal / request_to_internal 等方法。
"""
from __future__ import annotations
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer

View File

@@ -11,11 +11,9 @@ Gemini (GenerateContent / streamGenerateContent) Normalizer
- 响应/流式通常为 camelCasecandidates/finishReason/usageMetadata/modelVersion
"""
from __future__ import annotations
import json
import time
from typing import Any, Dict, List, Optional, Tuple
from typing import Any
from src.core.api_format.conversion.field_mappings import (
ERROR_TYPE_MAPPINGS,
@@ -68,7 +66,7 @@ class GeminiNormalizer(FormatNormalizer):
supports_images=True,
)
_FINISH_REASON_TO_STOP: Dict[str, StopReason] = {
_FINISH_REASON_TO_STOP: dict[str, StopReason] = {
"STOP": StopReason.END_TURN,
"MAX_TOKENS": StopReason.MAX_TOKENS,
"SAFETY": StopReason.CONTENT_FILTERED,
@@ -77,7 +75,7 @@ class GeminiNormalizer(FormatNormalizer):
"OTHER": StopReason.UNKNOWN,
}
_ERROR_TYPE_TO_GEMINI_STATUS: Dict[ErrorType, str] = {
_ERROR_TYPE_TO_GEMINI_STATUS: dict[ErrorType, str] = {
ErrorType.INVALID_REQUEST: "INVALID_ARGUMENT",
ErrorType.AUTHENTICATION: "UNAUTHENTICATED",
ErrorType.PERMISSION_DENIED: "PERMISSION_DENIED",
@@ -94,11 +92,11 @@ class GeminiNormalizer(FormatNormalizer):
# Requests
# =========================
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
model = str(request.get("model") or "")
dropped: Dict[str, int] = {}
dropped: dict[str, int] = {}
instructions: List[InstructionSegment] = []
instructions: list[InstructionSegment] = []
system_text, sys_dropped = self._collapse_system_instruction(
request.get("system_instruction")
if "system_instruction" in request
@@ -108,7 +106,7 @@ class GeminiNormalizer(FormatNormalizer):
if system_text:
instructions.append(InstructionSegment(role=Role.SYSTEM, text=system_text))
messages: List[InternalMessage] = []
messages: list[InternalMessage] = []
contents = request.get("contents") or []
if isinstance(contents, list):
for content in contents:
@@ -152,7 +150,7 @@ class GeminiNormalizer(FormatNormalizer):
)
# 构建 extra保留原始 gemini 字段
extra: Dict[str, Any] = {"gemini": self._extract_extra(request, {"contents"})}
extra: dict[str, Any] = {"gemini": self._extract_extra(request, {"contents"})}
# 保留 generationConfig 中的特殊字段responseModalities, thinkingConfig 等)
# 这些字段在 _get_generation_config 中已提取,需要单独存储以便转换时使用
@@ -160,7 +158,7 @@ class GeminiNormalizer(FormatNormalizer):
response_modalities = generation_config.get("response_modalities")
thinking_config = generation_config.get("thinking_config")
if response_modalities or thinking_config:
google_extra: Dict[str, Any] = {}
google_extra: dict[str, Any] = {}
if response_modalities:
google_extra["response_modalities"] = response_modalities
if thinking_config:
@@ -188,7 +186,7 @@ class GeminiNormalizer(FormatNormalizer):
return internal
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
system_text = internal.system or self._join_instructions(internal.instructions)
# tools/tool_choice
@@ -212,7 +210,7 @@ class GeminiNormalizer(FormatNormalizer):
if internal.tool_choice:
tool_config = self._tool_choice_to_gemini_tool_config(internal.tool_choice)
generation_config: Dict[str, Any] = {}
generation_config: dict[str, Any] = {}
if internal.max_tokens is not None:
generation_config["max_output_tokens"] = internal.max_tokens
if internal.temperature is not None:
@@ -231,7 +229,7 @@ class GeminiNormalizer(FormatNormalizer):
thinking_config = google_extra.get("thinking_config")
if isinstance(thinking_config, dict):
# snake_case -> camelCase 转换
gemini_thinking: Dict[str, Any] = {}
gemini_thinking: dict[str, Any] = {}
if "thinking_budget" in thinking_config:
gemini_thinking["thinkingBudget"] = thinking_config["thinking_budget"]
if "include_thoughts" in thinking_config:
@@ -265,11 +263,11 @@ class GeminiNormalizer(FormatNormalizer):
if "thinking_config" in orig_gc and "thinkingConfig" not in generation_config:
generation_config["thinkingConfig"] = orig_gc["thinking_config"]
contents: List[Dict[str, Any]] = []
contents: list[dict[str, Any]] = []
for msg in internal.messages:
contents.append(self._internal_message_to_content(msg))
result: Dict[str, Any] = {
result: dict[str, Any] = {
"contents": contents,
}
@@ -295,7 +293,7 @@ class GeminiNormalizer(FormatNormalizer):
# Responses
# =========================
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse:
def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
rid = str(response.get("id") or "")
model = str(response.get("modelVersion") or response.get("model") or "")
@@ -315,7 +313,7 @@ class GeminiNormalizer(FormatNormalizer):
usage_info = self._usage_metadata_to_internal(response.get("usageMetadata"))
extra: Dict[str, Any] = {}
extra: dict[str, Any] = {}
if finish_reason is not None:
extra.setdefault("raw", {})["finishReason"] = finish_reason
@@ -337,9 +335,9 @@ class GeminiNormalizer(FormatNormalizer):
self,
internal: InternalResponse,
*,
requested_model: Optional[str] = None,
) -> Dict[str, Any]:
parts: List[Dict[str, Any]] = []
requested_model: str | None = None,
) -> dict[str, Any]:
parts: list[dict[str, Any]] = []
for b in internal.content:
if isinstance(b, TextBlock):
if b.text:
@@ -374,7 +372,7 @@ class GeminiNormalizer(FormatNormalizer):
if internal.stop_reason is not None:
finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(internal.stop_reason.value, "OTHER")
usage_metadata: Dict[str, Any] = {}
usage_metadata: dict[str, Any] = {}
if internal.usage:
usage_metadata = {
"promptTokenCount": int(internal.usage.input_tokens),
@@ -384,7 +382,7 @@ class GeminiNormalizer(FormatNormalizer):
if internal.usage.cache_read_tokens:
usage_metadata["cachedContentTokenCount"] = int(internal.usage.cache_read_tokens)
candidate: Dict[str, Any] = {
candidate: dict[str, Any] = {
"content": {"parts": parts, "role": "model"},
"index": 0,
}
@@ -394,7 +392,7 @@ class GeminiNormalizer(FormatNormalizer):
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
model_name = requested_model if requested_model else (internal.model or "gemini")
out: Dict[str, Any] = {
out: dict[str, Any] = {
"candidates": [candidate],
"modelVersion": model_name,
}
@@ -412,9 +410,9 @@ class GeminiNormalizer(FormatNormalizer):
# Streaming
# =========================
def stream_chunk_to_internal(self, chunk: Dict[str, Any], state: StreamState) -> List[InternalStreamEvent]:
def stream_chunk_to_internal(self, chunk: dict[str, Any], state: StreamState) -> list[InternalStreamEvent]:
ss = state.substate(self.FORMAT_ID)
events: List[InternalStreamEvent] = []
events: list[InternalStreamEvent] = []
if not ss.get("message_started"):
# 保留初始化时设置的 model客户端请求的模型仅在空时用上游值
@@ -544,11 +542,11 @@ class GeminiNormalizer(FormatNormalizer):
self,
event: InternalStreamEvent,
state: StreamState,
) -> List[Dict[str, Any]]:
) -> list[dict[str, Any]]:
ss = state.substate(self.FORMAT_ID)
out: List[Dict[str, Any]] = []
out: list[dict[str, Any]] = []
def base_chunk(parts: List[Dict[str, Any]]) -> Dict[str, Any]:
def base_chunk(parts: list[dict[str, Any]]) -> dict[str, Any]:
return {
"candidates": [
{
@@ -622,7 +620,7 @@ class GeminiNormalizer(FormatNormalizer):
name = str(entry.get("name") or "")
raw_json = str(entry.get("json") or "")
args: Dict[str, Any] = {}
args: dict[str, Any] = {}
if raw_json:
try:
parsed = json.loads(raw_json)
@@ -639,7 +637,7 @@ class GeminiNormalizer(FormatNormalizer):
if event.stop_reason is not None:
finish_reason = STOP_REASON_MAPPINGS.get("GEMINI", {}).get(event.stop_reason.value, "OTHER")
chunk: Dict[str, Any] = base_chunk([])
chunk: dict[str, Any] = base_chunk([])
if finish_reason is not None:
chunk["candidates"][0]["finishReason"] = finish_reason
@@ -665,10 +663,10 @@ class GeminiNormalizer(FormatNormalizer):
# Error conversion
# =========================
def is_error_response(self, response: Dict[str, Any]) -> bool:
def is_error_response(self, response: dict[str, Any]) -> bool:
return isinstance(response, dict) and "error" in response
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError:
def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
err = error_response.get("error") if isinstance(error_response, dict) else None
err = err if isinstance(err, dict) else {}
@@ -691,9 +689,9 @@ class GeminiNormalizer(FormatNormalizer):
extra={"gemini": {"error": err}, "raw": {"status": raw_status}},
)
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]:
def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
status = self._ERROR_TYPE_TO_GEMINI_STATUS.get(internal.type, "INTERNAL")
payload: Dict[str, Any] = {
payload: dict[str, Any] = {
"code": 400 if internal.type == ErrorType.INVALID_REQUEST else 500,
"message": internal.message,
"status": status,
@@ -704,8 +702,8 @@ class GeminiNormalizer(FormatNormalizer):
# Helpers
# =========================
def _content_to_internal_message(self, content: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]:
dropped: Dict[str, int] = {}
def _content_to_internal_message(self, content: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
dropped: dict[str, int] = {}
role_raw = str(content.get("role") or "user")
if role_raw == "model":
@@ -727,15 +725,15 @@ class GeminiNormalizer(FormatNormalizer):
dropped,
)
def _parts_to_blocks(self, parts: Any) -> Tuple[List[ContentBlock], Dict[str, int]]:
dropped: Dict[str, int] = {}
def _parts_to_blocks(self, parts: Any) -> tuple[list[ContentBlock], dict[str, int]]:
dropped: dict[str, int] = {}
if parts is None:
return [], dropped
if not isinstance(parts, list):
dropped["gemini_parts_non_list"] = dropped.get("gemini_parts_non_list", 0) + 1
return [], dropped
blocks: List[ContentBlock] = []
blocks: list[ContentBlock] = []
for part in parts:
if not isinstance(part, dict):
dropped["gemini_part_non_dict"] = dropped.get("gemini_part_non_dict", 0) + 1
@@ -785,7 +783,7 @@ class GeminiNormalizer(FormatNormalizer):
name = str(func_resp.get("name") or "")
response = func_resp.get("response")
output: Any = None
content_text: Optional[str] = None
content_text: str | None = None
# 兼容历史response 常见结构为 {"result": ...}
if isinstance(response, dict) and "result" in response:
@@ -815,10 +813,10 @@ class GeminiNormalizer(FormatNormalizer):
return blocks, dropped
def _internal_message_to_content(self, msg: InternalMessage) -> Dict[str, Any]:
def _internal_message_to_content(self, msg: InternalMessage) -> dict[str, Any]:
role = "model" if msg.role == Role.ASSISTANT else "user"
parts: List[Dict[str, Any]] = []
parts: list[dict[str, Any]] = []
for b in msg.content:
if isinstance(b, UnknownBlock):
continue
@@ -861,8 +859,8 @@ class GeminiNormalizer(FormatNormalizer):
return {"role": role, "parts": parts}
def _collapse_system_instruction(self, system_instruction: Any) -> Tuple[Optional[str], Dict[str, int]]:
dropped: Dict[str, int] = {}
def _collapse_system_instruction(self, system_instruction: Any) -> tuple[str | None, dict[str, int]]:
dropped: dict[str, int] = {}
if system_instruction is None:
return None, dropped
@@ -870,7 +868,7 @@ class GeminiNormalizer(FormatNormalizer):
if isinstance(system_instruction, dict):
parts = system_instruction.get("parts")
if isinstance(parts, list):
texts: List[str] = []
texts: list[str] = []
for part in parts:
if isinstance(part, dict) and "text" in part and part.get("text"):
texts.append(str(part.get("text")))
@@ -880,7 +878,7 @@ class GeminiNormalizer(FormatNormalizer):
dropped["gemini_system_instruction_unsupported"] = dropped.get("gemini_system_instruction_unsupported", 0) + 1
return None, dropped
def _get_generation_config(self, request: Dict[str, Any]) -> Dict[str, Any]:
def _get_generation_config(self, request: dict[str, Any]) -> dict[str, Any]:
# 兼容 snake_case 与 camelCase
gc = request.get("generation_config") if "generation_config" in request else request.get("generationConfig")
if not isinstance(gc, dict):
@@ -893,7 +891,7 @@ class GeminiNormalizer(FormatNormalizer):
return gc.get(k)
return None
normalized: Dict[str, Any] = {}
normalized: dict[str, Any] = {}
normalized["max_output_tokens"] = pick("max_output_tokens", "maxOutputTokens")
normalized["temperature"] = pick("temperature")
normalized["top_p"] = pick("top_p", "topP")
@@ -912,11 +910,11 @@ class GeminiNormalizer(FormatNormalizer):
return {k: v for k, v in normalized.items() if v is not None}
def _gemini_tools_to_internal(self, tools: Any) -> Optional[List[ToolDefinition]]:
def _gemini_tools_to_internal(self, tools: Any) -> list[ToolDefinition] | None:
if not tools or not isinstance(tools, list):
return None
out: List[ToolDefinition] = []
out: list[ToolDefinition] = []
for tool in tools:
if not isinstance(tool, dict):
continue
@@ -945,7 +943,7 @@ class GeminiNormalizer(FormatNormalizer):
return out or None
def _gemini_tool_config_to_tool_choice(self, tool_config: Any) -> Optional[ToolChoice]:
def _gemini_tool_config_to_tool_choice(self, tool_config: Any) -> ToolChoice | None:
if tool_config is None:
return None
if not isinstance(tool_config, dict):
@@ -972,9 +970,9 @@ class GeminiNormalizer(FormatNormalizer):
return ToolChoice(type=ToolChoiceType.AUTO, extra={"gemini": tool_config})
def _tool_choice_to_gemini_tool_config(self, tool_choice: ToolChoice) -> Dict[str, Any]:
def _tool_choice_to_gemini_tool_config(self, tool_choice: ToolChoice) -> dict[str, Any]:
mode = "AUTO"
cfg: Dict[str, Any] = {}
cfg: dict[str, Any] = {}
if tool_choice.type == ToolChoiceType.NONE:
mode = "NONE"
@@ -987,12 +985,12 @@ class GeminiNormalizer(FormatNormalizer):
cfg["mode"] = mode
return {"function_calling_config": cfg}
def _usage_metadata_to_internal(self, usage_metadata: Any) -> Optional[UsageInfo]:
def _usage_metadata_to_internal(self, usage_metadata: Any) -> UsageInfo | None:
if not isinstance(usage_metadata, dict):
return None
mapping = USAGE_FIELD_MAPPINGS.get("GEMINI", {})
fields: Dict[str, int] = {}
fields: dict[str, int] = {}
extra = self._extract_extra(usage_metadata, set(mapping.keys()))
# promptTokenCount/candidatesTokenCount/totalTokenCount/cachedContentTokenCount
@@ -1023,7 +1021,7 @@ class GeminiNormalizer(FormatNormalizer):
extra={"gemini": extra} if extra else {},
)
def _join_instructions(self, instructions: List[InstructionSegment]) -> Optional[str]:
def _join_instructions(self, instructions: list[InstructionSegment]) -> str | None:
parts = [seg.text for seg in instructions if seg.text]
joined = "\n\n".join(parts)
return joined or None
@@ -1034,7 +1032,7 @@ class GeminiNormalizer(FormatNormalizer):
except ValueError:
return ErrorType.UNKNOWN
def _optional_int(self, value: Any) -> Optional[int]:
def _optional_int(self, value: Any) -> int | None:
if value is None:
return None
try:
@@ -1042,7 +1040,7 @@ class GeminiNormalizer(FormatNormalizer):
except (TypeError, ValueError):
return None
def _optional_float(self, value: Any) -> Optional[float]:
def _optional_float(self, value: Any) -> float | None:
if value is None:
return None
try:
@@ -1050,7 +1048,7 @@ class GeminiNormalizer(FormatNormalizer):
except (TypeError, ValueError):
return None
def _coerce_str_list(self, value: Any) -> Optional[List[str]]:
def _coerce_str_list(self, value: Any) -> list[str] | None:
if value is None:
return None
if isinstance(value, str):
@@ -1059,10 +1057,10 @@ class GeminiNormalizer(FormatNormalizer):
return [str(x) for x in value if x is not None]
return None
def _extract_extra(self, payload: Dict[str, Any], known_keys: set[str]) -> Dict[str, Any]:
def _extract_extra(self, payload: dict[str, Any], known_keys: set[str]) -> dict[str, Any]:
return {k: v for k, v in payload.items() if k not in known_keys}
def _merge_dropped(self, target: Dict[str, int], source: Dict[str, int]) -> None:
def _merge_dropped(self, target: dict[str, int], source: dict[str, int]) -> None:
for k, v in source.items():
target[k] = target.get(k, 0) + int(v)

View File

@@ -7,7 +7,6 @@ GEMINI_CLI 的请求/响应 body 与 GEMINI 一致Google Gemini API
如需 CLI 特殊处理,可覆盖 request_from_internal / request_to_internal 等方法。
"""
from __future__ import annotations
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer

View File

@@ -7,11 +7,10 @@ OpenAI Chat Completions Normalizer
- 可选OpenAI error <-> InternalError
"""
from __future__ import annotations
import json
import time
from typing import Any, Dict, List, Optional, Tuple, Union
from typing import Any
from src.core.logger import logger
from src.core.api_format.conversion.field_mappings import (
@@ -49,7 +48,6 @@ from src.core.api_format.conversion.stream_events import (
InternalStreamEvent,
MessageStartEvent,
MessageStopEvent,
StreamEventType,
ToolCallDeltaEvent,
)
from src.core.api_format.conversion.stream_state import StreamState
@@ -65,7 +63,7 @@ class OpenAINormalizer(FormatNormalizer):
)
# finish_reason -> StopReason
_FINISH_REASON_TO_STOP: Dict[str, StopReason] = {
_FINISH_REASON_TO_STOP: dict[str, StopReason] = {
"stop": StopReason.END_TURN,
"length": StopReason.MAX_TOKENS,
"tool_calls": StopReason.TOOL_USE,
@@ -74,7 +72,7 @@ class OpenAINormalizer(FormatNormalizer):
}
# StopReason -> finish_reason
_STOP_TO_FINISH_REASON: Dict[StopReason, str] = {
_STOP_TO_FINISH_REASON: dict[StopReason, str] = {
StopReason.END_TURN: "stop",
StopReason.MAX_TOKENS: "length",
StopReason.STOP_SEQUENCE: "stop",
@@ -84,7 +82,7 @@ class OpenAINormalizer(FormatNormalizer):
}
# InternalError.type -> OpenAI error.type最佳努力
_ERROR_TYPE_TO_OPENAI: Dict[ErrorType, str] = {
_ERROR_TYPE_TO_OPENAI: dict[ErrorType, str] = {
ErrorType.INVALID_REQUEST: "invalid_request_error",
ErrorType.AUTHENTICATION: "invalid_api_key",
ErrorType.PERMISSION_DENIED: "invalid_request_error",
@@ -101,13 +99,13 @@ class OpenAINormalizer(FormatNormalizer):
# Requests
# =========================
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
model = str(request.get("model") or "")
dropped: Dict[str, int] = {}
dropped: dict[str, int] = {}
instructions: List[InstructionSegment] = []
messages: List[InternalMessage] = []
instructions: list[InstructionSegment] = []
messages: list[InternalMessage] = []
for msg in request.get("messages") or []:
if not isinstance(msg, dict):
@@ -146,7 +144,7 @@ class OpenAINormalizer(FormatNormalizer):
)
# 构建 extra保留未识别字段
extra: Dict[str, Any] = {"openai": self._extract_extra(request, {"messages"})}
extra: dict[str, Any] = {"openai": self._extract_extra(request, {"messages"})}
# 处理 extra_body.google (用于 Gemini 特定功能透传,如 thinkingConfig, responseModalities)
extra_body = request.get("extra_body")
@@ -175,8 +173,8 @@ class OpenAINormalizer(FormatNormalizer):
return internal
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
out_messages: List[Dict[str, Any]] = []
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
out_messages: list[dict[str, Any]] = []
if internal.instructions:
for seg in internal.instructions:
@@ -189,7 +187,7 @@ class OpenAINormalizer(FormatNormalizer):
for msg in internal.messages:
out_messages.extend(self._internal_message_to_openai_messages(msg))
result: Dict[str, Any] = {
result: dict[str, Any] = {
"model": internal.model,
"messages": out_messages,
}
@@ -233,11 +231,11 @@ class OpenAINormalizer(FormatNormalizer):
# Responses
# =========================
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse:
def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
rid = str(response.get("id") or "")
model = str(response.get("model") or "")
extra: Dict[str, Any] = {}
extra: dict[str, Any] = {}
choices = response.get("choices") or []
if isinstance(choices, list) and len(choices) > 1:
@@ -285,12 +283,12 @@ class OpenAINormalizer(FormatNormalizer):
self,
internal: InternalResponse,
*,
requested_model: Optional[str] = None,
) -> Dict[str, Any]:
requested_model: str | None = None,
) -> dict[str, Any]:
# OpenAI Chat Completions response envelope
# 优先使用用户请求的原始模型名,回退到上游返回的模型名
model_name = requested_model if requested_model else internal.model
out: Dict[str, Any] = {
out: dict[str, Any] = {
"id": internal.id or "chatcmpl-unknown",
"object": "chat.completion",
"created": int(time.time()),
@@ -298,7 +296,7 @@ class OpenAINormalizer(FormatNormalizer):
"choices": [],
}
message: Dict[str, Any] = {"role": "assistant"}
message: dict[str, Any] = {"role": "assistant"}
content_blocks, tool_blocks = self._split_blocks(internal.content)
content_value = self._blocks_to_openai_content(content_blocks)
@@ -336,9 +334,9 @@ class OpenAINormalizer(FormatNormalizer):
# Streaming
# =========================
def stream_chunk_to_internal(self, chunk: Dict[str, Any], state: StreamState) -> List[InternalStreamEvent]:
def stream_chunk_to_internal(self, chunk: dict[str, Any], state: StreamState) -> list[InternalStreamEvent]:
ss = state.substate(self.FORMAT_ID)
events: List[InternalStreamEvent] = []
events: list[InternalStreamEvent] = []
# OpenAI streaming error通常是单个 {"error": {...}}
if isinstance(chunk, dict) and "error" in chunk:
@@ -437,11 +435,11 @@ class OpenAINormalizer(FormatNormalizer):
self,
event: InternalStreamEvent,
state: StreamState,
) -> List[Dict[str, Any]]:
) -> list[dict[str, Any]]:
ss = state.substate(self.FORMAT_ID)
out: List[Dict[str, Any]] = []
out: list[dict[str, Any]] = []
def base_chunk(delta: Dict[str, Any], finish_reason: Optional[str] = None) -> Dict[str, Any]:
def base_chunk(delta: dict[str, Any], finish_reason: str | None = None) -> dict[str, Any]:
return {
"id": state.message_id or "chatcmpl-stream",
"object": "chat.completion.chunk",
@@ -570,10 +568,10 @@ class OpenAINormalizer(FormatNormalizer):
# Error conversion
# =========================
def is_error_response(self, response: Dict[str, Any]) -> bool:
def is_error_response(self, response: dict[str, Any]) -> bool:
return isinstance(response, dict) and "error" in response
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError:
def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
err = error_response.get("error") if isinstance(error_response, dict) else None
err = err if isinstance(err, dict) else {}
@@ -592,9 +590,9 @@ class OpenAINormalizer(FormatNormalizer):
extra={"openai": {"error": err}, "raw": {"type": raw_type}},
)
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]:
def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
type_str = self._ERROR_TYPE_TO_OPENAI.get(internal.type, "server_error")
payload: Dict[str, Any] = {
payload: dict[str, Any] = {
"message": internal.message,
"type": type_str,
}
@@ -608,8 +606,8 @@ class OpenAINormalizer(FormatNormalizer):
# Helpers
# =========================
def _openai_message_to_internal(self, msg: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]:
dropped: Dict[str, int] = {}
def _openai_message_to_internal(self, msg: dict[str, Any]) -> tuple[InternalMessage | None, dict[str, int]]:
dropped: dict[str, int] = {}
role_raw = str(msg.get("role") or "unknown")
role = self._role_from_openai(role_raw)
@@ -656,8 +654,8 @@ class OpenAINormalizer(FormatNormalizer):
dropped,
)
def _openai_content_to_blocks(self, content: Any) -> Tuple[List[ContentBlock], Dict[str, int]]:
dropped: Dict[str, int] = {}
def _openai_content_to_blocks(self, content: Any) -> tuple[list[ContentBlock], dict[str, int]]:
dropped: dict[str, int] = {}
if content is None:
return [], dropped
@@ -667,7 +665,7 @@ class OpenAINormalizer(FormatNormalizer):
dropped["openai_content_non_list"] = dropped.get("openai_content_non_list", 0) + 1
return [], dropped
blocks: List[ContentBlock] = []
blocks: list[ContentBlock] = []
for part in content:
if not isinstance(part, dict):
dropped["openai_content_part_non_dict"] = dropped.get("openai_content_part_non_dict", 0) + 1
@@ -698,21 +696,21 @@ class OpenAINormalizer(FormatNormalizer):
return blocks, dropped
def _collapse_openai_text(self, content: Any) -> Tuple[str, Dict[str, int]]:
def _collapse_openai_text(self, content: Any) -> tuple[str, dict[str, int]]:
blocks, dropped = self._openai_content_to_blocks(content)
text_parts = [b.text for b in blocks if isinstance(b, TextBlock) and b.text]
return ("\n\n".join(text_parts), dropped)
def _join_instructions(self, instructions: List[InstructionSegment]) -> Optional[str]:
def _join_instructions(self, instructions: list[InstructionSegment]) -> str | None:
parts = [seg.text for seg in instructions if seg.text]
joined = "\n\n".join(parts)
return joined or None
def _openai_tools_to_internal(self, tools: Any) -> Optional[List[ToolDefinition]]:
def _openai_tools_to_internal(self, tools: Any) -> list[ToolDefinition] | None:
if not tools or not isinstance(tools, list):
return None
out: List[ToolDefinition] = []
out: list[ToolDefinition] = []
for tool in tools:
if not isinstance(tool, dict):
continue
@@ -720,7 +718,7 @@ class OpenAINormalizer(FormatNormalizer):
continue
function_raw = tool.get("function")
function: Dict[str, Any] = function_raw if isinstance(function_raw, dict) else {}
function: dict[str, Any] = function_raw if isinstance(function_raw, dict) else {}
name = str(function.get("name") or "")
if not name:
continue
@@ -740,7 +738,7 @@ class OpenAINormalizer(FormatNormalizer):
return out or None
def _openai_tool_choice_to_internal(self, tool_choice: Any) -> Optional[ToolChoice]:
def _openai_tool_choice_to_internal(self, tool_choice: Any) -> ToolChoice | None:
if tool_choice is None:
return None
@@ -756,13 +754,13 @@ class OpenAINormalizer(FormatNormalizer):
if isinstance(tool_choice, dict) and tool_choice.get("type") == "function":
fn_raw = tool_choice.get("function")
fn: Dict[str, Any] = fn_raw if isinstance(fn_raw, dict) else {}
fn: dict[str, Any] = fn_raw if isinstance(fn_raw, dict) else {}
name = str(fn.get("name") or "")
return ToolChoice(type=ToolChoiceType.TOOL, tool_name=name, extra={"openai": tool_choice})
return ToolChoice(type=ToolChoiceType.AUTO, extra={"openai": tool_choice})
def _tool_choice_to_openai(self, tool_choice: ToolChoice) -> Union[str, Dict[str, Any]]:
def _tool_choice_to_openai(self, tool_choice: ToolChoice) -> str | dict[str, Any]:
if tool_choice.type == ToolChoiceType.NONE:
return "none"
if tool_choice.type == ToolChoiceType.AUTO:
@@ -773,8 +771,8 @@ class OpenAINormalizer(FormatNormalizer):
return {"type": "function", "function": {"name": tool_choice.tool_name or ""}}
return "auto"
def _openai_tool_call_to_block(self, tool_call: Any) -> Tuple[Optional[ToolUseBlock], Dict[str, int]]:
dropped: Dict[str, int] = {}
def _openai_tool_call_to_block(self, tool_call: Any) -> tuple[ToolUseBlock | None, dict[str, int]]:
dropped: dict[str, int] = {}
if not isinstance(tool_call, dict):
dropped["openai_tool_call_non_dict"] = dropped.get("openai_tool_call_non_dict", 0) + 1
return None, dropped
@@ -785,12 +783,12 @@ class OpenAINormalizer(FormatNormalizer):
return None, dropped
fn_raw = tool_call.get("function")
fn: Dict[str, Any] = fn_raw if isinstance(fn_raw, dict) else {}
fn: dict[str, Any] = fn_raw if isinstance(fn_raw, dict) else {}
name = str(fn.get("name") or "")
args_str = str(fn.get("arguments") or "")
tool_id = str(tool_call.get("id") or "")
tool_input: Dict[str, Any]
tool_input: dict[str, Any]
if args_str:
try:
parsed = json.loads(args_str)
@@ -810,15 +808,15 @@ class OpenAINormalizer(FormatNormalizer):
dropped,
)
def _legacy_function_call_to_block(self, func_call: Dict[str, Any]) -> Tuple[Optional[ToolUseBlock], Dict[str, int]]:
dropped: Dict[str, int] = {}
def _legacy_function_call_to_block(self, func_call: dict[str, Any]) -> tuple[ToolUseBlock | None, dict[str, int]]:
dropped: dict[str, int] = {}
name = str(func_call.get("name") or "")
args_str = str(func_call.get("arguments") or "")
if not name:
dropped["openai_function_call_missing_name"] = dropped.get("openai_function_call_missing_name", 0) + 1
return None, dropped
tool_input: Dict[str, Any]
tool_input: dict[str, Any]
if args_str:
try:
parsed = json.loads(args_str)
@@ -840,10 +838,10 @@ class OpenAINormalizer(FormatNormalizer):
def _openai_tool_result_message_to_block(
self,
msg: Dict[str, Any],
msg: dict[str, Any],
tool_call_id: str,
) -> Tuple[Optional[ToolResultBlock], Dict[str, int]]:
dropped: Dict[str, int] = {}
) -> tuple[ToolResultBlock | None, dict[str, int]]:
dropped: dict[str, int] = {}
content = msg.get("content")
if content is None:
return ToolResultBlock(tool_use_id=tool_call_id, output=None, content_text=None, extra={"openai": msg}), dropped
@@ -887,7 +885,7 @@ class OpenAINormalizer(FormatNormalizer):
dropped,
)
def _openai_usage_to_internal(self, usage: Any) -> Optional[UsageInfo]:
def _openai_usage_to_internal(self, usage: Any) -> UsageInfo | None:
if not isinstance(usage, dict):
return None
@@ -908,9 +906,9 @@ class OpenAINormalizer(FormatNormalizer):
extra={"openai": extra} if extra else {},
)
def _blocks_to_openai_content(self, blocks: List[ContentBlock]) -> Optional[Union[str, List[Dict[str, Any]]]]:
parts: List[Dict[str, Any]] = []
text_parts: List[str] = []
def _blocks_to_openai_content(self, blocks: list[ContentBlock]) -> str | list[dict[str, Any]] | None:
parts: list[dict[str, Any]] = []
text_parts: list[str] = []
for b in blocks:
if isinstance(b, TextBlock):
@@ -948,9 +946,9 @@ class OpenAINormalizer(FormatNormalizer):
# OpenAI content 可以是空字符串;但作为响应 message.content 通常允许为 ""/None。
return ""
def _split_blocks(self, blocks: List[ContentBlock]) -> Tuple[List[ContentBlock], List[ToolUseBlock]]:
content_blocks: List[ContentBlock] = []
tool_blocks: List[ToolUseBlock] = []
def _split_blocks(self, blocks: list[ContentBlock]) -> tuple[list[ContentBlock], list[ToolUseBlock]]:
content_blocks: list[ContentBlock] = []
tool_blocks: list[ToolUseBlock] = []
for b in blocks:
if isinstance(b, ToolUseBlock):
tool_blocks.append(b)
@@ -963,7 +961,7 @@ class OpenAINormalizer(FormatNormalizer):
content_blocks.append(b)
return content_blocks, tool_blocks
def _internal_message_to_openai_messages(self, msg: InternalMessage) -> List[Dict[str, Any]]:
def _internal_message_to_openai_messages(self, msg: InternalMessage) -> list[dict[str, Any]]:
if msg.role == Role.USER:
return self._user_message_to_openai(msg)
if msg.role == Role.ASSISTANT:
@@ -974,9 +972,9 @@ class OpenAINormalizer(FormatNormalizer):
return [{"role": "tool", "content": content_value or ""}]
return [{"role": "user", "content": ""}]
def _user_message_to_openai(self, msg: InternalMessage) -> List[Dict[str, Any]]:
out: List[Dict[str, Any]] = []
pending: List[ContentBlock] = []
def _user_message_to_openai(self, msg: InternalMessage) -> list[dict[str, Any]]:
out: list[dict[str, Any]] = []
pending: list[ContentBlock] = []
def flush_user() -> None:
nonlocal pending
@@ -1008,9 +1006,9 @@ class OpenAINormalizer(FormatNormalizer):
return out
def _assistant_message_to_openai(self, msg: InternalMessage) -> Dict[str, Any]:
content_blocks: List[ContentBlock] = []
tool_blocks: List[ToolUseBlock] = []
def _assistant_message_to_openai(self, msg: InternalMessage) -> dict[str, Any]:
content_blocks: list[ContentBlock] = []
tool_blocks: list[ToolUseBlock] = []
for b in msg.content:
if isinstance(b, ToolUseBlock):
@@ -1022,7 +1020,7 @@ class OpenAINormalizer(FormatNormalizer):
continue
content_blocks.append(b)
out: Dict[str, Any] = {"role": "assistant"}
out: dict[str, Any] = {"role": "assistant"}
content_value = self._blocks_to_openai_content(content_blocks)
out["content"] = content_value if content_value is not None else ""
@@ -1031,7 +1029,7 @@ class OpenAINormalizer(FormatNormalizer):
return out
def _tool_result_block_to_openai_message(self, block: ToolResultBlock) -> Dict[str, Any]:
def _tool_result_block_to_openai_message(self, block: ToolResultBlock) -> dict[str, Any]:
content: str
if block.content_text is not None:
content = block.content_text
@@ -1048,7 +1046,7 @@ class OpenAINormalizer(FormatNormalizer):
"content": content,
}
def _tool_use_block_to_openai_call(self, block: ToolUseBlock, index: int) -> Dict[str, Any]:
def _tool_use_block_to_openai_call(self, block: ToolUseBlock, index: int) -> dict[str, Any]:
return {
"index": index,
"id": block.tool_id or f"call_{index}",
@@ -1085,7 +1083,7 @@ class OpenAINormalizer(FormatNormalizer):
except ValueError:
return ErrorType.UNKNOWN
def _optional_int(self, value: Any) -> Optional[int]:
def _optional_int(self, value: Any) -> int | None:
if value is None:
return None
try:
@@ -1093,7 +1091,7 @@ class OpenAINormalizer(FormatNormalizer):
except (TypeError, ValueError):
return None
def _optional_float(self, value: Any) -> Optional[float]:
def _optional_float(self, value: Any) -> float | None:
if value is None:
return None
try:
@@ -1101,7 +1099,7 @@ class OpenAINormalizer(FormatNormalizer):
except (TypeError, ValueError):
return None
def _coerce_str_list(self, value: Any) -> Optional[List[str]]:
def _coerce_str_list(self, value: Any) -> list[str] | None:
if value is None:
return None
if isinstance(value, str):
@@ -1110,14 +1108,14 @@ class OpenAINormalizer(FormatNormalizer):
return [str(x) for x in value if x is not None]
return None
def _extract_extra(self, payload: Dict[str, Any], known_keys: set[str]) -> Dict[str, Any]:
def _extract_extra(self, payload: dict[str, Any], known_keys: set[str]) -> dict[str, Any]:
return {k: v for k, v in payload.items() if k not in known_keys}
def _merge_dropped(self, target: Dict[str, int], source: Dict[str, int]) -> None:
def _merge_dropped(self, target: dict[str, int], source: dict[str, int]) -> None:
for k, v in source.items():
target[k] = target.get(k, 0) + int(v)
def _ensure_tool_block_index(self, ss: Dict[str, Any], tool_key: str) -> int:
def _ensure_tool_block_index(self, ss: dict[str, Any], tool_key: str) -> int:
mapping = ss.get("tool_id_to_block_index")
if not isinstance(mapping, dict):
mapping = {}
@@ -1131,7 +1129,7 @@ class OpenAINormalizer(FormatNormalizer):
ss["next_block_index"] = next_idx + 1
return next_idx
def _ensure_tool_call_index(self, ss: Dict[str, Any], tool_id: str) -> int:
def _ensure_tool_call_index(self, ss: dict[str, Any], tool_id: str) -> int:
mapping = ss.get("tool_id_to_index")
if not isinstance(mapping, dict):
mapping = {}

View File

@@ -10,11 +10,10 @@ OpenAI CLI / Responses Normalizer (OPENAI_CLI)
- 未识别的字段会进入 extra/raw未知内容块保留在 internal但默认输出阶段会丢弃。
"""
from __future__ import annotations
import json
import time
from typing import Any, Dict, List, Optional, Tuple, Union
from typing import Any
from src.core.api_format.conversion.field_mappings import (
ERROR_TYPE_MAPPINGS,
@@ -65,7 +64,7 @@ class OpenAICliNormalizer(FormatNormalizer):
supports_images=True,
)
_ERROR_TYPE_TO_OPENAI: Dict[ErrorType, str] = {
_ERROR_TYPE_TO_OPENAI: dict[ErrorType, str] = {
ErrorType.INVALID_REQUEST: "invalid_request_error",
ErrorType.AUTHENTICATION: "invalid_api_key",
ErrorType.PERMISSION_DENIED: "invalid_request_error",
@@ -82,12 +81,12 @@ class OpenAICliNormalizer(FormatNormalizer):
# Requests
# =========================
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
model = str(request.get("model") or "")
instructions_text = request.get("instructions")
instructions: List[InstructionSegment] = []
system_text: Optional[str] = None
instructions: list[InstructionSegment] = []
system_text: str | None = None
if isinstance(instructions_text, str) and instructions_text.strip():
system_text = instructions_text
instructions.append(InstructionSegment(role=Role.SYSTEM, text=instructions_text))
@@ -118,8 +117,8 @@ class OpenAICliNormalizer(FormatNormalizer):
return internal
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
result: Dict[str, Any] = {
def request_from_internal(self, internal: InternalRequest) -> dict[str, Any]:
result: dict[str, Any] = {
"model": internal.model,
"input": self._internal_messages_to_input(internal.messages),
}
@@ -164,7 +163,7 @@ class OpenAICliNormalizer(FormatNormalizer):
# Responses
# =========================
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse:
def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
payload = self._unwrap_response_object(response)
rid = str(payload.get("id") or "")
@@ -191,8 +190,8 @@ class OpenAICliNormalizer(FormatNormalizer):
self,
internal: InternalResponse,
*,
requested_model: Optional[str] = None,
) -> Dict[str, Any]:
requested_model: str | None = None,
) -> dict[str, Any]:
text = self._collapse_internal_text(internal.content)
output_message = {
@@ -203,7 +202,7 @@ class OpenAICliNormalizer(FormatNormalizer):
}
usage = internal.usage or UsageInfo()
usage_obj: Dict[str, Any] = {
usage_obj: dict[str, Any] = {
"input_tokens": usage.input_tokens,
"output_tokens": usage.output_tokens,
"total_tokens": usage.total_tokens or (usage.input_tokens + usage.output_tokens),
@@ -228,11 +227,11 @@ class OpenAICliNormalizer(FormatNormalizer):
def stream_chunk_to_internal(
self,
chunk: Dict[str, Any],
chunk: dict[str, Any],
state: StreamState,
) -> List[InternalStreamEvent]:
) -> list[InternalStreamEvent]:
ss = state.substate(self.FORMAT_ID)
events: List[InternalStreamEvent] = []
events: list[InternalStreamEvent] = []
# 统一错误结构(最佳努力)
if isinstance(chunk, dict) and "error" in chunk:
@@ -392,11 +391,11 @@ class OpenAICliNormalizer(FormatNormalizer):
self,
event: InternalStreamEvent,
state: StreamState,
) -> List[Dict[str, Any]]:
) -> list[dict[str, Any]]:
ss = state.substate(self.FORMAT_ID)
out: List[Dict[str, Any]] = []
out: list[dict[str, Any]] = []
def event_block(payload: Dict[str, Any]) -> Dict[str, Any]:
def event_block(payload: dict[str, Any]) -> dict[str, Any]:
# OpenAI Responses SSE 的 payload 通常自带 type 字段;这里强制保证
return payload
@@ -463,10 +462,10 @@ class OpenAICliNormalizer(FormatNormalizer):
# Error conversion
# =========================
def is_error_response(self, response: Dict[str, Any]) -> bool:
def is_error_response(self, response: dict[str, Any]) -> bool:
return isinstance(response, dict) and "error" in response
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError:
def error_to_internal(self, error_response: dict[str, Any]) -> InternalError:
err = error_response.get("error") if isinstance(error_response, dict) else None
err = err if isinstance(err, dict) else {}
@@ -484,9 +483,9 @@ class OpenAICliNormalizer(FormatNormalizer):
extra={"openai_cli": {"error": err}, "raw": {"type": raw_type}},
)
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]:
def error_from_internal(self, internal: InternalError) -> dict[str, Any]:
type_str = self._ERROR_TYPE_TO_OPENAI.get(internal.type, "server_error")
payload: Dict[str, Any] = {"type": type_str, "message": internal.message}
payload: dict[str, Any] = {"type": type_str, "message": internal.message}
if internal.code is not None:
payload["code"] = internal.code
if internal.param is not None:
@@ -497,7 +496,7 @@ class OpenAICliNormalizer(FormatNormalizer):
# Helpers
# =========================
def _unwrap_response_object(self, response: Dict[str, Any]) -> Dict[str, Any]:
def _unwrap_response_object(self, response: dict[str, Any]) -> dict[str, Any]:
if not isinstance(response, dict):
return {}
resp_inner = response.get("response")
@@ -506,8 +505,8 @@ class OpenAICliNormalizer(FormatNormalizer):
return resp_inner
return response
def _extract_output_text_blocks(self, payload: Dict[str, Any]) -> Tuple[List[ContentBlock], Dict[str, Any]]:
text_parts: List[str] = []
def _extract_output_text_blocks(self, payload: dict[str, Any]) -> tuple[list[ContentBlock], dict[str, Any]]:
text_parts: list[str] = []
output = payload.get("output")
if isinstance(output, list):
@@ -532,12 +531,12 @@ class OpenAICliNormalizer(FormatNormalizer):
if not text_parts and isinstance(payload.get("output_text"), str):
text_parts.append(payload.get("output_text") or "")
blocks: List[ContentBlock] = []
blocks: list[ContentBlock] = []
text = "".join(text_parts)
if text:
blocks.append(TextBlock(text=text))
extra: Dict[str, Any] = {"raw": {"openai_cli_output": output}} if output is not None else {}
extra: dict[str, Any] = {"raw": {"openai_cli_output": output}} if output is not None else {}
return blocks, extra
def _usage_to_internal(self, usage: Any) -> UsageInfo:
@@ -553,14 +552,14 @@ class OpenAICliNormalizer(FormatNormalizer):
extra={"openai_cli": {"usage": usage}},
)
def _collapse_internal_text(self, blocks: List[ContentBlock]) -> str:
parts: List[str] = []
def _collapse_internal_text(self, blocks: list[ContentBlock]) -> str:
parts: list[str] = []
for block in blocks:
if isinstance(block, TextBlock) and block.text:
parts.append(block.text)
return "".join(parts)
def _input_to_internal_messages(self, input_data: Any) -> List[InternalMessage]:
def _input_to_internal_messages(self, input_data: Any) -> list[InternalMessage]:
if input_data is None:
return []
@@ -575,7 +574,7 @@ class OpenAICliNormalizer(FormatNormalizer):
if not isinstance(input_data, list):
return [InternalMessage(role=Role.USER, content=[UnknownBlock(raw_type="input", payload={"input": input_data})])]
messages: List[InternalMessage] = []
messages: list[InternalMessage] = []
for item in input_data:
if not isinstance(item, dict):
continue
@@ -624,7 +623,7 @@ class OpenAICliNormalizer(FormatNormalizer):
# reasoning -> assistant 消息,提取 summary 作为文本
if item_type == "reasoning":
summary_parts: List[str] = []
summary_parts: list[str] = []
summary = item.get("summary")
if isinstance(summary, list):
for s in summary:
@@ -638,7 +637,7 @@ class OpenAICliNormalizer(FormatNormalizer):
summary_parts.append(summary)
# 如果有 summary 文本,创建一个 UnknownBlock 保留原始结构
reasoning_blocks: List[ContentBlock] = []
reasoning_blocks: list[ContentBlock] = []
if summary_parts:
# 保留 reasoning 的 summary 作为 UnknownBlock便于输出时决策
reasoning_blocks.append(UnknownBlock(
@@ -662,7 +661,7 @@ class OpenAICliNormalizer(FormatNormalizer):
return messages
def _responses_content_to_blocks(self, content: Any) -> List[ContentBlock]:
def _responses_content_to_blocks(self, content: Any) -> list[ContentBlock]:
if content is None:
return []
if isinstance(content, str):
@@ -673,7 +672,7 @@ class OpenAICliNormalizer(FormatNormalizer):
if not isinstance(content, list):
return [UnknownBlock(raw_type="content", payload={"content": content})]
blocks: List[ContentBlock] = []
blocks: list[ContentBlock] = []
for part in content:
if isinstance(part, str):
if part:
@@ -690,8 +689,8 @@ class OpenAICliNormalizer(FormatNormalizer):
blocks.append(UnknownBlock(raw_type=ptype or "unknown", payload=part))
return blocks
def _internal_messages_to_input(self, messages: List[InternalMessage]) -> List[Dict[str, Any]]:
out: List[Dict[str, Any]] = []
def _internal_messages_to_input(self, messages: list[InternalMessage]) -> list[dict[str, Any]]:
out: list[dict[str, Any]] = []
for msg in messages:
# ToolUseBlock -> function_call
for block in msg.content:
@@ -729,7 +728,7 @@ class OpenAICliNormalizer(FormatNormalizer):
# 普通 messageTextBlock
role = self._role_to_openai(msg.role)
content_items: List[Dict[str, Any]] = []
content_items: list[dict[str, Any]] = []
has_text = False
for block in msg.content:
@@ -748,10 +747,10 @@ class OpenAICliNormalizer(FormatNormalizer):
return out
def _tools_to_internal(self, tools: Any) -> Optional[List[ToolDefinition]]:
def _tools_to_internal(self, tools: Any) -> list[ToolDefinition] | None:
if not isinstance(tools, list):
return None
out: List[ToolDefinition] = []
out: list[ToolDefinition] = []
for tool in tools:
if not isinstance(tool, dict):
continue
@@ -783,7 +782,7 @@ class OpenAICliNormalizer(FormatNormalizer):
)
return out or None
def _tool_choice_to_internal(self, tool_choice: Any) -> Optional[ToolChoice]:
def _tool_choice_to_internal(self, tool_choice: Any) -> ToolChoice | None:
if tool_choice is None:
return None
if isinstance(tool_choice, str):
@@ -802,7 +801,7 @@ class OpenAICliNormalizer(FormatNormalizer):
return ToolChoice(type=ToolChoiceType.AUTO, extra={"raw": tool_choice})
def _tool_choice_to_openai(self, tool_choice: ToolChoice) -> Union[str, Dict[str, Any]]:
def _tool_choice_to_openai(self, tool_choice: ToolChoice) -> str | dict[str, Any]:
if tool_choice.type == ToolChoiceType.NONE:
return "none"
if tool_choice.type == ToolChoiceType.AUTO:
@@ -840,7 +839,7 @@ class OpenAICliNormalizer(FormatNormalizer):
return "tool"
return "user"
def _optional_int(self, value: Any) -> Optional[int]:
def _optional_int(self, value: Any) -> int | None:
if value is None:
return None
try:
@@ -848,7 +847,7 @@ class OpenAICliNormalizer(FormatNormalizer):
except (TypeError, ValueError):
return None
def _optional_float(self, value: Any) -> Optional[float]:
def _optional_float(self, value: Any) -> float | None:
if value is None:
return None
try:
@@ -856,13 +855,13 @@ class OpenAICliNormalizer(FormatNormalizer):
except (TypeError, ValueError):
return None
def _coerce_str_list(self, value: Any) -> Optional[List[str]]:
def _coerce_str_list(self, value: Any) -> list[str] | None:
if value is None:
return None
if isinstance(value, str):
return [value]
if isinstance(value, list):
out: List[str] = []
out: list[str] = []
for item in value:
if item is None:
continue
@@ -870,14 +869,14 @@ class OpenAICliNormalizer(FormatNormalizer):
return out
return [str(value)]
def _extract_extra(self, payload: Dict[str, Any], keep_keys: set[str]) -> Dict[str, Any]:
def _extract_extra(self, payload: dict[str, Any], keep_keys: set[str]) -> dict[str, Any]:
if not isinstance(payload, dict):
return {}
return {k: v for k, v in payload.items() if k not in keep_keys}
def _join_instructions(self, internal: InternalRequest) -> str:
if internal.instructions:
parts: List[str] = []
parts: list[str] = []
for seg in internal.instructions:
if seg.text:
parts.append(seg.text)

View File

@@ -9,12 +9,12 @@ source -> internal -> target
- 转换失败将抛出 `FormatConversionError`(不再静默回退)。
"""
from __future__ import annotations
import threading
import time
from contextlib import contextmanager
from typing import Any, Dict, Generator, List, Optional
from typing import Any
from collections.abc import Generator
from src.core.logger import logger
from src.core.metrics import format_conversion_duration_seconds, format_conversion_total
@@ -29,7 +29,7 @@ def _track_conversion_metrics(
direction: str,
source: str,
target: str,
) -> Generator[None, None, None]:
) -> Generator[None]:
start = time.perf_counter()
try:
yield
@@ -47,13 +47,13 @@ class FormatConversionRegistry:
"""基于 Normalizer 的格式转换注册表"""
def __init__(self) -> None:
self._normalizers: Dict[str, FormatNormalizer] = {}
self._normalizers: dict[str, FormatNormalizer] = {}
def register(self, normalizer: FormatNormalizer) -> None:
self._normalizers[str(normalizer.FORMAT_ID).upper()] = normalizer
logger.info(f"[FormatConversionRegistry] 注册 normalizer: {normalizer.FORMAT_ID}")
def get_normalizer(self, format_id: str) -> Optional[FormatNormalizer]:
def get_normalizer(self, format_id: str) -> FormatNormalizer | None:
return self._normalizers.get(str(format_id).upper())
def _require_normalizer(self, format_id: str) -> FormatNormalizer:
@@ -66,10 +66,10 @@ class FormatConversionRegistry:
def convert_request(
self,
request: Dict[str, Any],
request: dict[str, Any],
source_format: str,
target_format: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
if str(source_format).upper() == str(target_format).upper():
return request
@@ -85,12 +85,12 @@ class FormatConversionRegistry:
def convert_response(
self,
response: Dict[str, Any],
response: dict[str, Any],
source_format: str,
target_format: str,
*,
requested_model: Optional[str] = None,
) -> Dict[str, Any]:
requested_model: str | None = None,
) -> dict[str, Any]:
"""转换响应格式
Args:
@@ -124,10 +124,10 @@ class FormatConversionRegistry:
def convert_error_response(
self,
error_response: Dict[str, Any],
error_response: dict[str, Any],
source_format: str,
target_format: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
if str(source_format).upper() == str(target_format).upper():
return error_response
@@ -152,11 +152,11 @@ class FormatConversionRegistry:
def convert_stream_chunk(
self,
chunk: Dict[str, Any],
chunk: dict[str, Any],
source_format: str,
target_format: str,
state: Optional[StreamState] = None,
) -> List[Dict[str, Any]]:
state: StreamState | None = None,
) -> list[dict[str, Any]]:
if str(source_format).upper() == str(target_format).upper():
return [chunk]
@@ -182,7 +182,7 @@ class FormatConversionRegistry:
with _track_conversion_metrics("stream", str(source_format).upper(), str(target_format).upper()):
try:
events = src.stream_chunk_to_internal(chunk, state)
out: List[Dict[str, Any]] = []
out: list[dict[str, Any]] = []
for event in events:
out.extend(tgt.stream_event_from_internal(event, state))
return out
@@ -226,10 +226,10 @@ class FormatConversionRegistry:
return self.can_convert_stream(format_a, format_b) and self.can_convert_stream(format_b, format_a)
return True
def list_normalizers(self) -> List[str]:
def list_normalizers(self) -> list[str]:
return sorted(self._normalizers.keys())
def get_supported_targets(self, source_format: str) -> List[str]:
def get_supported_targets(self, source_format: str) -> list[str]:
src = str(source_format).upper()
if src not in self._normalizers:
return []

View File

@@ -4,11 +4,10 @@
用于把 OpenAI/Claude/Gemini 的流式协议映射为统一事件序列,再由目标格式 Normalizer 输出。
"""
from __future__ import annotations
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Dict, Optional, Union
from typing import Any
from .internal import ContentType, InternalError, StopReason, UsageInfo
@@ -34,8 +33,8 @@ class MessageStartEvent:
type: StreamEventType = field(default=StreamEventType.MESSAGE_START, init=False)
message_id: str = ""
model: str = ""
usage: Optional[UsageInfo] = None # Claude 流式响应的 message_start 可能包含 usage
extra: Dict[str, Any] = field(default_factory=dict)
usage: UsageInfo | None = None # Claude 流式响应的 message_start 可能包含 usage
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -46,9 +45,9 @@ class ContentBlockStartEvent:
block_index: int = 0
block_type: ContentType = ContentType.TEXT
# 工具调用时使用TOOL_USE block
tool_id: Optional[str] = None
tool_name: Optional[str] = None
extra: Dict[str, Any] = field(default_factory=dict)
tool_id: str | None = None
tool_name: str | None = None
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -58,7 +57,7 @@ class ContentDeltaEvent:
type: StreamEventType = field(default=StreamEventType.CONTENT_DELTA, init=False)
block_index: int = 0
text_delta: str = ""
extra: Dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -69,7 +68,7 @@ class ToolCallDeltaEvent:
block_index: int = 0
tool_id: str = ""
input_delta: str = "" # JSON 字符串片段
extra: Dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -78,7 +77,7 @@ class ContentBlockStopEvent:
type: StreamEventType = field(default=StreamEventType.CONTENT_BLOCK_STOP, init=False)
block_index: int = 0
extra: Dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -86,9 +85,9 @@ class MessageStopEvent:
"""消息结束事件"""
type: StreamEventType = field(default=StreamEventType.MESSAGE_STOP, init=False)
stop_reason: Optional[StopReason] = None
usage: Optional[UsageInfo] = None
extra: Dict[str, Any] = field(default_factory=dict)
stop_reason: StopReason | None = None
usage: UsageInfo | None = None
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -97,7 +96,7 @@ class UsageEvent:
type: StreamEventType = field(default=StreamEventType.USAGE, init=False)
usage: UsageInfo = field(default_factory=UsageInfo)
extra: Dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -106,7 +105,7 @@ class ErrorEvent:
type: StreamEventType = field(default=StreamEventType.ERROR, init=False)
error: InternalError
extra: Dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -115,21 +114,21 @@ class UnknownStreamEvent:
type: StreamEventType = field(default=StreamEventType.UNKNOWN, init=False)
raw_type: str = ""
payload: Dict[str, Any] = field(default_factory=dict)
extra: Dict[str, Any] = field(default_factory=dict)
payload: dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
InternalStreamEvent = Union[
MessageStartEvent,
ContentBlockStartEvent,
ContentDeltaEvent,
ToolCallDeltaEvent,
ContentBlockStopEvent,
MessageStopEvent,
UsageEvent,
ErrorEvent,
UnknownStreamEvent,
]
InternalStreamEvent = (
MessageStartEvent
| ContentBlockStartEvent
| ContentDeltaEvent
| ToolCallDeltaEvent
| ContentBlockStopEvent
| MessageStopEvent
| UsageEvent
| ErrorEvent
| UnknownStreamEvent
)
__all__ = [

View File

@@ -5,10 +5,9 @@
每个 Normalizer 通过 `substate(format_id)` 获取自己的隔离状态字典。
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Dict
from typing import Any
@dataclass
@@ -26,12 +25,12 @@ class StreamState:
message_id: str = ""
# Registry/调用层的通用扩展信息(与具体格式无关)
extra: Dict[str, Any] = field(default_factory=dict)
extra: dict[str, Any] = field(default_factory=dict)
# 各 Normalizer 的隔离状态key: FORMAT_ID
by_format: Dict[str, Dict[str, Any]] = field(default_factory=dict)
by_format: dict[str, dict[str, Any]] = field(default_factory=dict)
def substate(self, format_id: str) -> Dict[str, Any]:
def substate(self, format_id: str) -> dict[str, Any]:
"""获取指定格式的隔离子状态"""
key = str(format_id).upper()
return self.by_format.setdefault(key, {})

View File

@@ -4,9 +4,8 @@ API 格式检测
提供从请求头、响应内容等检测 API 格式的函数。
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Dict, Optional, Tuple
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from starlette.requests import Request
@@ -16,10 +15,10 @@ from src.core.api_format.metadata import API_FORMAT_DEFINITIONS, ApiFormatDefini
def _extract_api_key_by_definition(
headers: Dict[str, str],
query_params: Optional[Dict[str, str]],
headers: dict[str, str],
query_params: dict[str, str] | None,
definition: ApiFormatDefinition,
) -> Tuple[Optional[str], str]:
) -> tuple[str | None, str]:
"""
根据格式定义从请求中提取 API Key
@@ -64,9 +63,9 @@ def _extract_api_key_by_definition(
def detect_format_from_request(
headers: Dict[str, str],
query_params: Optional[Dict[str, str]] = None,
) -> Tuple[APIFormat, Optional[str], str]:
headers: dict[str, str],
query_params: dict[str, str] | None = None,
) -> tuple[APIFormat, str | None, str]:
"""
从请求头检测 API 格式和 API Key
@@ -107,8 +106,8 @@ def detect_format_from_request(
def detect_format_and_key_from_starlette(
request: "Request",
) -> Tuple[str, Optional[str], str]:
request: Request,
) -> tuple[str, str | None, str]:
"""
从 Starlette Request 对象检测 API 格式和 API Key
@@ -135,7 +134,7 @@ def detect_format_and_key_from_starlette(
def detect_format_from_response(
response_data: dict,
) -> Optional[APIFormat]:
) -> APIFormat | None:
"""
从响应内容检测 API 格式

View File

@@ -12,7 +12,8 @@
from __future__ import annotations
from typing import AbstractSet, Any, Dict, FrozenSet, Optional, Set
from collections.abc import Set as AbstractSet
from typing import Any
from src.core.api_format.enums import APIFormat
from src.core.api_format.metadata import get_auth_config, get_extra_headers, get_protected_keys
@@ -23,7 +24,7 @@ from src.core.api_format.metadata import get_auth_config, get_extra_headers, get
# =============================================================================
# 转发给上游时需要剔除的头部(系统管理 + 认证替换)
UPSTREAM_DROP_HEADERS: FrozenSet[str] = frozenset(
UPSTREAM_DROP_HEADERS: frozenset[str] = frozenset(
{
# 认证头 - 会被替换为 Provider 的认证
"authorization",
@@ -41,7 +42,7 @@ UPSTREAM_DROP_HEADERS: FrozenSet[str] = frozenset(
# 最小必脱敏集合(编译时常量,用于快速路径)
# 完整脱敏应使用 SystemConfigService.get_sensitive_headers()
CORE_REDACT_HEADERS: FrozenSet[str] = frozenset(
CORE_REDACT_HEADERS: frozenset[str] = frozenset(
{
"authorization",
"x-api-key",
@@ -50,7 +51,7 @@ CORE_REDACT_HEADERS: FrozenSet[str] = frozenset(
)
# Hop-by-hop 头部 (RFC 7230)
HOP_BY_HOP_HEADERS: FrozenSet[str] = frozenset(
HOP_BY_HOP_HEADERS: frozenset[str] = frozenset(
{
"connection",
"keep-alive",
@@ -64,7 +65,7 @@ HOP_BY_HOP_HEADERS: FrozenSet[str] = frozenset(
)
# 响应时需要过滤的头部body-dependent + hop-by-hop
RESPONSE_DROP_HEADERS: FrozenSet[str] = (
RESPONSE_DROP_HEADERS: frozenset[str] = (
frozenset(
{
"content-length",
@@ -82,7 +83,7 @@ RESPONSE_DROP_HEADERS: FrozenSet[str] = (
# =============================================================================
def normalize_headers(headers: Dict[str, str]) -> Dict[str, str]:
def normalize_headers(headers: dict[str, str]) -> dict[str, str]:
"""
将请求头 key 统一为小写
@@ -92,7 +93,7 @@ def normalize_headers(headers: Dict[str, str]) -> Dict[str, str]:
return {k.lower(): v for k, v in headers.items()}
def get_header_value(headers: Dict[str, str], key: str, default: str = "") -> str:
def get_header_value(headers: dict[str, str], key: str, default: str = "") -> str:
"""
大小写不敏感地获取请求头值
@@ -117,7 +118,7 @@ def get_header_value(headers: Dict[str, str], key: str, default: str = "") -> st
# =============================================================================
def extract_client_api_key(headers: Dict[str, str], api_format: APIFormat) -> Optional[str]:
def extract_client_api_key(headers: dict[str, str], api_format: APIFormat) -> str | None:
"""
从客户端请求头提取 API Key
@@ -147,10 +148,10 @@ def extract_client_api_key(headers: Dict[str, str], api_format: APIFormat) -> Op
def extract_client_api_key_with_query(
headers: Dict[str, str],
query_params: Optional[Dict[str, str]],
headers: dict[str, str],
query_params: dict[str, str] | None,
api_format: APIFormat,
) -> Optional[str]:
) -> str | None:
"""
从客户端请求头或 URL 参数提取 API Key
@@ -184,10 +185,10 @@ def extract_client_api_key_with_query(
def detect_capabilities(
headers: Dict[str, str],
headers: dict[str, str],
api_format: APIFormat,
request_body: Optional[Dict[str, Any]] = None, # noqa: ARG001 - 预留给部分格式使用
) -> Dict[str, bool]:
request_body: dict[str, Any] | None = None, # noqa: ARG001 - 预留给部分格式使用
) -> dict[str, bool]:
"""
从请求头检测能力需求
@@ -203,7 +204,7 @@ def detect_capabilities(
能力需求字典,如 {"context_1m": True}
"""
requirements: Dict[str, bool] = {}
requirements: dict[str, bool] = {}
if api_format in (APIFormat.CLAUDE, APIFormat.CLAUDE_CLI):
beta_header = get_header_value(headers, "anthropic-beta")
@@ -228,20 +229,20 @@ class HeaderBuilder:
def __init__(self) -> None:
# key: (original_case_key, value)
self._headers: Dict[str, tuple[str, str]] = {}
self._headers: dict[str, tuple[str, str]] = {}
def add(self, key: str, value: str) -> "HeaderBuilder":
def add(self, key: str, value: str) -> HeaderBuilder:
"""添加单个头部(会覆盖同名头部)"""
self._headers[key.lower()] = (key, value)
return self
def add_many(self, headers: Dict[str, str]) -> "HeaderBuilder":
def add_many(self, headers: dict[str, str]) -> HeaderBuilder:
"""批量添加头部"""
for k, v in headers.items():
self.add(k, v)
return self
def add_protected(self, headers: Dict[str, str], protected_keys: AbstractSet[str]) -> "HeaderBuilder":
def add_protected(self, headers: dict[str, str], protected_keys: AbstractSet[str]) -> HeaderBuilder:
"""
添加头部但保护指定的 key 不被覆盖
@@ -253,13 +254,13 @@ class HeaderBuilder:
self.add(k, v)
return self
def remove(self, keys: FrozenSet[str]) -> "HeaderBuilder":
def remove(self, keys: frozenset[str]) -> HeaderBuilder:
"""移除指定的头部"""
for k in keys:
self._headers.pop(k.lower(), None)
return self
def rename(self, from_key: str, to_key: str) -> "HeaderBuilder":
def rename(self, from_key: str, to_key: str) -> HeaderBuilder:
"""
重命名头部(保留原值)
@@ -273,9 +274,9 @@ class HeaderBuilder:
def apply_rules(
self,
rules: list[Dict[str, Any]],
protected_keys: Optional[AbstractSet[str]] = None,
) -> "HeaderBuilder":
rules: list[dict[str, Any]],
protected_keys: AbstractSet[str] | None = None,
) -> HeaderBuilder:
"""
应用请求头规则
@@ -314,20 +315,20 @@ class HeaderBuilder:
return self
def build(self) -> Dict[str, str]:
def build(self) -> dict[str, str]:
"""构建最终的头部字典"""
return {original_key: value for original_key, value in self._headers.values()}
def build_upstream_headers(
original_headers: Dict[str, str],
original_headers: dict[str, str],
api_format: APIFormat,
provider_api_key: str,
*,
endpoint_headers: Optional[Dict[str, str]] = None,
extra_headers: Optional[Dict[str, str]] = None,
drop_headers: Optional[FrozenSet[str]] = None,
) -> Dict[str, str]:
endpoint_headers: dict[str, str] | None = None,
extra_headers: dict[str, str] | None = None,
drop_headers: frozenset[str] | None = None,
) -> dict[str, str]:
"""
构建发送给上游 Provider 的请求头
@@ -386,10 +387,10 @@ def build_upstream_headers(
def merge_headers_with_protection(
base_headers: Dict[str, str],
extra_headers: Optional[Dict[str, str]],
protected_keys: FrozenSet[str] | Set[str],
) -> Dict[str, str]:
base_headers: dict[str, str],
extra_headers: dict[str, str] | None,
protected_keys: frozenset[str] | set[str],
) -> dict[str, str]:
"""
合并头部但保护指定的 key 不被覆盖
@@ -418,9 +419,9 @@ def merge_headers_with_protection(
def filter_response_headers(
headers: Optional[Dict[str, str]],
drop_headers: Optional[FrozenSet[str]] = None,
) -> Dict[str, str]:
headers: dict[str, str] | None,
drop_headers: frozenset[str] | None = None,
) -> dict[str, str]:
"""
过滤上游响应头中不应透传给客户端的字段
@@ -446,9 +447,9 @@ def filter_response_headers(
def redact_headers_for_log(
headers: Dict[str, str],
redact_keys: Optional[FrozenSet[str]] = None,
) -> Dict[str, str]:
headers: dict[str, str],
redact_keys: frozenset[str] | None = None,
) -> dict[str, str]:
"""
将敏感头部值替换为 *** 用于日志记录
@@ -487,7 +488,7 @@ def build_adapter_base_headers(
api_key: str,
*,
include_extra: bool = True,
) -> Dict[str, str]:
) -> dict[str, str]:
"""
根据 API 格式构建基础请求头
@@ -504,7 +505,7 @@ def build_adapter_base_headers(
auth_header, auth_type = get_auth_config(api_format)
auth_value = f"Bearer {api_key}" if auth_type == "bearer" else api_key
headers: Dict[str, str] = {
headers: dict[str, str] = {
auth_header: auth_value,
"Content-Type": "application/json",
}
@@ -520,8 +521,8 @@ def build_adapter_base_headers(
def build_adapter_headers(
api_format: APIFormat,
api_key: str,
extra_headers: Optional[Dict[str, str]] = None,
) -> Dict[str, str]:
extra_headers: dict[str, str] | None = None,
) -> dict[str, str]:
"""
构建完整的 Adapter 请求头
@@ -565,8 +566,8 @@ def get_adapter_protected_keys(api_format: APIFormat) -> tuple[str, ...]:
def extract_set_headers_from_rules(
header_rules: Optional[list[Dict[str, Any]]],
) -> Optional[Dict[str, str]]:
header_rules: list[dict[str, Any]] | None,
) -> dict[str, str] | None:
"""
从 header_rules 中提取 set 操作生成的头部字典
@@ -582,7 +583,7 @@ def extract_set_headers_from_rules(
if not header_rules:
return None
headers: Dict[str, str] = {}
headers: dict[str, str] = {}
for rule in header_rules:
if rule.get("action") == "set":
key = rule.get("key", "")
@@ -593,7 +594,7 @@ def extract_set_headers_from_rules(
return headers if headers else None
def get_extra_headers_from_endpoint(endpoint: Any) -> Optional[Dict[str, str]]:
def get_extra_headers_from_endpoint(endpoint: Any) -> dict[str, str] | None:
"""
从 endpoint 提取额外请求头

View File

@@ -13,13 +13,12 @@ API 格式元数据定义
definition = get_api_format_definition(APIFormat.CLAUDE)
"""
from __future__ import annotations
import re
from dataclasses import dataclass, field
from functools import lru_cache
from types import MappingProxyType
from typing import Dict, Iterable, List, Mapping, MutableMapping, Optional, Sequence, Union
from collections.abc import Iterable, Mapping, MutableMapping, Sequence
from .enums import APIFormat
@@ -64,7 +63,7 @@ class ApiFormatDefinition:
yield normalized
_DEFINITIONS: Dict[APIFormat, ApiFormatDefinition] = {
_DEFINITIONS: dict[APIFormat, ApiFormatDefinition] = {
APIFormat.CLAUDE: ApiFormatDefinition(
api_format=APIFormat.CLAUDE,
aliases=("claude", "anthropic", "claude_compatible"),
@@ -151,12 +150,12 @@ def get_api_format_definition(api_format: APIFormat) -> ApiFormatDefinition:
return API_FORMAT_DEFINITIONS[api_format]
def list_api_format_definitions() -> List[ApiFormatDefinition]:
def list_api_format_definitions() -> list[ApiFormatDefinition]:
"""返回所有定义的浅拷贝列表,供遍历使用。"""
return list(API_FORMAT_DEFINITIONS.values())
def build_alias_lookup() -> Dict[str, APIFormat]:
def build_alias_lookup() -> dict[str, APIFormat]:
"""
构建 alias -> APIFormat 的查找表。
每次调用都会返回新的 dict避免可变全局引发并发问题。
@@ -237,7 +236,7 @@ def get_protected_keys(api_format: APIFormat) -> frozenset[str]:
return frozenset({"authorization", "content-type"})
def get_data_format_id(api_format: Union[str, APIFormat]) -> str:
def get_data_format_id(api_format: str | APIFormat) -> str:
"""
获取格式的数据格式标识。
@@ -264,7 +263,7 @@ def get_data_format_id(api_format: Union[str, APIFormat]) -> str:
return api_format.value.lower()
def can_passthrough(client_format: Union[str, APIFormat], endpoint_format: Union[str, APIFormat]) -> bool:
def can_passthrough(client_format: str | APIFormat, endpoint_format: str | APIFormat) -> bool:
"""
判断两个格式之间是否可以透传(不需要数据转换)。
@@ -294,12 +293,12 @@ def can_passthrough(client_format: Union[str, APIFormat], endpoint_format: Union
@lru_cache(maxsize=1)
def _alias_lookup_cache() -> Dict[str, APIFormat]:
def _alias_lookup_cache() -> dict[str, APIFormat]:
"""缓存 alias -> APIFormat 查找表,减少重复构建。"""
return build_alias_lookup()
def resolve_api_format_alias(value: str) -> Optional[APIFormat]:
def resolve_api_format_alias(value: str) -> APIFormat | None:
"""根据别名查找 APIFormat找不到时返回 None。"""
if not value:
return None
@@ -310,9 +309,9 @@ def resolve_api_format_alias(value: str) -> Optional[APIFormat]:
def resolve_api_format(
value: Union[str, APIFormat, None],
default: Optional[APIFormat] = None,
) -> Optional[APIFormat]:
value: str | APIFormat | None,
default: APIFormat | None = None,
) -> APIFormat | None:
"""
将任意字符串/枚举值解析为 APIFormat。

View File

@@ -4,15 +4,14 @@ API 格式工具函数
提供格式判断、规范化等工具函数,供整个项目使用。
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Optional, Union
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from src.core.api_format.enums import APIFormat
def is_cli_format(format_id: Union[str, "APIFormat", None]) -> bool:
def is_cli_format(format_id: str | APIFormat | None) -> bool:
"""
判断是否为 CLI 透传格式
@@ -40,7 +39,7 @@ def is_cli_format(format_id: Union[str, "APIFormat", None]) -> bool:
return str(format_id).upper().endswith("_CLI")
def get_base_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]:
def get_base_format(format_id: str | APIFormat | None) -> str | None:
"""
获取基础格式(去除 _CLI 后缀)
@@ -66,7 +65,7 @@ def get_base_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]:
return format_str
def normalize_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]:
def normalize_format(format_id: str | APIFormat | None) -> str | None:
"""
规范化格式标识符
@@ -84,8 +83,8 @@ def normalize_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]:
def is_same_format(
format1: Union[str, "APIFormat", None],
format2: Union[str, "APIFormat", None],
format1: str | APIFormat | None,
format2: str | APIFormat | None,
) -> bool:
"""
判断两个格式是否相同
@@ -95,7 +94,7 @@ def is_same_format(
return normalize_format(format1) == normalize_format(format2)
def is_convertible_format(format_id: Union[str, "APIFormat", None]) -> bool:
def is_convertible_format(format_id: str | APIFormat | None) -> bool:
"""
判断是否为可转换格式

View File

@@ -8,7 +8,6 @@
"""
import asyncio
from typing import Set
from src.core.logger import logger
from sqlalchemy.orm import Session
@@ -23,7 +22,7 @@ class BatchCommitter:
interval_seconds: 批量提交间隔(秒)
"""
self.interval_seconds = interval_seconds
self._pending_sessions: Set[Session] = set()
self._pending_sessions: set[Session] = set()
self._lock = asyncio.Lock()
self._task = None

View File

@@ -3,8 +3,7 @@
"""
import json
from datetime import timedelta
from typing import Any, Optional
from typing import Any
from src.clients.redis_client import get_redis_client
from src.core.logger import logger
@@ -15,7 +14,7 @@ class CacheService:
"""缓存服务"""
@staticmethod
async def get(key: str) -> Optional[Any]:
async def get(key: str) -> Any | None:
"""
从缓存获取数据
@@ -170,7 +169,7 @@ class CacheService:
return False
@staticmethod
async def incr(key: str, ttl_seconds: Optional[int] = None) -> int:
async def incr(key: str, ttl_seconds: int | None = None) -> int:
"""
递增缓存值

View File

@@ -7,7 +7,7 @@
import threading
import time
from collections import OrderedDict
from typing import Any, Dict, Optional
from typing import Any
class SyncLRUCache:
@@ -26,7 +26,7 @@ class SyncLRUCache:
ttl: 过期时间(秒)
"""
self._cache: OrderedDict = OrderedDict()
self._expiry: Dict[Any, float] = {}
self._expiry: dict[Any, float] = {}
self.max_size = max_size
self.ttl = ttl
self._lock = threading.RLock()
@@ -57,7 +57,7 @@ class SyncLRUCache:
self._cache.move_to_end(key)
return self._cache[key]
def set(self, key: Any, value: Any, ttl: Optional[int] = None) -> None:
def set(self, key: Any, value: Any, ttl: int | None = None) -> None:
"""设置缓存值"""
with self._lock:
if ttl is None:
@@ -123,7 +123,7 @@ class SyncLRUCache:
k for k in self._cache.keys() if k not in self._expiry or now <= self._expiry[k]
]
def get_stats(self) -> Dict[str, Any]:
def get_stats(self) -> dict[str, Any]:
"""获取缓存统计信息"""
with self._lock:
return {

View File

@@ -8,6 +8,7 @@
- 使用 PBKDF2 派生密钥时会使用应用级 salt
"""
from __future__ import annotations
import base64
import hashlib
@@ -37,7 +38,7 @@ class CryptoService:
# 注意:更改此值会导致所有已加密数据无法解密
APP_SALT = hashlib.sha256(b"aether-v1").digest()[:16]
def __new__(cls) -> "CryptoService":
def __new__(cls) -> CryptoService:
if cls._instance is None:
cls._instance = super().__new__(cls)
cls._instance._initialize()

View File

@@ -2,10 +2,9 @@
错误消息处理工具函数
"""
from typing import Optional
def extract_error_message(error: Exception, status_code: Optional[int] = None) -> str:
def extract_error_message(error: Exception, status_code: int | None = None) -> str:
"""
从异常中提取错误消息,优先使用上游原始响应(用于链路追踪/调试)

View File

@@ -11,7 +11,7 @@ import asyncio
import re
import traceback
import uuid
from typing import Any, Dict, List, Optional
from typing import Any
import httpx
from fastapi import HTTPException, status
@@ -91,7 +91,7 @@ FIELD_NAME_TRANSLATIONS = {
}
def translate_pydantic_error(error: Dict[str, Any]) -> str:
def translate_pydantic_error(error: dict[str, Any]) -> str:
"""
将 Pydantic 验证错误翻译为中文
@@ -122,7 +122,7 @@ def translate_pydantic_error(error: Dict[str, Any]) -> str:
return translated_msg
def translate_pydantic_errors(errors: List[Dict[str, Any]]) -> str:
def translate_pydantic_errors(errors: list[dict[str, Any]]) -> str:
"""
翻译多个 Pydantic 验证错误
@@ -157,7 +157,7 @@ class ProxyException(HTTPException):
status_code: int,
error_type: str,
message: str,
details: Optional[Dict[str, Any]] = None,
details: dict[str, Any] | None = None,
):
self.error_type = error_type
self.message = message
@@ -171,8 +171,8 @@ class ProviderException(ProxyException):
def __init__(
self,
message: str,
provider_name: Optional[str] = None,
request_metadata: Optional[Any] = None,
provider_name: str | None = None,
request_metadata: Any | None = None,
**kwargs,
):
self.request_metadata = request_metadata # 保存元数据以便传递
@@ -192,10 +192,10 @@ class ProviderNotAvailableException(ProviderException):
def __init__(
self,
message: str,
provider_name: Optional[str] = None,
request_metadata: Optional[Any] = None,
upstream_status: Optional[int] = None,
upstream_response: Optional[str] = None,
provider_name: str | None = None,
request_metadata: Any | None = None,
upstream_status: int | None = None,
upstream_response: str | None = None,
):
super().__init__(
message=message,
@@ -209,7 +209,7 @@ class ProviderNotAvailableException(ProviderException):
class ProviderTimeoutException(ProviderException):
"""提供商请求超时"""
def __init__(self, provider_name: str, timeout: int, request_metadata: Optional[Any] = None):
def __init__(self, provider_name: str, timeout: int, request_metadata: Any | None = None):
super().__init__(
message=f"请求超时({timeout}秒)",
provider_name=provider_name,
@@ -221,7 +221,7 @@ class ProviderTimeoutException(ProviderException):
class ProviderAuthException(ProviderException):
"""提供商认证失败"""
def __init__(self, provider_name: str, request_metadata: Optional[Any] = None):
def __init__(self, provider_name: str, request_metadata: Any | None = None):
super().__init__(
message="上游服务认证失败",
provider_name=provider_name,
@@ -235,10 +235,10 @@ class ProviderRateLimitException(ProviderException):
def __init__(
self,
message: str,
provider_name: Optional[str] = None,
request_metadata: Optional[Any] = None,
response_headers: Optional[Dict[str, str]] = None, # 添加响应头
retry_after: Optional[int] = None, # 添加重试时间
provider_name: str | None = None,
request_metadata: Any | None = None,
response_headers: dict[str, str] | None = None, # 添加响应头
retry_after: int | None = None, # 添加重试时间
):
self.response_headers = response_headers or {} # 保存响应头
self.retry_after = retry_after # 保存重试时间
@@ -250,7 +250,7 @@ class ProviderRateLimitException(ProviderException):
class QuotaExceededException(ProxyException):
"""配额超限"""
def __init__(self, quota_type: str = "tokens", remaining: Optional[float] = None):
def __init__(self, quota_type: str = "tokens", remaining: float | None = None):
message = f"{quota_type}配额已用尽"
if remaining is not None:
message += f"(剩余: {remaining}"
@@ -278,7 +278,7 @@ class ConcurrencyLimitError(ProxyException):
"""并发限制异常"""
def __init__(
self, message: str, endpoint_id: Optional[str] = None, key_id: Optional[str] = None
self, message: str, endpoint_id: str | None = None, key_id: str | None = None
):
details = {}
if endpoint_id:
@@ -297,7 +297,7 @@ class ConcurrencyLimitError(ProxyException):
class ModelNotSupportedException(ProxyException):
"""模型不支持"""
def __init__(self, model: str, provider_name: Optional[str] = None):
def __init__(self, model: str, provider_name: str | None = None):
# 客户端消息不暴露提供商信息
message = f"模型 '{model}' 不受支持"
super().__init__(
@@ -311,7 +311,7 @@ class ModelNotSupportedException(ProxyException):
class StreamingNotSupportedException(ProxyException):
"""流式请求不支持"""
def __init__(self, model: str, provider_name: Optional[str] = None):
def __init__(self, model: str, provider_name: str | None = None):
# 客户端消息不暴露提供商信息
message = f"模型 '{model}' 不支持流式请求"
super().__init__(
@@ -325,7 +325,7 @@ class StreamingNotSupportedException(ProxyException):
class InvalidRequestException(ProxyException):
"""无效请求"""
def __init__(self, message: str, field: Optional[str] = None):
def __init__(self, message: str, field: str | None = None):
super().__init__(
status_code=status.HTTP_400_BAD_REQUEST,
error_type="invalid_request",
@@ -337,7 +337,7 @@ class InvalidRequestException(ProxyException):
class NotFoundException(ProxyException):
"""资源未找到"""
def __init__(self, message: str, resource_type: Optional[str] = None):
def __init__(self, message: str, resource_type: str | None = None):
super().__init__(
status_code=status.HTTP_404_NOT_FOUND,
error_type="not_found",
@@ -361,7 +361,7 @@ class ConfirmationRequiredException(ProxyException):
class ForbiddenException(ProxyException):
"""权限不足"""
def __init__(self, message: str, required_role: Optional[str] = None):
def __init__(self, message: str, required_role: str | None = None):
super().__init__(
status_code=status.HTTP_403_FORBIDDEN,
error_type="forbidden",
@@ -373,7 +373,7 @@ class ForbiddenException(ProxyException):
class DecryptionException(ProxyException):
"""解密失败异常 - 已知的配置问题,不需要打印堆栈"""
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
def __init__(self, message: str, details: dict[str, Any] | None = None):
super().__init__(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
error_type="decryption_error",
@@ -389,9 +389,9 @@ class JSONParseException(ProviderException):
self,
provider_name: str,
original_error: str,
response_content: Optional[str] = None,
content_type: Optional[str] = None,
request_metadata: Optional[Any] = None,
response_content: str | None = None,
content_type: str | None = None,
request_metadata: Any | None = None,
):
details = {
"original_error": original_error,
@@ -418,7 +418,7 @@ class EmptyStreamException(ProviderException):
self,
provider_name: str,
chunk_count: int = 0,
request_metadata: Optional[Any] = None,
request_metadata: Any | None = None,
):
super().__init__(
message="上游服务返回了空的流式响应",
@@ -438,10 +438,10 @@ class EmbeddedErrorException(ProviderException):
def __init__(
self,
provider_name: str,
error_code: Optional[int] = None,
error_message: Optional[str] = None,
error_status: Optional[str] = None,
request_metadata: Optional[Any] = None,
error_code: int | None = None,
error_message: str | None = None,
error_status: str | None = None,
request_metadata: Any | None = None,
):
# 客户端消息不暴露提供商信息
message = "上游服务返回了错误"
@@ -475,10 +475,10 @@ class ProviderCompatibilityException(ProviderException):
def __init__(
self,
message: str,
provider_name: Optional[str] = None,
provider_name: str | None = None,
status_code: int = 400,
upstream_error: Optional[str] = None,
request_metadata: Optional[Any] = None,
upstream_error: str | None = None,
request_metadata: Any | None = None,
):
self.upstream_error = upstream_error
super().__init__(
@@ -505,11 +505,11 @@ class UpstreamClientException(ProxyException):
def __init__(
self,
message: str,
provider_name: Optional[str] = None,
provider_name: str | None = None,
status_code: int = 400,
error_type: Optional[str] = None,
upstream_error: Optional[str] = None,
request_metadata: Optional[Any] = None,
error_type: str | None = None,
upstream_error: str | None = None,
request_metadata: Any | None = None,
):
self.upstream_error = upstream_error
self.request_metadata = request_metadata
@@ -535,8 +535,8 @@ class ThinkingSignatureException(UpstreamClientException):
def __init__(
self,
message: str,
provider_name: Optional[str] = None,
upstream_error: Optional[str] = None,
provider_name: str | None = None,
upstream_error: str | None = None,
request_metadata: Any = None,
):
super().__init__(
@@ -557,7 +557,7 @@ class ErrorResponse:
error_type: str,
message: str,
status_code: int = 500,
details: Optional[Dict[str, Any]] = None,
details: dict[str, Any] | None = None,
) -> JSONResponse:
"""创建标准错误响应"""
error_body = {"error": {"type": error_type, "message": message}}

View File

@@ -13,7 +13,6 @@ Key 能力系统
from dataclasses import dataclass, field
from enum import Enum
from typing import Dict, List, Optional, Tuple
class CapabilityMatchMode(Enum):
@@ -41,12 +40,12 @@ class CapabilityDefinition:
match_mode: CapabilityMatchMode
config_mode: CapabilityConfigMode
short_name: str = "" # 简短展示名称(用于列表等紧凑场景)
error_patterns: List[str] = field(default_factory=list) # 错误检测关键词组
error_patterns: list[str] = field(default_factory=list) # 错误检测关键词组
# ============ 能力注册表 ============
_capabilities: Dict[str, CapabilityDefinition] = {}
_capabilities: dict[str, CapabilityDefinition] = {}
def register_capability(
@@ -56,7 +55,7 @@ def register_capability(
match_mode: CapabilityMatchMode,
config_mode: CapabilityConfigMode,
short_name: str = "",
error_patterns: Optional[List[str]] = None,
error_patterns: list[str] | None = None,
) -> CapabilityDefinition:
"""注册能力"""
cap = CapabilityDefinition(
@@ -72,17 +71,17 @@ def register_capability(
return cap
def get_capability(name: str) -> Optional[CapabilityDefinition]:
def get_capability(name: str) -> CapabilityDefinition | None:
"""获取能力定义"""
return _capabilities.get(name)
def get_all_capabilities() -> List[CapabilityDefinition]:
def get_all_capabilities() -> list[CapabilityDefinition]:
"""获取所有能力定义"""
return list(_capabilities.values())
def get_user_configurable_capabilities() -> List[CapabilityDefinition]:
def get_user_configurable_capabilities() -> list[CapabilityDefinition]:
"""获取用户可配置的能力列表"""
return [c for c in _capabilities.values() if c.config_mode == CapabilityConfigMode.USER_CONFIGURABLE]
@@ -91,9 +90,9 @@ def get_user_configurable_capabilities() -> List[CapabilityDefinition]:
def check_capability_match(
key_capabilities: Optional[Dict[str, bool]],
requirements: Optional[Dict[str, bool]],
) -> Tuple[bool, Optional[str]]:
key_capabilities: dict[str, bool] | None,
requirements: dict[str, bool] | None,
) -> tuple[bool, str | None]:
"""
检查 Key 能力是否满足需求
@@ -157,7 +156,7 @@ def check_capability_match(
return True, None
def _match_error_patterns(error_msg: str, patterns: List[str]) -> bool:
def _match_error_patterns(error_msg: str, patterns: list[str]) -> bool:
"""检查错误信息是否匹配模式(所有关键词都要出现)"""
if not patterns:
return False
@@ -167,8 +166,8 @@ def _match_error_patterns(error_msg: str, patterns: List[str]) -> bool:
def detect_capability_upgrade_from_error(
error_msg: str,
current_requirements: Optional[Dict[str, bool]] = None,
) -> Optional[str]:
current_requirements: dict[str, bool] | None = None,
) -> str | None:
"""
从错误信息检测是否需要升级某能力
@@ -198,7 +197,7 @@ get_capability_definition = get_capability
class _CapabilityDefinitionsProxy:
"""CAPABILITY_DEFINITIONS 代理,提供字典式访问(兼容旧代码)"""
def get(self, name: str) -> Optional[CapabilityDefinition]:
def get(self, name: str) -> CapabilityDefinition | None:
return _capabilities.get(name)
def __getitem__(self, name: str) -> CapabilityDefinition:
@@ -210,10 +209,10 @@ class _CapabilityDefinitionsProxy:
def __contains__(self, name: str) -> bool:
return name in _capabilities
def values(self) -> List[CapabilityDefinition]:
def values(self) -> list[CapabilityDefinition]:
return list(_capabilities.values())
def items(self) -> List[Tuple[str, CapabilityDefinition]]:
def items(self) -> list[tuple[str, CapabilityDefinition]]:
return list(_capabilities.items())

View File

@@ -13,7 +13,6 @@ allowed_models 格式: ["claude-sonnet-4", "gpt-4o"]
import re
from functools import lru_cache
from typing import List, Optional, Tuple
import regex
@@ -26,10 +25,10 @@ MAX_MODEL_NAME_LENGTH = 200 # 与 MAX_MAPPING_LENGTH 保持一致
REGEX_MATCH_TIMEOUT_MS = 100 # 正则匹配超时(毫秒)
# 类型别名
AllowedModels = Optional[List[str]]
type AllowedModels = list[str] | None
def normalize_allowed_models(allowed_models: AllowedModels) -> Optional[set[str]]:
def normalize_allowed_models(allowed_models: AllowedModels) -> set[str] | None:
"""
将 allowed_models 规范化为模型名称集合
@@ -130,7 +129,7 @@ def get_allowed_models_preview(
return preview
def parse_allowed_models_to_list(allowed_models: AllowedModels) -> List[str]:
def parse_allowed_models_to_list(allowed_models: AllowedModels) -> list[str]:
"""
解析 allowed_models 为列表
@@ -146,7 +145,7 @@ def parse_allowed_models_to_list(allowed_models: AllowedModels) -> List[str]:
return list(allowed_models)
def validate_mapping_pattern(pattern: str) -> Tuple[bool, Optional[str]]:
def validate_mapping_pattern(pattern: str) -> tuple[bool, str | None]:
"""
验证映射模式是否安全
@@ -171,7 +170,7 @@ def validate_mapping_pattern(pattern: str) -> Tuple[bool, Optional[str]]:
return True, None
def validate_model_mappings(mappings: Optional[List[str]]) -> Tuple[bool, Optional[str]]:
def validate_model_mappings(mappings: list[str] | None) -> tuple[bool, str | None]:
"""
验证映射列表是否合法
@@ -196,8 +195,8 @@ def validate_model_mappings(mappings: Optional[List[str]]) -> Tuple[bool, Option
def validate_and_extract_model_mappings(
config: Optional[dict],
) -> Tuple[bool, Optional[str], Optional[List[str]]]:
config: dict | None,
) -> tuple[bool, str | None, list[str] | None]:
"""
从 config 中验证并提取 model_mappings
@@ -238,7 +237,7 @@ def validate_and_extract_model_mappings(
@lru_cache(maxsize=2000)
def _compile_pattern_cached(pattern: str) -> Optional[regex.Pattern]:
def _compile_pattern_cached(pattern: str) -> regex.Pattern | None:
"""
编译正则模式(带 LRU 缓存)
@@ -267,7 +266,7 @@ def clear_regex_cache() -> None:
def _match_with_timeout(
compiled_regex: regex.Pattern, text: str, timeout_ms: int = REGEX_MATCH_TIMEOUT_MS
) -> Optional[bool]:
) -> bool | None:
"""
带超时的正则匹配(使用 regex 库的原生超时支持)
@@ -342,9 +341,9 @@ def match_model_with_pattern(pattern: str, model_name: str) -> bool:
def check_model_allowed_with_mappings(
model_name: str,
allowed_models: AllowedModels,
model_mappings: Optional[List[str]] = None,
candidate_models: Optional[set[str]] = None,
) -> tuple[bool, Optional[str]]:
model_mappings: list[str] | None = None,
candidate_models: set[str] | None = None,
) -> tuple[bool, str | None]:
"""
检查模型是否被允许(支持映射通配符匹配)

View File

@@ -6,7 +6,10 @@
from dataclasses import dataclass, field
from enum import Enum
from typing import TYPE_CHECKING, Any, Awaitable, Callable, List, Optional, Tuple
from typing import TYPE_CHECKING, Any
from collections.abc import Callable
from collections.abc import Awaitable
if TYPE_CHECKING:
from fastapi import APIRouter
@@ -50,16 +53,16 @@ class ModuleMetadata:
# 可用性控制(部署级)
env_key: str # 环境变量名: LDAP_AVAILABLE
default_available: bool = False # 默认是否可用
required_packages: List[str] = field(default_factory=list) # 依赖的 Python 包
dependencies: List[str] = field(default_factory=list) # 依赖的其他模块
required_packages: list[str] = field(default_factory=list) # 依赖的 Python 包
dependencies: list[str] = field(default_factory=list) # 依赖的其他模块
# 路由配置 - 模块自定义前缀
api_prefix: Optional[str] = None # 如 "/api/admin/ldap"
api_prefix: str | None = None # 如 "/api/admin/ldap"
# 前端配置
admin_route: Optional[str] = None # 管理页面路由: "/admin/ldap"
admin_menu_icon: Optional[str] = None # 菜单图标
admin_menu_group: Optional[str] = None # 菜单分组: "system", "security"
admin_route: str | None = None # 管理页面路由: "/admin/ldap"
admin_menu_icon: str | None = None # 菜单图标
admin_menu_group: str | None = None # 菜单分组: "system", "security"
admin_menu_order: int = 100 # 菜单排序(越小越靠前)
@@ -74,19 +77,19 @@ class ModuleDefinition:
metadata: ModuleMetadata
# 工厂函数 - 内部再 import 重依赖
router_factory: Optional[Callable[[], "APIRouter"]] = None
service_factory: Optional[Callable[[], Any]] = None
router_factory: Callable[[], APIRouter] | None = None
service_factory: Callable[[], Any] | None = None
# 生命周期钩子
on_startup: Optional[Callable[[], Awaitable[None]]] = None
on_shutdown: Optional[Callable[[], Awaitable[None]]] = None
health_check: Optional[Callable[[], Awaitable[ModuleHealth]]] = None
on_startup: Callable[[], Awaitable[None]] | None = None
on_shutdown: Callable[[], Awaitable[None]] | None = None
health_check: Callable[[], Awaitable[ModuleHealth]] | None = None
# 自定义依赖检测(可选,用于检测 ldap3 等库是否安装)
check_dependencies: Optional[Callable[[], bool]] = None
check_dependencies: Callable[[], bool] | None = None
# 配置验证(可选,启用模块时调用,返回 (success, error_message)
validate_config: Optional[Callable[["Session"], Tuple[bool, str]]] = None
validate_config: Callable[[Session], tuple[bool, str]] | None = None
@dataclass
@@ -102,7 +105,7 @@ class ModuleStatus:
enabled: bool # 运行级启用(数据库配置)
active: bool # 最终激活状态 (available && enabled && dependencies_ok)
config_validated: bool # 配置验证通过(只有验证通过才允许启用)
config_error: Optional[str] # 配置验证失败的错误信息
config_error: str | None # 配置验证失败的错误信息
# 显示信息
display_name: str
@@ -110,9 +113,9 @@ class ModuleStatus:
category: ModuleCategory
# 前端配置
admin_route: Optional[str]
admin_menu_icon: Optional[str]
admin_menu_group: Optional[str]
admin_route: str | None
admin_menu_icon: str | None
admin_menu_group: str | None
admin_menu_order: int
# 健康状态

View File

@@ -4,9 +4,11 @@
负责模块的注册、状态管理和生命周期控制
"""
from __future__ import annotations
import importlib.util
import os
from typing import TYPE_CHECKING, Dict, List, Optional, Set
from typing import TYPE_CHECKING
from src.core.logger import logger
from src.core.modules.base import (
@@ -31,14 +33,14 @@ class ModuleRegistry:
- 提供模块状态查询
"""
_instance: Optional["ModuleRegistry"] = None
_instance: ModuleRegistry | None = None
def __init__(self):
self._modules: Dict[str, ModuleDefinition] = {}
self._initialized: Set[str] = set()
self._modules: dict[str, ModuleDefinition] = {}
self._initialized: set[str] = set()
@classmethod
def get_instance(cls) -> "ModuleRegistry":
def get_instance(cls) -> ModuleRegistry:
"""获取单例实例"""
if cls._instance is None:
cls._instance = cls()
@@ -63,11 +65,11 @@ class ModuleRegistry:
self._modules[name] = module
logger.debug(f"Module [{name}] registered")
def get_module(self, name: str) -> Optional[ModuleDefinition]:
def get_module(self, name: str) -> ModuleDefinition | None:
"""获取模块定义"""
return self._modules.get(name)
def get_all_modules(self) -> List[ModuleDefinition]:
def get_all_modules(self) -> list[ModuleDefinition]:
"""获取所有已注册模块"""
return list(self._modules.values())
@@ -115,13 +117,13 @@ class ModuleRegistry:
return True
def get_available_modules(self) -> List[ModuleDefinition]:
def get_available_modules(self) -> list[ModuleDefinition]:
"""获取所有部署可用的模块"""
return [m for m in self._modules.values() if self.is_available(m.metadata.name)]
# ========== 启用状态检查(运行级)==========
def is_enabled(self, name: str, db: "Session") -> bool:
def is_enabled(self, name: str, db: Session) -> bool:
"""
检查模块是否运行启用(数据库配置)
@@ -135,7 +137,7 @@ class ModuleRegistry:
value = SystemConfigService.get_config(db, config_key, default=False)
return bool(value)
def set_enabled(self, name: str, enabled: bool, db: "Session") -> None:
def set_enabled(self, name: str, enabled: bool, db: Session) -> None:
"""
设置模块启用状态
@@ -156,7 +158,7 @@ class ModuleRegistry:
# ========== 激活状态检查 ==========
def is_active(self, name: str, db: "Session") -> bool:
def is_active(self, name: str, db: Session) -> bool:
"""
检查模块是否最终激活
@@ -177,7 +179,7 @@ class ModuleRegistry:
# ========== 配置验证 ==========
def validate_config(self, name: str, db: "Session") -> tuple[bool, str]:
def validate_config(self, name: str, db: Session) -> tuple[bool, str]:
"""
验证模块配置是否有效
@@ -206,8 +208,8 @@ class ModuleRegistry:
# ========== 状态查询 ==========
def get_module_status(
self, name: str, db: "Session", health: Optional[ModuleHealth] = None
) -> Optional[ModuleStatus]:
self, name: str, db: Session, health: ModuleHealth | None = None
) -> ModuleStatus | None:
"""
获取单个模块状态
@@ -225,7 +227,7 @@ class ModuleRegistry:
# 获取配置验证状态
config_validated = False
config_error: Optional[str] = None
config_error: str | None = None
if available:
config_validated, config_error = self.validate_config(name, db)
if config_validated:
@@ -279,8 +281,8 @@ class ModuleRegistry:
return ModuleHealth.UNHEALTHY
async def get_module_status_async(
self, name: str, db: "Session"
) -> Optional[ModuleStatus]:
self, name: str, db: Session
) -> ModuleStatus | None:
"""异步获取模块状态(包含健康检查)"""
if name not in self._modules:
return None
@@ -288,7 +290,7 @@ class ModuleRegistry:
health = await self.check_health(name) if self.is_available(name) else ModuleHealth.UNKNOWN
return self.get_module_status(name, db, health=health)
async def get_all_status_async(self, db: "Session") -> Dict[str, ModuleStatus]:
async def get_all_status_async(self, db: Session) -> dict[str, ModuleStatus]:
"""异步获取所有模块状态(包含健康检查)"""
result = {}
for name in self._modules:
@@ -297,7 +299,7 @@ class ModuleRegistry:
result[name] = status
return result
def get_all_status(self, db: "Session") -> Dict[str, ModuleStatus]:
def get_all_status(self, db: Session) -> dict[str, ModuleStatus]:
"""获取所有模块状态(同步版本,不含健康检查)"""
result = {}
for name in self._modules:
@@ -306,7 +308,7 @@ class ModuleRegistry:
result[name] = status
return result
def get_available_status(self, db: "Session") -> Dict[str, ModuleStatus]:
def get_available_status(self, db: Session) -> dict[str, ModuleStatus]:
"""获取所有可用模块的状态"""
result = {}
for name, module in self._modules.items():
@@ -316,7 +318,7 @@ class ModuleRegistry:
result[name] = status
return result
def get_auth_modules_status(self, db: "Session") -> List[ModuleStatus]:
def get_auth_modules_status(self, db: Session) -> list[ModuleStatus]:
"""获取认证模块状态(供登录页使用)"""
result = []
for name, module in self._modules.items():

View File

@@ -2,7 +2,7 @@
优化工具类 - 包含Token计数和响应头管理
"""
from typing import Any, Dict, Optional
from typing import Any
import tiktoken

View File

@@ -5,8 +5,7 @@
import time
from collections import defaultdict
from datetime import datetime, timedelta
from typing import Any, Dict, Optional
from typing import Any
class ProviderHealthTracker:
@@ -26,11 +25,11 @@ class ProviderHealthTracker:
self.recovery_time = recovery_time
# 存储每个提供商的失败记录
self.failures: Dict[str, list] = defaultdict(list)
self.failures: dict[str, list] = defaultdict(list)
# 存储每个提供商的成功记录
self.successes: Dict[str, list] = defaultdict(list)
self.successes: dict[str, list] = defaultdict(list)
# 存储优先级调整
self.priority_adjustments: Dict[str, int] = {}
self.priority_adjustments: dict[str, int] = {}
def record_success(self, provider_name: str) -> None:
"""记录成功的请求"""
@@ -71,7 +70,7 @@ class ProviderHealthTracker:
"""
return self.priority_adjustments.get(provider_name, 0)
def get_health_status(self, provider_name: str) -> Dict:
def get_health_status(self, provider_name: str) -> dict:
"""
获取提供商的健康状态
"""
@@ -146,7 +145,7 @@ class SimpleProviderSelector:
def __init__(self, health_tracker: ProviderHealthTracker):
self.health_tracker = health_tracker
def select_provider(self, providers: list, specified_provider: Optional[str] = None) -> Any:
def select_provider(self, providers: list, specified_provider: str | None = None) -> Any:
"""
选择提供商

View File

@@ -10,9 +10,11 @@ import time
import traceback
import uuid
from contextlib import asynccontextmanager
from datetime import datetime, timedelta, timezone
from datetime import datetime, timezone
from enum import Enum
from typing import Any, Callable, Dict, List, Optional, Type, Union
from typing import Any
from collections.abc import Callable
from ..core.exceptions import ProxyException
from src.core.logger import logger
@@ -43,7 +45,7 @@ class ErrorPattern:
def __init__(
self,
error_types: List[Type[Exception]],
error_types: list[type[Exception]],
severity: ErrorSeverity,
recovery_strategy: RecoveryStrategy,
user_message: str,
@@ -113,10 +115,10 @@ class ResilienceManager:
"""系统韧性管理器"""
def __init__(self):
self.error_patterns: List[ErrorPattern] = []
self.circuit_breakers: Dict[str, CircuitBreaker] = {}
self.error_stats: Dict[str, int] = {}
self.last_errors: List[Dict[str, Any]] = []
self.error_patterns: list[ErrorPattern] = []
self.circuit_breakers: dict[str, CircuitBreaker] = {}
self.error_stats: dict[str, int] = {}
self.last_errors: list[dict[str, Any]] = []
self._setup_default_patterns()
def _setup_default_patterns(self):
@@ -196,8 +198,8 @@ class ResilienceManager:
return self.circuit_breakers[key]
def handle_error(
self, error: Exception, context: Dict[str, Any] = None, operation: str = "unknown"
) -> Dict[str, Any]:
self, error: Exception, context: dict[str, Any] = None, operation: str = "unknown"
) -> dict[str, Any]:
"""处理错误并返回处理结果"""
error_id = str(uuid.uuid4())[:8]
@@ -250,14 +252,14 @@ class ResilienceManager:
"pattern": None,
}
def _find_matching_pattern(self, error: Exception) -> Optional[ErrorPattern]:
def _find_matching_pattern(self, error: Exception) -> ErrorPattern | None:
"""查找匹配的错误处理模式"""
for pattern in self.error_patterns:
if any(isinstance(error, error_type) for error_type in pattern.error_types):
return pattern
return None
def get_error_stats(self) -> Dict[str, Any]:
def get_error_stats(self) -> dict[str, Any]:
"""获取错误统计"""
return {
"total_errors": sum(self.error_stats.values()),
@@ -279,7 +281,7 @@ def resilient_operation(
max_retries: int = None,
retry_delay: float = None,
circuit_breaker_key: str = None,
context: Dict[str, Any] = None,
context: dict[str, Any] = None,
):
"""
韧性操作装饰器
@@ -354,7 +356,7 @@ def resilient_operation(
@asynccontextmanager
async def safe_operation(operation_name: str, context: Dict[str, Any] = None):
async def safe_operation(operation_name: str, context: dict[str, Any] = None):
"""
安全操作上下文管理器
自动处理异常并提供用户友好的错误信息

View File

@@ -4,7 +4,6 @@
"""
import re
from typing import List, Optional
class PasswordValidator:
@@ -14,7 +13,7 @@ class PasswordValidator:
MAX_LENGTH = 128
@classmethod
def validate(cls, password: str) -> tuple[bool, Optional[str]]:
def validate(cls, password: str) -> tuple[bool, str | None]:
"""
验证密码复杂度
@@ -109,7 +108,7 @@ class EmailValidator:
EMAIL_REGEX = re.compile(r"^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$")
@classmethod
def validate(cls, email: str) -> tuple[bool, Optional[str]]:
def validate(cls, email: str) -> tuple[bool, str | None]:
"""
验证邮箱格式
@@ -139,7 +138,7 @@ class UsernameValidator:
USERNAME_REGEX = re.compile(r"^[a-zA-Z0-9_.\-]+$")
@classmethod
def validate(cls, username: str) -> tuple[bool, Optional[str]]:
def validate(cls, username: str) -> tuple[bool, str | None]:
"""
验证用户名