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:
"""
判断是否为可转换格式