mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
Merge branch 'fix/python314-upgrade'
# Conflicts: # src/api/handlers/base/base_handler.py # src/api/handlers/base/request_builder.py # src/models/endpoint_models.py # src/services/orchestration/candidate_resolver.py # src/services/orchestration/fallback_orchestrator.py
This commit is contained in:
@@ -12,8 +12,9 @@
|
||||
|
||||
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 +27,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]:
|
||||
"""
|
||||
检查端点是否兼容客户端格式
|
||||
|
||||
|
||||
@@ -4,7 +4,6 @@
|
||||
用于严格模式转换失败时抛出,让编排器可以尝试下一个候选。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class FormatConversionError(Exception):
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
# 工具输出可能是纯文本,也可能是结构化 JSON(Gemini 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__ = [
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -6,7 +6,6 @@ Normalizers
|
||||
本目录在 Phase 1 仅创建结构;具体实现将在 Phase 2+ 补齐。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
__all__: list[str] = []
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -11,11 +11,9 @@ Gemini (GenerateContent / streamGenerateContent) Normalizer
|
||||
- 响应/流式通常为 camelCase(candidates/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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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 = {}
|
||||
|
||||
@@ -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):
|
||||
|
||||
# 普通 message(TextBlock)
|
||||
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)
|
||||
|
||||
@@ -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 []
|
||||
|
||||
@@ -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__ = [
|
||||
|
||||
@@ -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, {})
|
||||
|
||||
@@ -6,7 +6,8 @@ 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 +17,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 +65,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 +108,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 +136,7 @@ def detect_format_and_key_from_starlette(
|
||||
|
||||
def detect_format_from_response(
|
||||
response_data: dict,
|
||||
) -> Optional[APIFormat]:
|
||||
) -> APIFormat | None:
|
||||
"""
|
||||
从响应内容检测 API 格式
|
||||
|
||||
|
||||
@@ -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 提取额外请求头
|
||||
|
||||
|
||||
@@ -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。
|
||||
|
||||
|
||||
@@ -6,13 +6,13 @@ 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 +40,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 +66,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 +84,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 +95,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:
|
||||
"""
|
||||
判断是否为可转换格式
|
||||
|
||||
|
||||
@@ -8,7 +8,6 @@
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import Set
|
||||
|
||||
from src.core.logger import logger
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -23,7 +22,7 @@ class BatchCommitter:
|
||||
interval_seconds: 批量提交间隔(秒)
|
||||
"""
|
||||
self.interval_seconds = interval_seconds
|
||||
self._pending_sessions: Set[Session] = set()
|
||||
self._pending_sessions: set[Session] = set()
|
||||
self._lock = asyncio.Lock()
|
||||
self._task = None
|
||||
|
||||
|
||||
@@ -3,8 +3,7 @@
|
||||
"""
|
||||
|
||||
import json
|
||||
from datetime import timedelta
|
||||
from typing import Any, Optional
|
||||
from typing import Any
|
||||
|
||||
from src.clients.redis_client import get_redis_client
|
||||
from src.core.logger import logger
|
||||
@@ -15,7 +14,7 @@ class CacheService:
|
||||
"""缓存服务"""
|
||||
|
||||
@staticmethod
|
||||
async def get(key: str) -> Optional[Any]:
|
||||
async def get(key: str) -> Any | None:
|
||||
"""
|
||||
从缓存获取数据
|
||||
|
||||
@@ -170,7 +169,7 @@ class CacheService:
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
async def incr(key: str, ttl_seconds: Optional[int] = None) -> int:
|
||||
async def incr(key: str, ttl_seconds: int | None = None) -> int:
|
||||
"""
|
||||
递增缓存值
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
import threading
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
|
||||
class SyncLRUCache:
|
||||
@@ -26,7 +26,7 @@ class SyncLRUCache:
|
||||
ttl: 过期时间(秒)
|
||||
"""
|
||||
self._cache: OrderedDict = OrderedDict()
|
||||
self._expiry: Dict[Any, float] = {}
|
||||
self._expiry: dict[Any, float] = {}
|
||||
self.max_size = max_size
|
||||
self.ttl = ttl
|
||||
self._lock = threading.RLock()
|
||||
@@ -57,7 +57,7 @@ class SyncLRUCache:
|
||||
self._cache.move_to_end(key)
|
||||
return self._cache[key]
|
||||
|
||||
def set(self, key: Any, value: Any, ttl: Optional[int] = None) -> None:
|
||||
def set(self, key: Any, value: Any, ttl: int | None = None) -> None:
|
||||
"""设置缓存值"""
|
||||
with self._lock:
|
||||
if ttl is None:
|
||||
@@ -123,7 +123,7 @@ class SyncLRUCache:
|
||||
k for k in self._cache.keys() if k not in self._expiry or now <= self._expiry[k]
|
||||
]
|
||||
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
def get_stats(self) -> dict[str, Any]:
|
||||
"""获取缓存统计信息"""
|
||||
with self._lock:
|
||||
return {
|
||||
|
||||
@@ -8,6 +8,7 @@
|
||||
- 使用 PBKDF2 派生密钥时会使用应用级 salt
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
import base64
|
||||
import hashlib
|
||||
|
||||
@@ -37,7 +38,7 @@ class CryptoService:
|
||||
# 注意:更改此值会导致所有已加密数据无法解密
|
||||
APP_SALT = hashlib.sha256(b"aether-v1").digest()[:16]
|
||||
|
||||
def __new__(cls) -> "CryptoService":
|
||||
def __new__(cls) -> CryptoService:
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
cls._instance._initialize()
|
||||
|
||||
@@ -2,10 +2,9 @@
|
||||
错误消息处理工具函数
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
|
||||
def extract_error_message(error: Exception, status_code: Optional[int] = None) -> str:
|
||||
def extract_error_message(error: Exception, status_code: int | None = None) -> str:
|
||||
"""
|
||||
从异常中提取错误消息,优先使用上游原始响应(用于链路追踪/调试)
|
||||
|
||||
|
||||
@@ -11,7 +11,7 @@ import asyncio
|
||||
import re
|
||||
import traceback
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, status
|
||||
@@ -91,7 +91,7 @@ FIELD_NAME_TRANSLATIONS = {
|
||||
}
|
||||
|
||||
|
||||
def translate_pydantic_error(error: Dict[str, Any]) -> str:
|
||||
def translate_pydantic_error(error: dict[str, Any]) -> str:
|
||||
"""
|
||||
将 Pydantic 验证错误翻译为中文
|
||||
|
||||
@@ -122,7 +122,7 @@ def translate_pydantic_error(error: Dict[str, Any]) -> str:
|
||||
return translated_msg
|
||||
|
||||
|
||||
def translate_pydantic_errors(errors: List[Dict[str, Any]]) -> str:
|
||||
def translate_pydantic_errors(errors: list[dict[str, Any]]) -> str:
|
||||
"""
|
||||
翻译多个 Pydantic 验证错误
|
||||
|
||||
@@ -157,7 +157,7 @@ class ProxyException(HTTPException):
|
||||
status_code: int,
|
||||
error_type: str,
|
||||
message: str,
|
||||
details: Optional[Dict[str, Any]] = None,
|
||||
details: dict[str, Any] | None = None,
|
||||
):
|
||||
self.error_type = error_type
|
||||
self.message = message
|
||||
@@ -171,8 +171,8 @@ class ProviderException(ProxyException):
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
provider_name: Optional[str] = None,
|
||||
request_metadata: Optional[Any] = None,
|
||||
provider_name: str | None = None,
|
||||
request_metadata: Any | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
self.request_metadata = request_metadata # 保存元数据以便传递
|
||||
@@ -192,10 +192,10 @@ class ProviderNotAvailableException(ProviderException):
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
provider_name: Optional[str] = None,
|
||||
request_metadata: Optional[Any] = None,
|
||||
upstream_status: Optional[int] = None,
|
||||
upstream_response: Optional[str] = None,
|
||||
provider_name: str | None = None,
|
||||
request_metadata: Any | None = None,
|
||||
upstream_status: int | None = None,
|
||||
upstream_response: str | None = None,
|
||||
):
|
||||
super().__init__(
|
||||
message=message,
|
||||
@@ -209,7 +209,7 @@ class ProviderNotAvailableException(ProviderException):
|
||||
class ProviderTimeoutException(ProviderException):
|
||||
"""提供商请求超时"""
|
||||
|
||||
def __init__(self, provider_name: str, timeout: int, request_metadata: Optional[Any] = None):
|
||||
def __init__(self, provider_name: str, timeout: int, request_metadata: Any | None = None):
|
||||
super().__init__(
|
||||
message=f"请求超时({timeout}秒)",
|
||||
provider_name=provider_name,
|
||||
@@ -221,7 +221,7 @@ class ProviderTimeoutException(ProviderException):
|
||||
class ProviderAuthException(ProviderException):
|
||||
"""提供商认证失败"""
|
||||
|
||||
def __init__(self, provider_name: str, request_metadata: Optional[Any] = None):
|
||||
def __init__(self, provider_name: str, request_metadata: Any | None = None):
|
||||
super().__init__(
|
||||
message="上游服务认证失败",
|
||||
provider_name=provider_name,
|
||||
@@ -235,10 +235,10 @@ class ProviderRateLimitException(ProviderException):
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
provider_name: Optional[str] = None,
|
||||
request_metadata: Optional[Any] = None,
|
||||
response_headers: Optional[Dict[str, str]] = None, # 添加响应头
|
||||
retry_after: Optional[int] = None, # 添加重试时间
|
||||
provider_name: str | None = None,
|
||||
request_metadata: Any | None = None,
|
||||
response_headers: dict[str, str] | None = None, # 添加响应头
|
||||
retry_after: int | None = None, # 添加重试时间
|
||||
):
|
||||
self.response_headers = response_headers or {} # 保存响应头
|
||||
self.retry_after = retry_after # 保存重试时间
|
||||
@@ -250,7 +250,7 @@ class ProviderRateLimitException(ProviderException):
|
||||
class QuotaExceededException(ProxyException):
|
||||
"""配额超限"""
|
||||
|
||||
def __init__(self, quota_type: str = "tokens", remaining: Optional[float] = None):
|
||||
def __init__(self, quota_type: str = "tokens", remaining: float | None = None):
|
||||
message = f"{quota_type}配额已用尽"
|
||||
if remaining is not None:
|
||||
message += f"(剩余: {remaining})"
|
||||
@@ -278,7 +278,7 @@ class ConcurrencyLimitError(ProxyException):
|
||||
"""并发限制异常"""
|
||||
|
||||
def __init__(
|
||||
self, message: str, endpoint_id: Optional[str] = None, key_id: Optional[str] = None
|
||||
self, message: str, endpoint_id: str | None = None, key_id: str | None = None
|
||||
):
|
||||
details = {}
|
||||
if endpoint_id:
|
||||
@@ -297,7 +297,7 @@ class ConcurrencyLimitError(ProxyException):
|
||||
class ModelNotSupportedException(ProxyException):
|
||||
"""模型不支持"""
|
||||
|
||||
def __init__(self, model: str, provider_name: Optional[str] = None):
|
||||
def __init__(self, model: str, provider_name: str | None = None):
|
||||
# 客户端消息不暴露提供商信息
|
||||
message = f"模型 '{model}' 不受支持"
|
||||
super().__init__(
|
||||
@@ -311,7 +311,7 @@ class ModelNotSupportedException(ProxyException):
|
||||
class StreamingNotSupportedException(ProxyException):
|
||||
"""流式请求不支持"""
|
||||
|
||||
def __init__(self, model: str, provider_name: Optional[str] = None):
|
||||
def __init__(self, model: str, provider_name: str | None = None):
|
||||
# 客户端消息不暴露提供商信息
|
||||
message = f"模型 '{model}' 不支持流式请求"
|
||||
super().__init__(
|
||||
@@ -325,7 +325,7 @@ class StreamingNotSupportedException(ProxyException):
|
||||
class InvalidRequestException(ProxyException):
|
||||
"""无效请求"""
|
||||
|
||||
def __init__(self, message: str, field: Optional[str] = None):
|
||||
def __init__(self, message: str, field: str | None = None):
|
||||
super().__init__(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
error_type="invalid_request",
|
||||
@@ -337,7 +337,7 @@ class InvalidRequestException(ProxyException):
|
||||
class NotFoundException(ProxyException):
|
||||
"""资源未找到"""
|
||||
|
||||
def __init__(self, message: str, resource_type: Optional[str] = None):
|
||||
def __init__(self, message: str, resource_type: str | None = None):
|
||||
super().__init__(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
error_type="not_found",
|
||||
@@ -361,7 +361,7 @@ class ConfirmationRequiredException(ProxyException):
|
||||
class ForbiddenException(ProxyException):
|
||||
"""权限不足"""
|
||||
|
||||
def __init__(self, message: str, required_role: Optional[str] = None):
|
||||
def __init__(self, message: str, required_role: str | None = None):
|
||||
super().__init__(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
error_type="forbidden",
|
||||
@@ -373,7 +373,7 @@ class ForbiddenException(ProxyException):
|
||||
class DecryptionException(ProxyException):
|
||||
"""解密失败异常 - 已知的配置问题,不需要打印堆栈"""
|
||||
|
||||
def __init__(self, message: str, details: Optional[Dict[str, Any]] = None):
|
||||
def __init__(self, message: str, details: dict[str, Any] | None = None):
|
||||
super().__init__(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
error_type="decryption_error",
|
||||
@@ -389,9 +389,9 @@ class JSONParseException(ProviderException):
|
||||
self,
|
||||
provider_name: str,
|
||||
original_error: str,
|
||||
response_content: Optional[str] = None,
|
||||
content_type: Optional[str] = None,
|
||||
request_metadata: Optional[Any] = None,
|
||||
response_content: str | None = None,
|
||||
content_type: str | None = None,
|
||||
request_metadata: Any | None = None,
|
||||
):
|
||||
details = {
|
||||
"original_error": original_error,
|
||||
@@ -418,7 +418,7 @@ class EmptyStreamException(ProviderException):
|
||||
self,
|
||||
provider_name: str,
|
||||
chunk_count: int = 0,
|
||||
request_metadata: Optional[Any] = None,
|
||||
request_metadata: Any | None = None,
|
||||
):
|
||||
super().__init__(
|
||||
message="上游服务返回了空的流式响应",
|
||||
@@ -438,10 +438,10 @@ class EmbeddedErrorException(ProviderException):
|
||||
def __init__(
|
||||
self,
|
||||
provider_name: str,
|
||||
error_code: Optional[int] = None,
|
||||
error_message: Optional[str] = None,
|
||||
error_status: Optional[str] = None,
|
||||
request_metadata: Optional[Any] = None,
|
||||
error_code: int | None = None,
|
||||
error_message: str | None = None,
|
||||
error_status: str | None = None,
|
||||
request_metadata: Any | None = None,
|
||||
):
|
||||
# 客户端消息不暴露提供商信息
|
||||
message = "上游服务返回了错误"
|
||||
@@ -475,10 +475,10 @@ class ProviderCompatibilityException(ProviderException):
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
provider_name: Optional[str] = None,
|
||||
provider_name: str | None = None,
|
||||
status_code: int = 400,
|
||||
upstream_error: Optional[str] = None,
|
||||
request_metadata: Optional[Any] = None,
|
||||
upstream_error: str | None = None,
|
||||
request_metadata: Any | None = None,
|
||||
):
|
||||
self.upstream_error = upstream_error
|
||||
super().__init__(
|
||||
@@ -505,11 +505,11 @@ class UpstreamClientException(ProxyException):
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
provider_name: Optional[str] = None,
|
||||
provider_name: str | None = None,
|
||||
status_code: int = 400,
|
||||
error_type: Optional[str] = None,
|
||||
upstream_error: Optional[str] = None,
|
||||
request_metadata: Optional[Any] = None,
|
||||
error_type: str | None = None,
|
||||
upstream_error: str | None = None,
|
||||
request_metadata: Any | None = None,
|
||||
):
|
||||
self.upstream_error = upstream_error
|
||||
self.request_metadata = request_metadata
|
||||
@@ -535,8 +535,8 @@ class ThinkingSignatureException(UpstreamClientException):
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
provider_name: Optional[str] = None,
|
||||
upstream_error: Optional[str] = None,
|
||||
provider_name: str | None = None,
|
||||
upstream_error: str | None = None,
|
||||
request_metadata: Any = None,
|
||||
):
|
||||
super().__init__(
|
||||
@@ -557,7 +557,7 @@ class ErrorResponse:
|
||||
error_type: str,
|
||||
message: str,
|
||||
status_code: int = 500,
|
||||
details: Optional[Dict[str, Any]] = None,
|
||||
details: dict[str, Any] | None = None,
|
||||
) -> JSONResponse:
|
||||
"""创建标准错误响应"""
|
||||
error_body = {"error": {"type": error_type, "message": message}}
|
||||
|
||||
@@ -13,7 +13,6 @@ Key 能力系统
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
|
||||
|
||||
class CapabilityMatchMode(Enum):
|
||||
@@ -41,12 +40,12 @@ class CapabilityDefinition:
|
||||
match_mode: CapabilityMatchMode
|
||||
config_mode: CapabilityConfigMode
|
||||
short_name: str = "" # 简短展示名称(用于列表等紧凑场景)
|
||||
error_patterns: List[str] = field(default_factory=list) # 错误检测关键词组
|
||||
error_patterns: list[str] = field(default_factory=list) # 错误检测关键词组
|
||||
|
||||
|
||||
# ============ 能力注册表 ============
|
||||
|
||||
_capabilities: Dict[str, CapabilityDefinition] = {}
|
||||
_capabilities: dict[str, CapabilityDefinition] = {}
|
||||
|
||||
|
||||
def register_capability(
|
||||
@@ -56,7 +55,7 @@ def register_capability(
|
||||
match_mode: CapabilityMatchMode,
|
||||
config_mode: CapabilityConfigMode,
|
||||
short_name: str = "",
|
||||
error_patterns: Optional[List[str]] = None,
|
||||
error_patterns: list[str] | None = None,
|
||||
) -> CapabilityDefinition:
|
||||
"""注册能力"""
|
||||
cap = CapabilityDefinition(
|
||||
@@ -72,17 +71,17 @@ def register_capability(
|
||||
return cap
|
||||
|
||||
|
||||
def get_capability(name: str) -> Optional[CapabilityDefinition]:
|
||||
def get_capability(name: str) -> CapabilityDefinition | None:
|
||||
"""获取能力定义"""
|
||||
return _capabilities.get(name)
|
||||
|
||||
|
||||
def get_all_capabilities() -> List[CapabilityDefinition]:
|
||||
def get_all_capabilities() -> list[CapabilityDefinition]:
|
||||
"""获取所有能力定义"""
|
||||
return list(_capabilities.values())
|
||||
|
||||
|
||||
def get_user_configurable_capabilities() -> List[CapabilityDefinition]:
|
||||
def get_user_configurable_capabilities() -> list[CapabilityDefinition]:
|
||||
"""获取用户可配置的能力列表"""
|
||||
return [c for c in _capabilities.values() if c.config_mode == CapabilityConfigMode.USER_CONFIGURABLE]
|
||||
|
||||
@@ -91,9 +90,9 @@ def get_user_configurable_capabilities() -> List[CapabilityDefinition]:
|
||||
|
||||
|
||||
def check_capability_match(
|
||||
key_capabilities: Optional[Dict[str, bool]],
|
||||
requirements: Optional[Dict[str, bool]],
|
||||
) -> Tuple[bool, Optional[str]]:
|
||||
key_capabilities: dict[str, bool] | None,
|
||||
requirements: dict[str, bool] | None,
|
||||
) -> tuple[bool, str | None]:
|
||||
"""
|
||||
检查 Key 能力是否满足需求
|
||||
|
||||
@@ -157,7 +156,7 @@ def check_capability_match(
|
||||
return True, None
|
||||
|
||||
|
||||
def _match_error_patterns(error_msg: str, patterns: List[str]) -> bool:
|
||||
def _match_error_patterns(error_msg: str, patterns: list[str]) -> bool:
|
||||
"""检查错误信息是否匹配模式(所有关键词都要出现)"""
|
||||
if not patterns:
|
||||
return False
|
||||
@@ -167,8 +166,8 @@ def _match_error_patterns(error_msg: str, patterns: List[str]) -> bool:
|
||||
|
||||
def detect_capability_upgrade_from_error(
|
||||
error_msg: str,
|
||||
current_requirements: Optional[Dict[str, bool]] = None,
|
||||
) -> Optional[str]:
|
||||
current_requirements: dict[str, bool] | None = None,
|
||||
) -> str | None:
|
||||
"""
|
||||
从错误信息检测是否需要升级某能力
|
||||
|
||||
@@ -198,7 +197,7 @@ get_capability_definition = get_capability
|
||||
class _CapabilityDefinitionsProxy:
|
||||
"""CAPABILITY_DEFINITIONS 代理,提供字典式访问(兼容旧代码)"""
|
||||
|
||||
def get(self, name: str) -> Optional[CapabilityDefinition]:
|
||||
def get(self, name: str) -> CapabilityDefinition | None:
|
||||
return _capabilities.get(name)
|
||||
|
||||
def __getitem__(self, name: str) -> CapabilityDefinition:
|
||||
@@ -210,10 +209,10 @@ class _CapabilityDefinitionsProxy:
|
||||
def __contains__(self, name: str) -> bool:
|
||||
return name in _capabilities
|
||||
|
||||
def values(self) -> List[CapabilityDefinition]:
|
||||
def values(self) -> list[CapabilityDefinition]:
|
||||
return list(_capabilities.values())
|
||||
|
||||
def items(self) -> List[Tuple[str, CapabilityDefinition]]:
|
||||
def items(self) -> list[tuple[str, CapabilityDefinition]]:
|
||||
return list(_capabilities.items())
|
||||
|
||||
|
||||
|
||||
@@ -13,7 +13,6 @@ allowed_models 格式: ["claude-sonnet-4", "gpt-4o"]
|
||||
|
||||
import re
|
||||
from functools import lru_cache
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import regex
|
||||
|
||||
@@ -26,10 +25,10 @@ MAX_MODEL_NAME_LENGTH = 200 # 与 MAX_MAPPING_LENGTH 保持一致
|
||||
REGEX_MATCH_TIMEOUT_MS = 100 # 正则匹配超时(毫秒)
|
||||
|
||||
# 类型别名
|
||||
AllowedModels = Optional[List[str]]
|
||||
type AllowedModels = list[str] | None
|
||||
|
||||
|
||||
def normalize_allowed_models(allowed_models: AllowedModels) -> Optional[set[str]]:
|
||||
def normalize_allowed_models(allowed_models: AllowedModels) -> set[str] | None:
|
||||
"""
|
||||
将 allowed_models 规范化为模型名称集合
|
||||
|
||||
@@ -130,7 +129,7 @@ def get_allowed_models_preview(
|
||||
return preview
|
||||
|
||||
|
||||
def parse_allowed_models_to_list(allowed_models: AllowedModels) -> List[str]:
|
||||
def parse_allowed_models_to_list(allowed_models: AllowedModels) -> list[str]:
|
||||
"""
|
||||
解析 allowed_models 为列表
|
||||
|
||||
@@ -146,7 +145,7 @@ def parse_allowed_models_to_list(allowed_models: AllowedModels) -> List[str]:
|
||||
return list(allowed_models)
|
||||
|
||||
|
||||
def validate_mapping_pattern(pattern: str) -> Tuple[bool, Optional[str]]:
|
||||
def validate_mapping_pattern(pattern: str) -> tuple[bool, str | None]:
|
||||
"""
|
||||
验证映射模式是否安全
|
||||
|
||||
@@ -171,7 +170,7 @@ def validate_mapping_pattern(pattern: str) -> Tuple[bool, Optional[str]]:
|
||||
return True, None
|
||||
|
||||
|
||||
def validate_model_mappings(mappings: Optional[List[str]]) -> Tuple[bool, Optional[str]]:
|
||||
def validate_model_mappings(mappings: list[str] | None) -> tuple[bool, str | None]:
|
||||
"""
|
||||
验证映射列表是否合法
|
||||
|
||||
@@ -196,8 +195,8 @@ def validate_model_mappings(mappings: Optional[List[str]]) -> Tuple[bool, Option
|
||||
|
||||
|
||||
def validate_and_extract_model_mappings(
|
||||
config: Optional[dict],
|
||||
) -> Tuple[bool, Optional[str], Optional[List[str]]]:
|
||||
config: dict | None,
|
||||
) -> tuple[bool, str | None, list[str] | None]:
|
||||
"""
|
||||
从 config 中验证并提取 model_mappings
|
||||
|
||||
@@ -238,7 +237,7 @@ def validate_and_extract_model_mappings(
|
||||
|
||||
|
||||
@lru_cache(maxsize=2000)
|
||||
def _compile_pattern_cached(pattern: str) -> Optional[regex.Pattern]:
|
||||
def _compile_pattern_cached(pattern: str) -> regex.Pattern | None:
|
||||
"""
|
||||
编译正则模式(带 LRU 缓存)
|
||||
|
||||
@@ -267,7 +266,7 @@ def clear_regex_cache() -> None:
|
||||
|
||||
def _match_with_timeout(
|
||||
compiled_regex: regex.Pattern, text: str, timeout_ms: int = REGEX_MATCH_TIMEOUT_MS
|
||||
) -> Optional[bool]:
|
||||
) -> bool | None:
|
||||
"""
|
||||
带超时的正则匹配(使用 regex 库的原生超时支持)
|
||||
|
||||
@@ -342,9 +341,9 @@ def match_model_with_pattern(pattern: str, model_name: str) -> bool:
|
||||
def check_model_allowed_with_mappings(
|
||||
model_name: str,
|
||||
allowed_models: AllowedModels,
|
||||
model_mappings: Optional[List[str]] = None,
|
||||
candidate_models: Optional[set[str]] = None,
|
||||
) -> tuple[bool, Optional[str]]:
|
||||
model_mappings: list[str] | None = None,
|
||||
candidate_models: set[str] | None = None,
|
||||
) -> tuple[bool, str | None]:
|
||||
"""
|
||||
检查模型是否被允许(支持映射通配符匹配)
|
||||
|
||||
|
||||
@@ -4,9 +4,14 @@
|
||||
包含模块元数据、定义和状态的数据结构
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable, List, Optional, Tuple
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Awaitable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastapi import APIRouter
|
||||
@@ -50,16 +55,16 @@ class ModuleMetadata:
|
||||
# 可用性控制(部署级)
|
||||
env_key: str # 环境变量名: LDAP_AVAILABLE
|
||||
default_available: bool = False # 默认是否可用
|
||||
required_packages: List[str] = field(default_factory=list) # 依赖的 Python 包
|
||||
dependencies: List[str] = field(default_factory=list) # 依赖的其他模块
|
||||
required_packages: list[str] = field(default_factory=list) # 依赖的 Python 包
|
||||
dependencies: list[str] = field(default_factory=list) # 依赖的其他模块
|
||||
|
||||
# 路由配置 - 模块自定义前缀
|
||||
api_prefix: Optional[str] = None # 如 "/api/admin/ldap"
|
||||
api_prefix: str | None = None # 如 "/api/admin/ldap"
|
||||
|
||||
# 前端配置
|
||||
admin_route: Optional[str] = None # 管理页面路由: "/admin/ldap"
|
||||
admin_menu_icon: Optional[str] = None # 菜单图标
|
||||
admin_menu_group: Optional[str] = None # 菜单分组: "system", "security"
|
||||
admin_route: str | None = None # 管理页面路由: "/admin/ldap"
|
||||
admin_menu_icon: str | None = None # 菜单图标
|
||||
admin_menu_group: str | None = None # 菜单分组: "system", "security"
|
||||
admin_menu_order: int = 100 # 菜单排序(越小越靠前)
|
||||
|
||||
|
||||
@@ -74,19 +79,19 @@ class ModuleDefinition:
|
||||
metadata: ModuleMetadata
|
||||
|
||||
# 工厂函数 - 内部再 import 重依赖
|
||||
router_factory: Optional[Callable[[], "APIRouter"]] = None
|
||||
service_factory: Optional[Callable[[], Any]] = None
|
||||
router_factory: Callable[[], APIRouter] | None = None
|
||||
service_factory: Callable[[], Any] | None = None
|
||||
|
||||
# 生命周期钩子
|
||||
on_startup: Optional[Callable[[], Awaitable[None]]] = None
|
||||
on_shutdown: Optional[Callable[[], Awaitable[None]]] = None
|
||||
health_check: Optional[Callable[[], Awaitable[ModuleHealth]]] = None
|
||||
on_startup: Callable[[], Awaitable[None]] | None = None
|
||||
on_shutdown: Callable[[], Awaitable[None]] | None = None
|
||||
health_check: Callable[[], Awaitable[ModuleHealth]] | None = None
|
||||
|
||||
# 自定义依赖检测(可选,用于检测 ldap3 等库是否安装)
|
||||
check_dependencies: Optional[Callable[[], bool]] = None
|
||||
check_dependencies: Callable[[], bool] | None = None
|
||||
|
||||
# 配置验证(可选,启用模块时调用,返回 (success, error_message))
|
||||
validate_config: Optional[Callable[["Session"], Tuple[bool, str]]] = None
|
||||
validate_config: Callable[[Session], tuple[bool, str]] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -102,7 +107,7 @@ class ModuleStatus:
|
||||
enabled: bool # 运行级启用(数据库配置)
|
||||
active: bool # 最终激活状态 (available && enabled && dependencies_ok)
|
||||
config_validated: bool # 配置验证通过(只有验证通过才允许启用)
|
||||
config_error: Optional[str] # 配置验证失败的错误信息
|
||||
config_error: str | None # 配置验证失败的错误信息
|
||||
|
||||
# 显示信息
|
||||
display_name: str
|
||||
@@ -110,9 +115,9 @@ class ModuleStatus:
|
||||
category: ModuleCategory
|
||||
|
||||
# 前端配置
|
||||
admin_route: Optional[str]
|
||||
admin_menu_icon: Optional[str]
|
||||
admin_menu_group: Optional[str]
|
||||
admin_route: str | None
|
||||
admin_menu_icon: str | None
|
||||
admin_menu_group: str | None
|
||||
admin_menu_order: int
|
||||
|
||||
# 健康状态
|
||||
|
||||
@@ -4,9 +4,11 @@
|
||||
负责模块的注册、状态管理和生命周期控制
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Dict, List, Optional, Set
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.core.modules.base import (
|
||||
@@ -31,14 +33,14 @@ class ModuleRegistry:
|
||||
- 提供模块状态查询
|
||||
"""
|
||||
|
||||
_instance: Optional["ModuleRegistry"] = None
|
||||
_instance: ModuleRegistry | None = None
|
||||
|
||||
def __init__(self):
|
||||
self._modules: Dict[str, ModuleDefinition] = {}
|
||||
self._initialized: Set[str] = set()
|
||||
self._modules: dict[str, ModuleDefinition] = {}
|
||||
self._initialized: set[str] = set()
|
||||
|
||||
@classmethod
|
||||
def get_instance(cls) -> "ModuleRegistry":
|
||||
def get_instance(cls) -> ModuleRegistry:
|
||||
"""获取单例实例"""
|
||||
if cls._instance is None:
|
||||
cls._instance = cls()
|
||||
@@ -63,11 +65,11 @@ class ModuleRegistry:
|
||||
self._modules[name] = module
|
||||
logger.debug(f"Module [{name}] registered")
|
||||
|
||||
def get_module(self, name: str) -> Optional[ModuleDefinition]:
|
||||
def get_module(self, name: str) -> ModuleDefinition | None:
|
||||
"""获取模块定义"""
|
||||
return self._modules.get(name)
|
||||
|
||||
def get_all_modules(self) -> List[ModuleDefinition]:
|
||||
def get_all_modules(self) -> list[ModuleDefinition]:
|
||||
"""获取所有已注册模块"""
|
||||
return list(self._modules.values())
|
||||
|
||||
@@ -115,13 +117,13 @@ class ModuleRegistry:
|
||||
|
||||
return True
|
||||
|
||||
def get_available_modules(self) -> List[ModuleDefinition]:
|
||||
def get_available_modules(self) -> list[ModuleDefinition]:
|
||||
"""获取所有部署可用的模块"""
|
||||
return [m for m in self._modules.values() if self.is_available(m.metadata.name)]
|
||||
|
||||
# ========== 启用状态检查(运行级)==========
|
||||
|
||||
def is_enabled(self, name: str, db: "Session") -> bool:
|
||||
def is_enabled(self, name: str, db: Session) -> bool:
|
||||
"""
|
||||
检查模块是否运行启用(数据库配置)
|
||||
|
||||
@@ -135,7 +137,7 @@ class ModuleRegistry:
|
||||
value = SystemConfigService.get_config(db, config_key, default=False)
|
||||
return bool(value)
|
||||
|
||||
def set_enabled(self, name: str, enabled: bool, db: "Session") -> None:
|
||||
def set_enabled(self, name: str, enabled: bool, db: Session) -> None:
|
||||
"""
|
||||
设置模块启用状态
|
||||
|
||||
@@ -156,7 +158,7 @@ class ModuleRegistry:
|
||||
|
||||
# ========== 激活状态检查 ==========
|
||||
|
||||
def is_active(self, name: str, db: "Session") -> bool:
|
||||
def is_active(self, name: str, db: Session) -> bool:
|
||||
"""
|
||||
检查模块是否最终激活
|
||||
|
||||
@@ -177,7 +179,7 @@ class ModuleRegistry:
|
||||
|
||||
# ========== 配置验证 ==========
|
||||
|
||||
def validate_config(self, name: str, db: "Session") -> tuple[bool, str]:
|
||||
def validate_config(self, name: str, db: Session) -> tuple[bool, str]:
|
||||
"""
|
||||
验证模块配置是否有效
|
||||
|
||||
@@ -206,8 +208,8 @@ class ModuleRegistry:
|
||||
# ========== 状态查询 ==========
|
||||
|
||||
def get_module_status(
|
||||
self, name: str, db: "Session", health: Optional[ModuleHealth] = None
|
||||
) -> Optional[ModuleStatus]:
|
||||
self, name: str, db: Session, health: ModuleHealth | None = None
|
||||
) -> ModuleStatus | None:
|
||||
"""
|
||||
获取单个模块状态
|
||||
|
||||
@@ -225,7 +227,7 @@ class ModuleRegistry:
|
||||
|
||||
# 获取配置验证状态
|
||||
config_validated = False
|
||||
config_error: Optional[str] = None
|
||||
config_error: str | None = None
|
||||
if available:
|
||||
config_validated, config_error = self.validate_config(name, db)
|
||||
if config_validated:
|
||||
@@ -279,8 +281,8 @@ class ModuleRegistry:
|
||||
return ModuleHealth.UNHEALTHY
|
||||
|
||||
async def get_module_status_async(
|
||||
self, name: str, db: "Session"
|
||||
) -> Optional[ModuleStatus]:
|
||||
self, name: str, db: Session
|
||||
) -> ModuleStatus | None:
|
||||
"""异步获取模块状态(包含健康检查)"""
|
||||
if name not in self._modules:
|
||||
return None
|
||||
@@ -288,7 +290,7 @@ class ModuleRegistry:
|
||||
health = await self.check_health(name) if self.is_available(name) else ModuleHealth.UNKNOWN
|
||||
return self.get_module_status(name, db, health=health)
|
||||
|
||||
async def get_all_status_async(self, db: "Session") -> Dict[str, ModuleStatus]:
|
||||
async def get_all_status_async(self, db: Session) -> dict[str, ModuleStatus]:
|
||||
"""异步获取所有模块状态(包含健康检查)"""
|
||||
result = {}
|
||||
for name in self._modules:
|
||||
@@ -297,7 +299,7 @@ class ModuleRegistry:
|
||||
result[name] = status
|
||||
return result
|
||||
|
||||
def get_all_status(self, db: "Session") -> Dict[str, ModuleStatus]:
|
||||
def get_all_status(self, db: Session) -> dict[str, ModuleStatus]:
|
||||
"""获取所有模块状态(同步版本,不含健康检查)"""
|
||||
result = {}
|
||||
for name in self._modules:
|
||||
@@ -306,7 +308,7 @@ class ModuleRegistry:
|
||||
result[name] = status
|
||||
return result
|
||||
|
||||
def get_available_status(self, db: "Session") -> Dict[str, ModuleStatus]:
|
||||
def get_available_status(self, db: Session) -> dict[str, ModuleStatus]:
|
||||
"""获取所有可用模块的状态"""
|
||||
result = {}
|
||||
for name, module in self._modules.items():
|
||||
@@ -316,7 +318,7 @@ class ModuleRegistry:
|
||||
result[name] = status
|
||||
return result
|
||||
|
||||
def get_auth_modules_status(self, db: "Session") -> List[ModuleStatus]:
|
||||
def get_auth_modules_status(self, db: Session) -> list[ModuleStatus]:
|
||||
"""获取认证模块状态(供登录页使用)"""
|
||||
result = []
|
||||
for name, module in self._modules.items():
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
优化工具类 - 包含Token计数和响应头管理
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
import tiktoken
|
||||
|
||||
|
||||
@@ -5,8 +5,7 @@
|
||||
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any
|
||||
|
||||
|
||||
class ProviderHealthTracker:
|
||||
@@ -26,11 +25,11 @@ class ProviderHealthTracker:
|
||||
self.recovery_time = recovery_time
|
||||
|
||||
# 存储每个提供商的失败记录
|
||||
self.failures: Dict[str, list] = defaultdict(list)
|
||||
self.failures: dict[str, list] = defaultdict(list)
|
||||
# 存储每个提供商的成功记录
|
||||
self.successes: Dict[str, list] = defaultdict(list)
|
||||
self.successes: dict[str, list] = defaultdict(list)
|
||||
# 存储优先级调整
|
||||
self.priority_adjustments: Dict[str, int] = {}
|
||||
self.priority_adjustments: dict[str, int] = {}
|
||||
|
||||
def record_success(self, provider_name: str) -> None:
|
||||
"""记录成功的请求"""
|
||||
@@ -71,7 +70,7 @@ class ProviderHealthTracker:
|
||||
"""
|
||||
return self.priority_adjustments.get(provider_name, 0)
|
||||
|
||||
def get_health_status(self, provider_name: str) -> Dict:
|
||||
def get_health_status(self, provider_name: str) -> dict:
|
||||
"""
|
||||
获取提供商的健康状态
|
||||
"""
|
||||
@@ -146,7 +145,7 @@ class SimpleProviderSelector:
|
||||
def __init__(self, health_tracker: ProviderHealthTracker):
|
||||
self.health_tracker = health_tracker
|
||||
|
||||
def select_provider(self, providers: list, specified_provider: Optional[str] = None) -> Any:
|
||||
def select_provider(self, providers: list, specified_provider: str | None = None) -> Any:
|
||||
"""
|
||||
选择提供商
|
||||
|
||||
|
||||
@@ -10,9 +10,11 @@ import time
|
||||
import traceback
|
||||
import uuid
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from datetime import datetime, timezone
|
||||
from enum import Enum
|
||||
from typing import Any, Callable, Dict, List, Optional, Type, Union
|
||||
from typing import Any
|
||||
|
||||
from collections.abc import Callable
|
||||
|
||||
from ..core.exceptions import ProxyException
|
||||
from src.core.logger import logger
|
||||
@@ -43,7 +45,7 @@ class ErrorPattern:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
error_types: List[Type[Exception]],
|
||||
error_types: list[type[Exception]],
|
||||
severity: ErrorSeverity,
|
||||
recovery_strategy: RecoveryStrategy,
|
||||
user_message: str,
|
||||
@@ -113,10 +115,10 @@ class ResilienceManager:
|
||||
"""系统韧性管理器"""
|
||||
|
||||
def __init__(self):
|
||||
self.error_patterns: List[ErrorPattern] = []
|
||||
self.circuit_breakers: Dict[str, CircuitBreaker] = {}
|
||||
self.error_stats: Dict[str, int] = {}
|
||||
self.last_errors: List[Dict[str, Any]] = []
|
||||
self.error_patterns: list[ErrorPattern] = []
|
||||
self.circuit_breakers: dict[str, CircuitBreaker] = {}
|
||||
self.error_stats: dict[str, int] = {}
|
||||
self.last_errors: list[dict[str, Any]] = []
|
||||
self._setup_default_patterns()
|
||||
|
||||
def _setup_default_patterns(self):
|
||||
@@ -196,8 +198,8 @@ class ResilienceManager:
|
||||
return self.circuit_breakers[key]
|
||||
|
||||
def handle_error(
|
||||
self, error: Exception, context: Dict[str, Any] = None, operation: str = "unknown"
|
||||
) -> Dict[str, Any]:
|
||||
self, error: Exception, context: dict[str, Any] = None, operation: str = "unknown"
|
||||
) -> dict[str, Any]:
|
||||
"""处理错误并返回处理结果"""
|
||||
|
||||
error_id = str(uuid.uuid4())[:8]
|
||||
@@ -250,14 +252,14 @@ class ResilienceManager:
|
||||
"pattern": None,
|
||||
}
|
||||
|
||||
def _find_matching_pattern(self, error: Exception) -> Optional[ErrorPattern]:
|
||||
def _find_matching_pattern(self, error: Exception) -> ErrorPattern | None:
|
||||
"""查找匹配的错误处理模式"""
|
||||
for pattern in self.error_patterns:
|
||||
if any(isinstance(error, error_type) for error_type in pattern.error_types):
|
||||
return pattern
|
||||
return None
|
||||
|
||||
def get_error_stats(self) -> Dict[str, Any]:
|
||||
def get_error_stats(self) -> dict[str, Any]:
|
||||
"""获取错误统计"""
|
||||
return {
|
||||
"total_errors": sum(self.error_stats.values()),
|
||||
@@ -279,7 +281,7 @@ def resilient_operation(
|
||||
max_retries: int = None,
|
||||
retry_delay: float = None,
|
||||
circuit_breaker_key: str = None,
|
||||
context: Dict[str, Any] = None,
|
||||
context: dict[str, Any] = None,
|
||||
):
|
||||
"""
|
||||
韧性操作装饰器
|
||||
@@ -354,7 +356,7 @@ def resilient_operation(
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def safe_operation(operation_name: str, context: Dict[str, Any] = None):
|
||||
async def safe_operation(operation_name: str, context: dict[str, Any] = None):
|
||||
"""
|
||||
安全操作上下文管理器
|
||||
自动处理异常并提供用户友好的错误信息
|
||||
|
||||
@@ -4,7 +4,6 @@
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import List, Optional
|
||||
|
||||
|
||||
class PasswordValidator:
|
||||
@@ -14,7 +13,7 @@ class PasswordValidator:
|
||||
MAX_LENGTH = 128
|
||||
|
||||
@classmethod
|
||||
def validate(cls, password: str) -> tuple[bool, Optional[str]]:
|
||||
def validate(cls, password: str) -> tuple[bool, str | None]:
|
||||
"""
|
||||
验证密码复杂度
|
||||
|
||||
@@ -109,7 +108,7 @@ class EmailValidator:
|
||||
EMAIL_REGEX = re.compile(r"^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$")
|
||||
|
||||
@classmethod
|
||||
def validate(cls, email: str) -> tuple[bool, Optional[str]]:
|
||||
def validate(cls, email: str) -> tuple[bool, str | None]:
|
||||
"""
|
||||
验证邮箱格式
|
||||
|
||||
@@ -139,7 +138,7 @@ class UsernameValidator:
|
||||
USERNAME_REGEX = re.compile(r"^[a-zA-Z0-9_.\-]+$")
|
||||
|
||||
@classmethod
|
||||
def validate(cls, username: str) -> tuple[bool, Optional[str]]:
|
||||
def validate(cls, username: str) -> tuple[bool, str | None]:
|
||||
"""
|
||||
验证用户名
|
||||
|
||||
|
||||
Reference in New Issue
Block a user