refactor(conversion): body_rules 保护 cache-sensitive 字段,normalizer 保真优化与诊断日志

- RequestBuilder 新增 protected_body_keys 机制,按 provider API 格式阻止 body_rules
  改写 prompt cache 相关的顶层请求字段(messages/tools/system 等)
- Claude/Gemini normalizer 优先复用原始 raw tool_choice,避免 round-trip 丢失信息
- Claude normalizer 修复 content blocks 输出顺序(flush_text_parts),
  _coerce_claude_message_sequence 返回结构化诊断
- OpenAI normalizer 保留 raw tool call arguments 字符串与原始 tool 定义 extra 字段
- schema_utils allOf 合并保持 required 字段插入顺序
- 各转换环节增加结构化 debug 日志用于调试
This commit is contained in:
fawney19
2026-03-17 01:25:29 +08:00
parent c97c9332eb
commit 4ecaefbade
22 changed files with 1191 additions and 115 deletions

View File

@@ -44,7 +44,11 @@ from src.api.handlers.base.chat_error_utils import (
_resolve_dynamic_format, _resolve_dynamic_format,
) )
from src.api.handlers.base.parsers import get_parser_for_format from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.request_builder import PassthroughRequestBuilder, get_provider_auth from src.api.handlers.base.request_builder import (
PassthroughRequestBuilder,
get_cache_sensitive_protected_body_keys,
get_provider_auth,
)
from src.api.handlers.base.response_parser import ResponseParser from src.api.handlers.base.response_parser import ResponseParser
from src.api.handlers.base.stream_context import ( from src.api.handlers.base.stream_context import (
StreamContext, StreamContext,
@@ -98,6 +102,7 @@ class ProviderRequestResult:
mapped_model: str | None mapped_model: str | None
envelope: Any # ProviderEnvelope | None envelope: Any # ProviderEnvelope | None
extra_headers: dict[str, str] = field(default_factory=dict) extra_headers: dict[str, str] = field(default_factory=dict)
protected_body_keys: frozenset[str] = field(default_factory=frozenset)
upstream_is_stream: bool = True upstream_is_stream: bool = True
needs_conversion: bool = False needs_conversion: bool = False
provider_api_format: str = "" provider_api_format: str = ""
@@ -828,6 +833,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
mapped_model=mapped_model, mapped_model=mapped_model,
envelope=envelope, envelope=envelope,
extra_headers=extra_headers, extra_headers=extra_headers,
protected_body_keys=get_cache_sensitive_protected_body_keys(provider_api_format),
upstream_is_stream=upstream_is_stream, upstream_is_stream=upstream_is_stream,
needs_conversion=needs_conversion, needs_conversion=needs_conversion,
provider_api_format=provider_api_format, provider_api_format=provider_api_format,
@@ -907,6 +913,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
extra_headers=prep.extra_headers if prep.extra_headers else None, extra_headers=prep.extra_headers if prep.extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None, pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=envelope, envelope=envelope,
protected_body_keys=prep.protected_body_keys,
provider_api_format=prep.provider_api_format,
) )
if upstream_is_stream: if upstream_is_stream:
from src.core.api_format.headers import set_accept_if_absent from src.core.api_format.headers import set_accept_if_absent

View File

@@ -488,6 +488,8 @@ class ChatSyncExecutor:
extra_headers=prep.extra_headers if prep.extra_headers else None, extra_headers=prep.extra_headers if prep.extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None, pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=envelope, envelope=envelope,
protected_body_keys=prep.protected_body_keys,
provider_api_format=prep.provider_api_format,
) )
if upstream_is_stream: if upstream_is_stream:
from src.core.api_format.headers import set_accept_if_absent from src.core.api_format.headers import set_accept_if_absent

View File

@@ -18,7 +18,10 @@ from src.api.handlers.base.base_handler import (
wait_for_with_disconnect_detection, wait_for_with_disconnect_detection,
) )
from src.api.handlers.base.parsers import get_parser_for_format from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.request_builder import get_provider_auth from src.api.handlers.base.request_builder import (
get_cache_sensitive_protected_body_keys,
get_provider_auth,
)
from src.api.handlers.base.stream_context import StreamContext from src.api.handlers.base.stream_context import StreamContext
from src.api.handlers.base.utils import ( from src.api.handlers.base.utils import (
build_sse_headers, build_sse_headers,
@@ -427,6 +430,8 @@ class CliStreamMixin:
extra_headers=extra_headers if extra_headers else None, extra_headers=extra_headers if extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None, pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=envelope, envelope=envelope,
protected_body_keys=get_cache_sensitive_protected_body_keys(provider_api_format),
provider_api_format=provider_api_format,
) )
if upstream_is_stream: if upstream_is_stream:
from src.core.api_format.headers import set_accept_if_absent from src.core.api_format.headers import set_accept_if_absent

View File

@@ -10,7 +10,10 @@ import httpx
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from src.api.handlers.base.parsers import get_parser_for_format from src.api.handlers.base.parsers import get_parser_for_format
from src.api.handlers.base.request_builder import get_provider_auth from src.api.handlers.base.request_builder import (
get_cache_sensitive_protected_body_keys,
get_provider_auth,
)
from src.api.handlers.base.stream_context import extract_proxy_timing, is_format_converted from src.api.handlers.base.stream_context import extract_proxy_timing, is_format_converted
from src.api.handlers.base.upstream_stream_bridge import ( from src.api.handlers.base.upstream_stream_bridge import (
aggregate_upstream_stream_to_internal_response, aggregate_upstream_stream_to_internal_response,
@@ -248,6 +251,8 @@ class CliSyncMixin:
extra_headers=extra_headers if extra_headers else None, extra_headers=extra_headers if extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None, pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=envelope, envelope=envelope,
protected_body_keys=get_cache_sensitive_protected_body_keys(provider_api_format),
provider_api_format=provider_api_format,
) )
if upstream_is_stream: if upstream_is_stream:
from src.core.api_format.headers import set_accept_if_absent from src.core.api_format.headers import set_accept_if_absent

View File

@@ -366,7 +366,11 @@ class HandlerAdapterBase(ApiAdapter):
统一的 endpoint 测试方法,支持 OAuth/Antigravity/Kiro 等特殊路由。 统一的 endpoint 测试方法,支持 OAuth/Antigravity/Kiro 等特殊路由。
""" """
from src.api.handlers.base.endpoint_checker import run_endpoint_check from src.api.handlers.base.endpoint_checker import run_endpoint_check
from src.api.handlers.base.request_builder import apply_body_rules, evaluate_condition from src.api.handlers.base.request_builder import (
apply_body_rules,
evaluate_condition,
get_cache_sensitive_protected_body_keys,
)
from src.core.api_format.headers import HeaderBuilder from src.core.api_format.headers import HeaderBuilder
from src.core.provider_types import ProviderType from src.core.provider_types import ProviderType
@@ -507,7 +511,12 @@ class HandlerAdapterBase(ApiAdapter):
) )
if body_rules: if body_rules:
body = apply_body_rules(body, body_rules, original_body=body) body = apply_body_rules(
body,
body_rules,
protected_keys=get_cache_sensitive_protected_body_keys(cls.FORMAT_ID),
original_body=body,
)
if is_antigravity: if is_antigravity:
from src.services.provider.adapters.antigravity.envelope import ( from src.services.provider.adapters.antigravity.envelope import (

View File

@@ -27,10 +27,70 @@ from src.core.api_format import (
resolve_header_name_case, resolve_header_name_case,
) )
from src.core.crypto import crypto_service from src.core.crypto import crypto_service
from src.core.logger import logger
from src.models.endpoint_models import _CONDITION_OPS, _TYPE_IS_VALUES, parse_re_flags from src.models.endpoint_models import _CONDITION_OPS, _TYPE_IS_VALUES, parse_re_flags
from src.services.provider.auth import get_provider_auth # noqa: F401 from src.services.provider.auth import get_provider_auth # noqa: F401
from src.services.provider.envelope import ProviderEnvelope from src.services.provider.envelope import ProviderEnvelope
def _payload_item_count(value: Any) -> int | None:
"""统计顶层 prompt-bearing 容器项数量;标量按 1 处理。"""
if isinstance(value, list):
return len(value)
if value is None:
return None
if isinstance(value, (str, bytes)):
return 1 if value else 0
if isinstance(value, dict):
return 1 if value else 0
return 1
def summarize_request_payload_shape(
payload: dict[str, Any],
*,
provider_api_format: str | None,
protected_body_keys: frozenset[str] | None,
body_rules: Any,
) -> dict[str, Any]:
"""生成最终出站 payload 的结构化摘要,避免记录正文。"""
tools = payload.get("tools")
tool_count = len(tools) if isinstance(tools, list) else None
function_declaration_count = 0
if isinstance(tools, list):
for tool in tools:
if not isinstance(tool, dict):
continue
decls = tool.get("function_declarations") or tool.get("functionDeclarations")
if isinstance(decls, list):
function_declaration_count += len(decls)
return {
"format": str(provider_api_format or "").strip().lower() or None,
"top_level_keys": sorted(payload.keys()),
"message_count": _payload_item_count(payload.get("messages")),
"input_count": _payload_item_count(payload.get("input")),
"contents_count": _payload_item_count(payload.get("contents")),
"tool_count": tool_count,
"function_declaration_count": function_declaration_count,
"has_system": any(
k in payload for k in ("system", "system_instruction", "systemInstruction")
),
"has_instructions": "instructions" in payload,
"has_tool_choice": any(
k in payload for k in ("tool_choice", "toolChoice", "tool_config", "toolConfig")
),
"has_generation_config": any(
k in payload for k in ("generation_config", "generationConfig")
),
"has_prompt_cache_key": bool(str(payload.get("prompt_cache_key") or "").strip()),
"protected_body_keys_enabled": bool(protected_body_keys),
"protected_body_keys": sorted(protected_body_keys or ()),
"body_rule_count": len(body_rules) if isinstance(body_rules, list) else 0,
}
# ============================================================================== # ==============================================================================
# 统一的头部配置常量 # 统一的头部配置常量
# ============================================================================== # ==============================================================================
@@ -46,6 +106,38 @@ PROTECTED_BODY_FIELDS: frozenset[str] = frozenset(
} }
) )
# cache-sensitive 顶层字段:用于阻止 endpoint body_rules 在最终出站前
# 二次改写 prompt-bearing 结构,破坏上游 prompt cache 一致性。
_CACHE_SENSITIVE_BODY_FIELDS_BY_FORMAT: dict[str, frozenset[str]] = {
"openai:chat": frozenset({"messages", "tools", "tool_choice"}),
"openai:cli": frozenset({"input", "instructions", "tools", "tool_choice", "prompt_cache_key"}),
"openai:compact": frozenset(
{"input", "instructions", "tools", "tool_choice", "prompt_cache_key"}
),
"claude:chat": frozenset({"messages", "system", "tools", "tool_choice"}),
"gemini:chat": frozenset(
{
"contents",
"system_instruction",
"systemInstruction",
"tools",
"tool_config",
"toolConfig",
"generation_config",
"generationConfig",
}
),
}
def get_cache_sensitive_protected_body_keys(provider_api_format: str | None) -> frozenset[str]:
"""根据目标 Provider API 格式返回需要保护的顶层请求字段。"""
fmt = str(provider_api_format or "").strip().lower()
extra = _CACHE_SENSITIVE_BODY_FIELDS_BY_FORMAT.get(fmt, frozenset())
if not extra:
return PROTECTED_BODY_FIELDS
return frozenset({*PROTECTED_BODY_FIELDS, *extra})
# ============================================================================== # ==============================================================================
# 测试请求常量与辅助函数 # 测试请求常量与辅助函数
@@ -1129,6 +1221,8 @@ class RequestBuilder(ABC):
extra_headers: dict[str, str] | None = None, extra_headers: dict[str, str] | None = None,
pre_computed_auth: tuple[str, str] | None = None, pre_computed_auth: tuple[str, str] | None = None,
envelope: ProviderEnvelope | None = None, envelope: ProviderEnvelope | None = None,
protected_body_keys: frozenset[str] | None = None,
provider_api_format: str | None = None,
) -> tuple[dict[str, Any], dict[str, str]]: ) -> tuple[dict[str, Any], dict[str, str]]:
""" """
构建完整的请求(请求体 + 请求头) 构建完整的请求(请求体 + 请求头)
@@ -1142,6 +1236,8 @@ class RequestBuilder(ABC):
is_stream: 是否为流式请求 is_stream: 是否为流式请求
extra_headers: 额外请求头 extra_headers: 额外请求头
pre_computed_auth: 预先计算的认证信息 (auth_header, auth_value) pre_computed_auth: 预先计算的认证信息 (auth_header, auth_value)
protected_body_keys: body_rules 不允许修改的顶层请求字段
provider_api_format: 运行时实际生效的 Provider API 格式,用于准确记录摘要日志
Returns: Returns:
Tuple[payload, headers] Tuple[payload, headers]
@@ -1155,7 +1251,23 @@ class RequestBuilder(ABC):
# 应用请求体规则(如果 endpoint 配置了 body_rules # 应用请求体规则(如果 endpoint 配置了 body_rules
body_rules = getattr(endpoint, "body_rules", None) body_rules = getattr(endpoint, "body_rules", None)
if body_rules: if body_rules:
payload = apply_body_rules(payload, body_rules, original_body=original_body) payload = apply_body_rules(
payload,
body_rules,
protected_keys=protected_body_keys,
original_body=original_body,
)
effective_provider_api_format = provider_api_format or getattr(endpoint, "api_format", None)
logger.debug(
"[RequestBuilder] outbound payload summary: {}",
summarize_request_payload_shape(
payload,
provider_api_format=effective_provider_api_format,
protected_body_keys=protected_body_keys,
body_rules=body_rules,
),
)
headers = self.build_headers( headers = self.build_headers(
original_headers, original_headers,
@@ -1183,8 +1295,8 @@ class PassthroughRequestBuilder(RequestBuilder):
self, self,
original_body: dict[str, Any], original_body: dict[str, Any],
*, *,
mapped_model: str | None = None, # noqa: ARG002 - 由 apply_mapped_model 处理 mapped_model: str | None = None,
is_stream: bool = False, # noqa: ARG002 - 保留原始值,不自动添加 is_stream: bool = False,
) -> dict[str, Any]: ) -> dict[str, Any]:
""" """
透传请求体 - 原样复制,不做任何修改 透传请求体 - 原样复制,不做任何修改
@@ -1193,6 +1305,7 @@ class PassthroughRequestBuilder(RequestBuilder):
- model: 由各 handler 的 apply_mapped_model 方法处理 - model: 由各 handler 的 apply_mapped_model 方法处理
- stream: 保留客户端原始值(不同 API 处理方式不同) - stream: 保留客户端原始值(不同 API 处理方式不同)
""" """
del mapped_model, is_stream
return dict(original_body) return dict(original_body)
@staticmethod @staticmethod

View File

@@ -248,7 +248,11 @@ class GeminiChatAdapter(ChatAdapterBase):
) -> dict[str, Any]: ) -> dict[str, Any]:
"""测试 Gemini API 模型连接性(非流式)""" """测试 Gemini API 模型连接性(非流式)"""
from src.api.handlers.base.endpoint_checker import run_endpoint_check from src.api.handlers.base.endpoint_checker import run_endpoint_check
from src.api.handlers.base.request_builder import apply_body_rules, evaluate_condition from src.api.handlers.base.request_builder import (
apply_body_rules,
evaluate_condition,
get_cache_sensitive_protected_body_keys,
)
from src.core.api_format.headers import HeaderBuilder from src.core.api_format.headers import HeaderBuilder
# Gemini需要从request_data或model_name参数获取model名称 # Gemini需要从request_data或model_name参数获取model名称
@@ -344,7 +348,12 @@ class GeminiChatAdapter(ChatAdapterBase):
# 应用请求体规则(在格式转换后应用,确保规则效果不被覆盖) # 应用请求体规则(在格式转换后应用,确保规则效果不被覆盖)
if body_rules: if body_rules:
body = apply_body_rules(body, body_rules, original_body=body) body = apply_body_rules(
body,
body_rules,
protected_keys=get_cache_sensitive_protected_body_keys(cls.FORMAT_ID),
original_body=body,
)
# Antigravity 需要将请求体包装为 v1internal 信封格式 # Antigravity 需要将请求体包装为 v1internal 信封格式
if is_antigravity: if is_antigravity:

View File

@@ -7,8 +7,15 @@ OpenAI Chat <-> Responses API 的工具 / tool_choice / web_search 双向转换
from __future__ import annotations from __future__ import annotations
import json
from typing import Any from typing import Any
def stable_json_dumps(value: Any) -> str:
"""Serialize JSON deterministically for cache-sensitive fallback generation."""
return json.dumps(value, ensure_ascii=False, sort_keys=True)
# Thinking 签名验证的跳过标记 # Thinking 签名验证的跳过标记
# 当无法获取真实签名时,使用此值作为占位符 # 当无法获取真实签名时,使用此值作为占位符
DUMMY_THOUGHT_SIGNATURE = "skip_thought_signature_validator" DUMMY_THOUGHT_SIGNATURE = "skip_thought_signature_validator"

View File

@@ -7,6 +7,7 @@ Claude Messages API Normalizer
- 可选Claude error <-> InternalError - 可选Claude error <-> InternalError
""" """
import copy
import json import json
from typing import Any from typing import Any
@@ -60,6 +61,7 @@ from src.core.api_format.conversion.stream_events import (
ToolCallDeltaEvent, ToolCallDeltaEvent,
) )
from src.core.api_format.conversion.stream_state import StreamState from src.core.api_format.conversion.stream_state import StreamState
from src.core.logger import logger
class ClaudeNormalizer(FormatNormalizer): class ClaudeNormalizer(FormatNormalizer):
@@ -213,6 +215,9 @@ class ClaudeNormalizer(FormatNormalizer):
*, *,
target_variant: str | None = None, target_variant: str | None = None,
) -> dict[str, Any]: ) -> dict[str, Any]:
target_variant_norm = str(target_variant or "").strip().lower()
system_shape = "none"
# system: 如果任一 InstructionSegment 有 cache_control输出数组格式 # system: 如果任一 InstructionSegment 有 cache_control输出数组格式
has_cache_control = any(seg.extra.get("cache_control") for seg in internal.instructions) has_cache_control = any(seg.extra.get("cache_control") for seg in internal.instructions)
if has_cache_control and internal.instructions: if has_cache_control and internal.instructions:
@@ -231,16 +236,43 @@ class ClaudeNormalizer(FormatNormalizer):
] ]
if not system_value: if not system_value:
system_value = None system_value = None
else:
system_shape = "blocks"
else: else:
system_value = internal.system or self._join_instructions(internal.instructions) system_value = internal.system or self._join_instructions(internal.instructions)
if system_value:
system_shape = "string"
# Claude Messages API: messages[] 仅允许 user/assistant且需要交替这里做最小修复 # Claude Messages API: messages[] 仅允许 user/assistant且需要交替这里做最小修复
fixed_messages = self._coerce_claude_message_sequence(internal.messages) fixed_messages, coerce_diagnostics = self._coerce_claude_message_sequence(internal.messages)
out_messages: list[dict[str, Any]] = [ out_messages: list[dict[str, Any]] = [
self._internal_message_to_claude(m) for m in fixed_messages self._internal_message_to_claude(m) for m in fixed_messages
] ]
if coerce_diagnostics["changed"]:
logger.debug(
"[ClaudeNormalizer] normalized message sequence: variant={}, input_count={}, output_count={}, prepended_empty_user={}, merged_adjacent={}, coerced_roles={}",
target_variant_norm or "default",
len(internal.messages),
len(fixed_messages),
coerce_diagnostics["prepended_empty_user"],
coerce_diagnostics["merged_adjacent"],
coerce_diagnostics["coerced_roles"],
)
if system_shape == "blocks":
logger.debug(
"[ClaudeNormalizer] emitted block system payload for cache_control instructions: variant={}, segment_count={}",
target_variant_norm or "default",
len(internal.instructions),
)
elif system_shape == "string" and internal.instructions:
logger.debug(
"[ClaudeNormalizer] emitted string system payload: variant={}, segment_count={}",
target_variant_norm or "default",
len(internal.instructions),
)
# max_tokens: 优先使用请求中的值, 其次 GlobalModel.output_limit, 最后硬编码默认值 # max_tokens: 优先使用请求中的值, 其次 GlobalModel.output_limit, 最后硬编码默认值
effective_max_tokens: int effective_max_tokens: int
if internal.max_tokens is not None: if internal.max_tokens is not None:
@@ -288,11 +320,18 @@ class ClaudeNormalizer(FormatNormalizer):
result["tools"] = claude_tools result["tools"] = claude_tools
if internal.tool_choice: if internal.tool_choice:
reused_raw_tool_choice = isinstance(internal.tool_choice.extra.get("claude"), dict)
tc = self._tool_choice_to_claude(internal.tool_choice) tc = self._tool_choice_to_claude(internal.tool_choice)
# parallel_tool_calls=False -> disable_parallel_tool_use=True # parallel_tool_calls=False -> disable_parallel_tool_use=True
if internal.parallel_tool_calls is False and tc.get("type") != "none": if internal.parallel_tool_calls is False and tc.get("type") != "none":
tc["disable_parallel_tool_use"] = True tc["disable_parallel_tool_use"] = True
result["tool_choice"] = tc result["tool_choice"] = tc
logger.debug(
"[ClaudeNormalizer] {} Claude tool_choice: variant={}, type={}",
"reused raw" if reused_raw_tool_choice else "rebuilt",
target_variant_norm or "default",
tc.get("type"),
)
# thinking 配置 # thinking 配置
if internal.thinking and internal.thinking.enabled: if internal.thinking and internal.thinking.enabled:
@@ -1180,6 +1219,11 @@ class ClaudeNormalizer(FormatNormalizer):
return ToolChoice(type=ToolChoiceType.AUTO, extra={"claude": tool_choice}) 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]:
raw_choice = (
tool_choice.extra.get("claude") if isinstance(tool_choice.extra, dict) else None
)
if isinstance(raw_choice, dict):
return copy.deepcopy(raw_choice)
if tool_choice.type == ToolChoiceType.NONE: if tool_choice.type == ToolChoiceType.NONE:
return {"type": "none"} return {"type": "none"}
if tool_choice.type == ToolChoiceType.AUTO: if tool_choice.type == ToolChoiceType.AUTO:
@@ -1201,6 +1245,13 @@ class ClaudeNormalizer(FormatNormalizer):
blocks: list[dict[str, Any]] = [] blocks: list[dict[str, Any]] = []
text_parts: list[str] = [] text_parts: list[str] = []
def flush_text_parts() -> None:
nonlocal text_parts
if not text_parts:
return
blocks.append({"type": "text", "text": "\n".join(text_parts)})
text_parts = []
for b in msg.content: for b in msg.content:
if isinstance(b, UnknownBlock): if isinstance(b, UnknownBlock):
continue continue
@@ -1208,6 +1259,7 @@ class ClaudeNormalizer(FormatNormalizer):
if isinstance(b, TextBlock): if isinstance(b, TextBlock):
if b.text: if b.text:
if force_structured_text: if force_structured_text:
flush_text_parts()
text_block: dict[str, Any] = {"type": "text", "text": b.text} text_block: dict[str, Any] = {"type": "text", "text": b.text}
cc = b.extra.get("cache_control") if b.extra else None cc = b.extra.get("cache_control") if b.extra else None
if isinstance(cc, dict): if isinstance(cc, dict):
@@ -1225,6 +1277,7 @@ class ClaudeNormalizer(FormatNormalizer):
text_parts.append("[Image]") text_parts.append("[Image]")
continue continue
flush_text_parts()
if b.data and b.media_type: if b.data and b.media_type:
blocks.append( blocks.append(
{ {
@@ -1247,6 +1300,7 @@ class ClaudeNormalizer(FormatNormalizer):
if isinstance(b, FileBlock): if isinstance(b, FileBlock):
if b.data and b.media_type: if b.data and b.media_type:
flush_text_parts()
blocks.append( blocks.append(
{ {
"type": "document", "type": "document",
@@ -1258,6 +1312,7 @@ class ClaudeNormalizer(FormatNormalizer):
} }
) )
elif b.file_url: elif b.file_url:
flush_text_parts()
blocks.append( blocks.append(
{ {
"type": "document", "type": "document",
@@ -1273,6 +1328,7 @@ class ClaudeNormalizer(FormatNormalizer):
if isinstance(b, AudioBlock): if isinstance(b, AudioBlock):
if b.data and b.media_type: if b.data and b.media_type:
flush_text_parts()
blocks.append( blocks.append(
{ {
"type": "document", "type": "document",
@@ -1286,6 +1342,7 @@ class ClaudeNormalizer(FormatNormalizer):
continue continue
if isinstance(b, ToolUseBlock) and role == "assistant": if isinstance(b, ToolUseBlock) and role == "assistant":
flush_text_parts()
blocks.append( blocks.append(
{ {
"type": "tool_use", "type": "tool_use",
@@ -1297,6 +1354,7 @@ class ClaudeNormalizer(FormatNormalizer):
continue continue
if isinstance(b, ToolResultBlock) and role == "user": if isinstance(b, ToolResultBlock) and role == "user":
flush_text_parts()
if b.content_text is not None: if b.content_text is not None:
content: Any = b.content_text content: Any = b.content_text
elif b.output is None: elif b.output is None:
@@ -1316,38 +1374,50 @@ class ClaudeNormalizer(FormatNormalizer):
) )
continue continue
if text_parts: if text_parts and not blocks:
if blocks: return {"role": role, "content": "\n".join(text_parts)}
blocks = [{"type": "text", "text": "\n".join(text_parts)}] + blocks
else:
return {"role": role, "content": "\n".join(text_parts)}
flush_text_parts()
return {"role": role, "content": blocks} return {"role": role, "content": blocks}
def _coerce_claude_message_sequence( def _coerce_claude_message_sequence(
self, messages: list[InternalMessage] self, messages: list[InternalMessage]
) -> list[InternalMessage]: ) -> tuple[list[InternalMessage], dict[str, Any]]:
diagnostics: dict[str, Any] = {
"changed": False,
"prepended_empty_user": False,
"merged_adjacent": 0,
"coerced_roles": 0,
}
normalized: list[InternalMessage] = [] normalized: list[InternalMessage] = []
for m in messages: for m in messages:
role = m.role role = m.role
if role not in (Role.USER, Role.ASSISTANT): if role not in (Role.USER, Role.ASSISTANT):
diagnostics["coerced_roles"] += 1
role = Role.USER role = Role.USER
normalized.append(InternalMessage(role=role, content=m.content, extra=m.extra)) normalized.append(InternalMessage(role=role, content=list(m.content), extra=m.extra))
if not normalized: if not normalized:
return [] return [], diagnostics
if normalized[0].role != Role.USER: if normalized[0].role != Role.USER:
normalized = [InternalMessage(role=Role.USER, content=[])] + normalized normalized = [InternalMessage(role=Role.USER, content=[])] + normalized
diagnostics["prepended_empty_user"] = True
merged: list[InternalMessage] = [] merged: list[InternalMessage] = []
for m in normalized: for m in normalized:
if merged and merged[-1].role == m.role: if merged and merged[-1].role == m.role:
merged[-1].content.extend(m.content) merged[-1].content.extend(m.content)
diagnostics["merged_adjacent"] += 1
continue continue
merged.append(m) merged.append(m)
return merged diagnostics["changed"] = bool(
diagnostics["prepended_empty_user"]
or diagnostics["merged_adjacent"]
or diagnostics["coerced_roles"]
)
return merged, diagnostics
def _claude_usage_to_internal(self, usage: Any) -> UsageInfo | None: def _claude_usage_to_internal(self, usage: Any) -> UsageInfo | None:
if not isinstance(usage, dict): if not isinstance(usage, dict):

View File

@@ -11,6 +11,7 @@ Gemini (GenerateContent / streamGenerateContent) Normalizer
- 响应/流式通常为 camelCasecandidates/finishReason/usageMetadata/modelVersion - 响应/流式通常为 camelCasecandidates/finishReason/usageMetadata/modelVersion
""" """
import copy
import json import json
import re import re
from typing import Any from typing import Any
@@ -67,6 +68,7 @@ from src.core.api_format.conversion.stream_events import (
) )
from src.core.api_format.conversion.stream_state import StreamState from src.core.api_format.conversion.stream_state import StreamState
from src.core.api_format.schema_utils import clean_gemini_schema as _clean_gemini_schema from src.core.api_format.schema_utils import clean_gemini_schema as _clean_gemini_schema
from src.core.logger import logger
# Valid Gemini Part data-oneof field names (camelCase + snake_case). # Valid Gemini Part data-oneof field names (camelCase + snake_case).
_VALID_PART_DATA_FIELDS = frozenset( _VALID_PART_DATA_FIELDS = frozenset(
@@ -106,20 +108,37 @@ def compact_gemini_contents(contents: list[dict[str, Any]]) -> list[dict[str, An
contents with no Gemini-compatible parts, or consecutive same-role entries contents with no Gemini-compatible parts, or consecutive same-role entries
after filtering. after filtering.
""" """
dropped_non_list_parts = 0
dropped_invalid_parts = 0
dropped_empty_contents = 0
merged_same_role = 0
# 1. Strip invalid parts, then drop contents with no valid parts remaining. # 1. Strip invalid parts, then drop contents with no valid parts remaining.
# Use shallow copy to avoid mutating the caller's original dicts. # Use shallow copy to avoid mutating the caller's original dicts.
non_empty: list[dict[str, Any]] = [] non_empty: list[dict[str, Any]] = []
for c in contents: for c in contents:
parts = c.get("parts") parts = c.get("parts")
if not isinstance(parts, list): if not isinstance(parts, list):
dropped_non_list_parts += 1
continue continue
valid_parts = [p for p in parts if _is_valid_gemini_part(p)] valid_parts = [p for p in parts if _is_valid_gemini_part(p)]
dropped_invalid_parts += max(0, len(parts) - len(valid_parts))
if valid_parts: if valid_parts:
c = {**c, "parts": valid_parts} c = {**c, "parts": valid_parts}
non_empty.append(c) non_empty.append(c)
else:
dropped_empty_contents += 1
# 2. Merge consecutive same-role entries. # 2. Merge consecutive same-role entries.
if not non_empty: if not non_empty:
if dropped_non_list_parts or dropped_invalid_parts or dropped_empty_contents:
logger.debug(
"[GeminiNormalizer] compact_gemini_contents dropped all contents: input_count={}, dropped_non_list_parts={}, dropped_invalid_parts={}, dropped_empty_contents={}",
len(contents),
dropped_non_list_parts,
dropped_invalid_parts,
dropped_empty_contents,
)
return non_empty return non_empty
merged: list[dict[str, Any]] = [non_empty[0]] merged: list[dict[str, Any]] = [non_empty[0]]
@@ -130,9 +149,26 @@ def compact_gemini_contents(contents: list[dict[str, Any]]) -> list[dict[str, An
prev_parts.extend(c.get("parts") or []) prev_parts.extend(c.get("parts") or [])
else: else:
merged[-1]["parts"] = list(c.get("parts") or []) merged[-1]["parts"] = list(c.get("parts") or [])
merged_same_role += 1
else: else:
merged.append(c) merged.append(c)
if (
dropped_non_list_parts
or dropped_invalid_parts
or dropped_empty_contents
or merged_same_role
):
logger.debug(
"[GeminiNormalizer] compacted contents: input_count={}, output_count={}, dropped_non_list_parts={}, dropped_invalid_parts={}, dropped_empty_contents={}, merged_same_role={}",
len(contents),
len(merged),
dropped_non_list_parts,
dropped_invalid_parts,
dropped_empty_contents,
merged_same_role,
)
return merged return merged
@@ -453,6 +489,12 @@ class GeminiNormalizer(FormatNormalizer):
# tools/tool_choice — clean unsupported JSON Schema fields from parameters # tools/tool_choice — clean unsupported JSON Schema fields from parameters
# Gemini 特殊内置工具名称需要单独处理 # Gemini 特殊内置工具名称需要单独处理
emitted_cleaned_schemas = 0
reused_raw_tool_choice = (
isinstance(internal.tool_choice.extra.get("gemini"), dict)
if internal.tool_choice
else False
)
tools = None tools = None
if internal.tools: if internal.tools:
func_decls: list[dict[str, Any]] = [] func_decls: list[dict[str, Any]] = []
@@ -467,6 +509,7 @@ class GeminiNormalizer(FormatNormalizer):
params = dict(t.parameters) if t.parameters else {} params = dict(t.parameters) if t.parameters else {}
if params: if params:
_clean_gemini_schema(params) _clean_gemini_schema(params)
emitted_cleaned_schemas += 1
decl: dict[str, Any] = { decl: dict[str, Any] = {
"name": t.name, "name": t.name,
"parameters": params, "parameters": params,
@@ -486,6 +529,26 @@ class GeminiNormalizer(FormatNormalizer):
tool_config = None tool_config = None
if internal.tool_choice: if internal.tool_choice:
tool_config = self._tool_choice_to_gemini_tool_config(internal.tool_choice) tool_config = self._tool_choice_to_gemini_tool_config(internal.tool_choice)
logger.debug(
"[GeminiNormalizer] {} tool_config: variant={}, mode={}",
"reused raw" if reused_raw_tool_choice else "rebuilt",
target_variant_norm or "default",
(
(
tool_config.get("function_calling_config")
or tool_config.get("functionCallingConfig")
or {}
).get("mode")
if isinstance(tool_config, dict)
else None
),
)
if emitted_cleaned_schemas:
logger.debug(
"[GeminiNormalizer] cleaned Gemini schemas: variant={}, count={}",
target_variant_norm or "default",
emitted_cleaned_schemas,
)
generation_config: dict[str, Any] = {} generation_config: dict[str, Any] = {}
if internal.max_tokens is not None: if internal.max_tokens is not None:
@@ -524,6 +587,7 @@ class GeminiNormalizer(FormatNormalizer):
if isinstance(wrapped_schema, dict): if isinstance(wrapped_schema, dict):
schema = dict(wrapped_schema) schema = dict(wrapped_schema)
_clean_gemini_schema(schema) _clean_gemini_schema(schema)
emitted_cleaned_schemas += 1
generation_config["responseSchema"] = schema generation_config["responseSchema"] = schema
elif internal.response_format.type == "json_object": elif internal.response_format.type == "json_object":
generation_config["responseMimeType"] = "application/json" generation_config["responseMimeType"] = "application/json"
@@ -648,6 +712,13 @@ class GeminiNormalizer(FormatNormalizer):
result: dict[str, Any] = { result: dict[str, Any] = {
"contents": contents, "contents": contents,
} }
if len(contents) != len(raw_contents):
logger.debug(
"[GeminiNormalizer] contents count changed after compaction: variant={}, before={}, after={}",
target_variant_norm or "default",
len(raw_contents),
len(contents),
)
# Gemini Chat 模式 model 可能在 URL 路径中;这里仅在 internal.model 存在时回写 # Gemini Chat 模式 model 可能在 URL 路径中;这里仅在 internal.model 存在时回写
if internal.model: if internal.model:
@@ -2113,6 +2184,12 @@ class GeminiNormalizer(FormatNormalizer):
return ToolChoice(type=ToolChoiceType.AUTO, extra={"gemini": tool_config}) 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]:
raw_tool_config = (
tool_choice.extra.get("gemini") if isinstance(tool_choice.extra, dict) else None
)
if isinstance(raw_tool_config, dict):
return copy.deepcopy(raw_tool_config)
mode = "AUTO" mode = "AUTO"
cfg: dict[str, Any] = {} cfg: dict[str, Any] = {}

View File

@@ -25,6 +25,7 @@ from src.core.api_format.conversion.constants import (
from src.core.api_format.conversion.constants import ( from src.core.api_format.conversion.constants import (
responses_web_search_tool_to_chat_options as _responses_web_search_tool_to_chat_options, responses_web_search_tool_to_chat_options as _responses_web_search_tool_to_chat_options,
) )
from src.core.api_format.conversion.constants import stable_json_dumps as _stable_json_dumps
from src.core.api_format.conversion.field_mappings import ( from src.core.api_format.conversion.field_mappings import (
ERROR_TYPE_MAPPINGS, ERROR_TYPE_MAPPINGS,
REASONING_EFFORT_TO_THINKING_BUDGET, REASONING_EFFORT_TO_THINKING_BUDGET,
@@ -355,12 +356,13 @@ class OpenAINormalizer(FormatNormalizer):
openai_tools.append(chat_tool) openai_tools.append(chat_tool)
continue continue
func: dict[str, Any] = { func: dict[str, Any] = {
"name": t.name,
"parameters": t.parameters or {},
**(t.extra.get("openai_function") or {}), **(t.extra.get("openai_function") or {}),
"name": t.name,
} }
if t.description is not None: if t.description is not None:
func["description"] = t.description func["description"] = t.description
if t.parameters is not None:
func["parameters"] = t.parameters
openai_tools.append( openai_tools.append(
{ {
"type": "function", "type": "function",
@@ -1439,6 +1441,7 @@ class OpenAINormalizer(FormatNormalizer):
description=function.get("description"), description=function.get("description"),
parameters=params_raw if isinstance(params_raw, dict) else None, parameters=params_raw if isinstance(params_raw, dict) else None,
extra={ extra={
"openai_chat_raw_tool": tool,
"openai_tool": self._extract_extra(tool, {"type", "function"}), "openai_tool": self._extract_extra(tool, {"type", "function"}),
"openai_function": self._extract_extra( "openai_function": self._extract_extra(
function, {"name", "description", "parameters"} function, {"name", "description", "parameters"}
@@ -1514,25 +1517,32 @@ class OpenAINormalizer(FormatNormalizer):
fn_raw = tool_call.get("function") 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 "") name = str(fn.get("name") or "")
args_str = str(fn.get("arguments") or "") args_raw = fn.get("arguments") if isinstance(fn, dict) else None
tool_id = str(tool_call.get("id") or "") tool_id = str(tool_call.get("id") or "")
tool_input: dict[str, Any] tool_input: dict[str, Any]
if args_str: if isinstance(args_raw, str):
try: if args_raw:
parsed = json.loads(args_str) try:
tool_input = parsed if isinstance(parsed, dict) else {"raw": parsed} parsed = json.loads(args_raw)
except json.JSONDecodeError: tool_input = parsed if isinstance(parsed, dict) else {"raw": parsed}
tool_input = {"raw": args_str} except json.JSONDecodeError:
tool_input = {"raw": args_raw}
else:
tool_input = {}
else: else:
tool_input = {} tool_input = {}
extra: dict[str, Any] = {"openai": tool_call}
if isinstance(args_raw, str):
extra["raw"] = {"arguments": args_raw}
return ( return (
ToolUseBlock( ToolUseBlock(
tool_id=tool_id, tool_id=tool_id,
tool_name=name, tool_name=name,
tool_input=tool_input, tool_input=tool_input,
extra={"openai": tool_call}, extra=extra,
), ),
dropped, dropped,
) )
@@ -1542,7 +1552,7 @@ class OpenAINormalizer(FormatNormalizer):
) -> tuple[ToolUseBlock | None, dict[str, int]]: ) -> tuple[ToolUseBlock | None, dict[str, int]]:
dropped: dict[str, int] = {} dropped: dict[str, int] = {}
name = str(func_call.get("name") or "") name = str(func_call.get("name") or "")
args_str = str(func_call.get("arguments") or "") args_raw = func_call.get("arguments")
if not name: if not name:
dropped["openai_function_call_missing_name"] = ( dropped["openai_function_call_missing_name"] = (
dropped.get("openai_function_call_missing_name", 0) + 1 dropped.get("openai_function_call_missing_name", 0) + 1
@@ -1550,21 +1560,28 @@ class OpenAINormalizer(FormatNormalizer):
return None, dropped return None, dropped
tool_input: dict[str, Any] tool_input: dict[str, Any]
if args_str: if isinstance(args_raw, str):
try: if args_raw:
parsed = json.loads(args_str) try:
tool_input = parsed if isinstance(parsed, dict) else {"raw": parsed} parsed = json.loads(args_raw)
except json.JSONDecodeError: tool_input = parsed if isinstance(parsed, dict) else {"raw": parsed}
tool_input = {"raw": args_str} except json.JSONDecodeError:
tool_input = {"raw": args_raw}
else:
tool_input = {}
else: else:
tool_input = {} tool_input = {}
extra: dict[str, Any] = {"openai": {"function_call": func_call}}
if isinstance(args_raw, str):
extra["raw"] = {"arguments": args_raw}
return ( return (
ToolUseBlock( ToolUseBlock(
tool_id="call_0", tool_id="call_0",
tool_name=name, tool_name=name,
tool_input=tool_input, tool_input=tool_input,
extra={"openai": {"function_call": func_call}}, extra=extra,
), ),
dropped, dropped,
) )
@@ -1591,13 +1608,14 @@ class OpenAINormalizer(FormatNormalizer):
except json.JSONDecodeError: except json.JSONDecodeError:
parsed = None parsed = None
raw_extra = {"raw": {"content": content}, "openai": msg}
if parsed is not None: if parsed is not None:
return ( return (
ToolResultBlock( ToolResultBlock(
tool_use_id=tool_call_id, tool_use_id=tool_call_id,
output=parsed, output=parsed,
content_text=None, content_text=None,
extra={"raw": {"content": content}, "openai": msg}, extra=raw_extra,
), ),
dropped, dropped,
) )
@@ -1607,11 +1625,10 @@ class OpenAINormalizer(FormatNormalizer):
tool_use_id=tool_call_id, tool_use_id=tool_call_id,
output=None, output=None,
content_text=content, content_text=content,
extra={"openai": msg}, extra=raw_extra,
), ),
dropped, dropped,
) )
# 非字符串:尽量保留为 output # 非字符串:尽量保留为 output
return ( return (
ToolResultBlock( ToolResultBlock(
@@ -1858,14 +1875,18 @@ class OpenAINormalizer(FormatNormalizer):
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 content: str
if block.content_text is not None: raw = block.extra.get("raw") if isinstance(block.extra, dict) else None
raw_content = raw.get("content") if isinstance(raw, dict) else None
if isinstance(raw_content, str):
content = raw_content
elif block.content_text is not None:
content = block.content_text content = block.content_text
elif block.output is None: elif block.output is None:
content = "" content = ""
elif isinstance(block.output, str): elif isinstance(block.output, str):
content = block.output content = block.output
else: else:
content = json.dumps(block.output, ensure_ascii=False) content = _stable_json_dumps(block.output)
return { return {
"role": "tool", "role": "tool",
@@ -1875,12 +1896,19 @@ class OpenAINormalizer(FormatNormalizer):
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]:
# index 参数仅用于 fallback id 生成;非流式响应的 tool_calls 数组不应包含 index 字段 # index 参数仅用于 fallback id 生成;非流式响应的 tool_calls 数组不应包含 index 字段
raw = block.extra.get("raw") if isinstance(block.extra, dict) else None
raw_arguments = raw.get("arguments") if isinstance(raw, dict) else None
arguments = (
raw_arguments
if isinstance(raw_arguments, str)
else _stable_json_dumps(block.tool_input or {})
)
return { return {
"id": block.tool_id or f"call_{index}", "id": block.tool_id or f"call_{index}",
"type": "function", "type": "function",
"function": { "function": {
"name": block.tool_name, "name": block.tool_name,
"arguments": json.dumps(block.tool_input or {}, ensure_ascii=False), "arguments": arguments,
}, },
} }

View File

@@ -28,6 +28,7 @@ from src.core.api_format.conversion.constants import (
from src.core.api_format.conversion.constants import ( from src.core.api_format.conversion.constants import (
chat_web_search_options_to_responses_tools as _chat_web_search_options_to_responses_tools, chat_web_search_options_to_responses_tools as _chat_web_search_options_to_responses_tools,
) )
from src.core.api_format.conversion.constants import stable_json_dumps as _stable_json_dumps
from src.core.api_format.conversion.field_mappings import ( from src.core.api_format.conversion.field_mappings import (
ERROR_TYPE_MAPPINGS, ERROR_TYPE_MAPPINGS,
REASONING_EFFORT_TO_THINKING_BUDGET, REASONING_EFFORT_TO_THINKING_BUDGET,
@@ -196,7 +197,11 @@ class OpenAICliNormalizer(FormatNormalizer):
"parallel_tool_calls", "parallel_tool_calls",
"reasoning", "reasoning",
}, },
) ),
"openai_cli_request_flags": {
"has_instructions": "instructions" in request,
"has_stream": "stream" in request,
},
} }
if isinstance(text_config, dict): if isinstance(text_config, dict):
fmt = text_config.get("format") fmt = text_config.get("format")
@@ -246,24 +251,32 @@ class OpenAICliNormalizer(FormatNormalizer):
*, *,
target_variant: str | None = None, target_variant: str | None = None,
) -> dict[str, Any]: ) -> dict[str, Any]:
_ = target_variant
openai_extra = internal.extra.get("openai", {}) openai_extra = internal.extra.get("openai", {})
openai_cli_extra = internal.extra.get("openai_cli", {}) openai_cli_extra = internal.extra.get("openai_cli", {})
request_flags = internal.extra.get("openai_cli_request_flags", {}) if internal.extra else {}
is_codex_variant = (target_variant or "").lower() == "codex"
has_explicit_instructions = bool(
isinstance(request_flags, dict) and request_flags.get("has_instructions")
)
has_explicit_stream = bool(
isinstance(request_flags, dict) and request_flags.get("has_stream")
)
result: dict[str, Any] = { result: dict[str, Any] = {
"model": internal.model, "model": internal.model,
"input": self._internal_messages_to_input(internal.messages, system_to_developer=False), "input": self._internal_messages_to_input(internal.messages, system_to_developer=False),
} }
# 合并 instructions如果没有则使用 system # 合并 instructions如果没有则使用 system
# 仅在 internal 中确有内容或原请求显式提供过时输出,
# 让 Codex 默认 body_rules 仍可在字段缺失时注入默认 instructions。
instructions_text = ( instructions_text = (
self._join_instructions(internal.instructions) self._join_instructions(internal.instructions)
if internal.instructions if internal.instructions
else internal.system else internal.system
) )
# Responses API 兼容 instructions 字段Codex 强制要求 if instructions_text or has_explicit_instructions:
# 统一添加该字段以确保兼容性 result["instructions"] = instructions_text or ""
result["instructions"] = instructions_text or ""
if internal.max_tokens is not None: if internal.max_tokens is not None:
# Responses API 使用 max_output_tokens # Responses API 使用 max_output_tokens
@@ -275,7 +288,8 @@ class OpenAICliNormalizer(FormatNormalizer):
if internal.stop_sequences: if internal.stop_sequences:
result["stop"] = list(internal.stop_sequences) result["stop"] = list(internal.stop_sequences)
result["stream"] = bool(internal.stream) if internal.stream or has_explicit_stream:
result["stream"] = bool(internal.stream)
if internal.tools: if internal.tools:
# Responses API 使用扁平结构: {type, name, description, parameters} # Responses API 使用扁平结构: {type, name, description, parameters}
@@ -291,15 +305,18 @@ class OpenAICliNormalizer(FormatNormalizer):
if translated_tool := _chat_tool_to_responses_tool(chat_tool): if translated_tool := _chat_tool_to_responses_tool(chat_tool):
rebuilt_tools.append(translated_tool) rebuilt_tools.append(translated_tool)
else: else:
rebuilt_tools.append( rebuilt_tool: dict[str, Any] = {
{ "type": "function",
"type": "function", **(t.extra.get("openai_cli") or {}),
"name": t.name, **(t.extra.get("openai_tool") or {}),
"description": t.description or "", **(t.extra.get("openai_function") or {}),
"parameters": t.parameters or {}, "name": t.name,
**(t.extra.get("openai_tool") or {}), }
} if t.description is not None:
) rebuilt_tool["description"] = t.description
if t.parameters is not None:
rebuilt_tool["parameters"] = t.parameters
rebuilt_tools.append(rebuilt_tool)
if rebuilt_tools: if rebuilt_tools:
result["tools"] = rebuilt_tools result["tools"] = rebuilt_tools
@@ -382,8 +399,7 @@ class OpenAICliNormalizer(FormatNormalizer):
): ):
result[key] = value result[key] = value
# 标准 Responses API 默认设置 store=false if is_codex_variant and "store" not in result:
if "store" not in result:
result["store"] = False result["store"] = False
return result return result
@@ -465,17 +481,20 @@ class OpenAICliNormalizer(FormatNormalizer):
for block in internal.content: for block in internal.content:
if isinstance(block, ToolUseBlock): if isinstance(block, ToolUseBlock):
raw = block.extra.get("raw") if isinstance(block.extra, dict) else None
raw_arguments = raw.get("arguments") if isinstance(raw, dict) else None
arguments = (
raw_arguments
if isinstance(raw_arguments, str)
else _stable_json_dumps(block.tool_input or {})
)
output_items.append( output_items.append(
{ {
"type": "function_call", "type": "function_call",
"call_id": block.tool_id, "call_id": block.tool_id,
"id": block.tool_id, "id": block.tool_id,
"name": block.tool_name, "name": block.tool_name,
"arguments": ( "arguments": arguments,
json.dumps(block.tool_input, ensure_ascii=False)
if block.tool_input
else "{}"
),
"status": "completed", "status": "completed",
} }
) )
@@ -1491,7 +1510,9 @@ class OpenAICliNormalizer(FormatNormalizer):
has_tool_use = True has_tool_use = True
tool_id = str(item.get("call_id") or item.get("id") or "") tool_id = str(item.get("call_id") or item.get("id") or "")
tool_name = str(item.get("name") or "") tool_name = str(item.get("name") or "")
args_raw = item.get("arguments") or "{}" args_raw = item.get("arguments")
if args_raw is None:
args_raw = "{}"
try: try:
tool_input = ( tool_input = (
json.loads(args_raw) json.loads(args_raw)
@@ -1500,11 +1521,19 @@ class OpenAICliNormalizer(FormatNormalizer):
) )
except (json.JSONDecodeError, TypeError): except (json.JSONDecodeError, TypeError):
tool_input = {"_raw": args_raw} tool_input = {"_raw": args_raw}
extra: dict[str, Any] = {
"openai_cli": self._extract_extra(
item, {"type", "call_id", "id", "name", "arguments", "status"}
)
}
if isinstance(args_raw, str):
extra["raw"] = {"arguments": args_raw}
blocks.append( blocks.append(
ToolUseBlock( ToolUseBlock(
tool_id=tool_id, tool_id=tool_id,
tool_name=tool_name, tool_name=tool_name,
tool_input=tool_input, tool_input=tool_input,
extra=extra,
) )
) )
continue continue
@@ -1631,36 +1660,65 @@ class OpenAICliNormalizer(FormatNormalizer):
def _parse_function_call_item(self, item: dict[str, Any]) -> InternalMessage: def _parse_function_call_item(self, item: dict[str, Any]) -> InternalMessage:
tool_id = str(item.get("call_id") or item.get("id") or "") tool_id = str(item.get("call_id") or item.get("id") or "")
tool_name = str(item.get("name") or "") tool_name = str(item.get("name") or "")
args_raw = item.get("arguments") or "{}" args_raw = item.get("arguments")
try: if args_raw is None:
tool_input = ( args_raw = "{}"
json.loads(args_raw)
if isinstance(args_raw, str) if isinstance(args_raw, str):
else (args_raw if isinstance(args_raw, dict) else {}) if args_raw:
) try:
except (json.JSONDecodeError, TypeError): parsed = json.loads(args_raw)
tool_input = {"_raw": args_raw} tool_input = parsed if isinstance(parsed, dict) else {"_raw": parsed}
except (json.JSONDecodeError, TypeError):
tool_input = {"_raw": args_raw}
else:
tool_input = {}
elif isinstance(args_raw, dict):
tool_input = args_raw
else:
tool_input = {}
extra: dict[str, Any] = {
"openai_cli": self._extract_extra(item, {"type", "call_id", "id", "name", "arguments"})
}
if isinstance(args_raw, str):
extra["raw"] = {"arguments": args_raw}
tool_block = ToolUseBlock( tool_block = ToolUseBlock(
tool_id=tool_id, tool_id=tool_id,
tool_name=tool_name, tool_name=tool_name,
tool_input=tool_input, tool_input=tool_input,
extra={ extra=extra,
"openai_cli": self._extract_extra(
item, {"type", "call_id", "id", "name", "arguments"}
)
},
) )
return InternalMessage(role=Role.ASSISTANT, content=[tool_block]) return InternalMessage(role=Role.ASSISTANT, content=[tool_block])
def _parse_function_call_output_item(self, item: dict[str, Any]) -> InternalMessage: def _parse_function_call_output_item(self, item: dict[str, Any]) -> InternalMessage:
tool_use_id = str(item.get("call_id") or item.get("id") or "") tool_use_id = str(item.get("call_id") or item.get("id") or "")
output = item.get("output") output = item.get("output")
content_text = output if isinstance(output, str) else None
extra: dict[str, Any] = {
"openai_cli": self._extract_extra(item, {"type", "call_id", "id", "output"})
}
if isinstance(output, str):
extra["raw"] = {"content": output, "output": output}
try:
parsed = json.loads(output)
except (json.JSONDecodeError, TypeError):
parsed = None
result_block = ToolResultBlock(
tool_use_id=tool_use_id,
output=parsed,
content_text=None if parsed is not None else output,
extra=extra,
)
return InternalMessage(role=Role.USER, content=[result_block])
result_block = ToolResultBlock( result_block = ToolResultBlock(
tool_use_id=tool_use_id, tool_use_id=tool_use_id,
output=output, output=output,
content_text=content_text, content_text=None,
extra={"openai_cli": self._extract_extra(item, {"type", "call_id", "id", "output"})}, extra=extra,
) )
return InternalMessage(role=Role.USER, content=[result_block]) return InternalMessage(role=Role.USER, content=[result_block])
@@ -1760,28 +1818,38 @@ class OpenAICliNormalizer(FormatNormalizer):
# ToolUseBlock -> function_call # ToolUseBlock -> function_call
for block in msg.content: for block in msg.content:
if isinstance(block, ToolUseBlock): if isinstance(block, ToolUseBlock):
raw = block.extra.get("raw") if isinstance(block.extra, dict) else None
raw_arguments = raw.get("arguments") if isinstance(raw, dict) else None
arguments = (
raw_arguments
if isinstance(raw_arguments, str)
else _stable_json_dumps(block.tool_input or {})
)
out.append( out.append(
{ {
"type": "function_call", "type": "function_call",
"call_id": block.tool_id, "call_id": block.tool_id,
"name": block.tool_name, "name": block.tool_name,
"arguments": ( "arguments": arguments,
json.dumps(block.tool_input, ensure_ascii=False)
if block.tool_input
else "{}"
),
} }
) )
continue continue
if isinstance(block, ToolResultBlock): if isinstance(block, ToolResultBlock):
# Responses API function_call_output.output 必须是字符串 # Responses API function_call_output.output 必须是字符串
if block.content_text is not None: raw = block.extra.get("raw") if isinstance(block.extra, dict) else None
raw_output = raw.get("output") if isinstance(raw, dict) else None
raw_content = raw.get("content") if isinstance(raw, dict) else None
if isinstance(raw_output, str):
output_str = raw_output
elif isinstance(raw_content, str):
output_str = raw_content
elif block.content_text is not None:
output_str = block.content_text output_str = block.content_text
elif isinstance(block.output, str): elif isinstance(block.output, str):
output_str = block.output output_str = block.output
elif block.output is not None: elif block.output is not None:
output_str = json.dumps(block.output, ensure_ascii=False) output_str = _stable_json_dumps(block.output)
else: else:
output_str = "" output_str = ""
out.append( out.append(
@@ -1881,6 +1949,7 @@ class OpenAICliNormalizer(FormatNormalizer):
fn.get("parameters") if isinstance(fn.get("parameters"), dict) else None fn.get("parameters") if isinstance(fn.get("parameters"), dict) else None
), ),
extra={ extra={
"openai_cli_raw_tool": tool,
"openai_tool": self._extract_extra(tool, {"type", "function"}), "openai_tool": self._extract_extra(tool, {"type", "function"}),
"openai_function": self._extract_extra( "openai_function": self._extract_extra(
fn, {"name", "description", "parameters"} fn, {"name", "description", "parameters"}
@@ -1932,9 +2001,10 @@ class OpenAICliNormalizer(FormatNormalizer):
else None else None
), ),
extra={ extra={
"openai_cli_raw_tool": tool,
"openai_cli": self._extract_extra( "openai_cli": self._extract_extra(
tool, {"name", "description", "parameters"} tool, {"name", "description", "parameters"}
) ),
}, },
) )
) )
@@ -1961,15 +2031,24 @@ class OpenAICliNormalizer(FormatNormalizer):
if isinstance(tool_choice, dict): if isinstance(tool_choice, dict):
# OpenAI 兼容结构:{"type":"function","function":{"name":"..."}} # OpenAI 兼容结构:{"type":"function","function":{"name":"..."}}
if tool_choice.get("type") == "function" and isinstance( if tool_choice.get("type") == "function":
tool_choice.get("function"), dict if isinstance(tool_choice.get("function"), dict):
): name = str(tool_choice["function"].get("name") or "")
name = str(tool_choice["function"].get("name") or "") else:
name = str(tool_choice.get("name") or "")
return ToolChoice( return ToolChoice(
type=ToolChoiceType.TOOL, tool_name=name, extra={"openai_cli": tool_choice} type=ToolChoiceType.TOOL, tool_name=name, extra={"openai_cli": tool_choice}
) )
if tool_choice.get("type") == "custom": if tool_choice.get("type") == "custom":
name = str(tool_choice.get("name") or "") name = str(
tool_choice.get("name")
or (
(tool_choice.get("custom") or {}).get("name")
if isinstance(tool_choice.get("custom"), dict)
else ""
)
or ""
)
return ToolChoice( return ToolChoice(
type=ToolChoiceType.TOOL, tool_name=name, extra={"openai_cli": tool_choice} type=ToolChoiceType.TOOL, tool_name=name, extra={"openai_cli": tool_choice}
) )

View File

@@ -195,11 +195,13 @@ class FormatConversionRegistry:
tgt = self.get_normalizer(target_format) tgt = self.get_normalizer(target_format)
return src is not None and src is tgt return src is not None and src is tgt
def _repair_internal_tool_call_ids(self, internal: InternalRequest) -> None: def _repair_internal_tool_call_ids(self, internal: InternalRequest) -> dict[str, int]:
"""修复 InternalRequest 中空的 tool id/tool_use_id避免上游校验报错。""" """修复 InternalRequest 中空的 tool id/tool_use_id避免上游校验报错。"""
pending_tool_ids: list[str] = [] pending_tool_ids: list[str] = []
auto_counter = 0 auto_counter = 0
generated_tool_use_ids = 0
filled_tool_result_ids = 0
def next_tool_id() -> str: def next_tool_id() -> str:
nonlocal auto_counter nonlocal auto_counter
@@ -213,6 +215,7 @@ class FormatConversionRegistry:
if not tool_id: if not tool_id:
tool_id = next_tool_id() tool_id = next_tool_id()
block.tool_id = tool_id block.tool_id = tool_id
generated_tool_use_ids += 1
pending_tool_ids.append(tool_id) pending_tool_ids.append(tool_id)
continue continue
@@ -224,11 +227,17 @@ class FormatConversionRegistry:
pending_tool_ids.remove(tool_use_id) pending_tool_ids.remove(tool_use_id)
continue continue
filled_tool_result_ids += 1
if pending_tool_ids: if pending_tool_ids:
block.tool_use_id = pending_tool_ids.pop(0) block.tool_use_id = pending_tool_ids.pop(0)
else: else:
block.tool_use_id = next_tool_id() block.tool_use_id = next_tool_id()
return {
"generated_tool_use_ids": generated_tool_use_ids,
"filled_tool_result_ids": filled_tool_result_ids,
}
# ==================== 请求/响应转换(严格) ==================== # ==================== 请求/响应转换(严格) ====================
def convert_request( def convert_request(
@@ -262,7 +271,15 @@ class FormatConversionRegistry:
try: try:
internal = src.request_to_internal(request) internal = src.request_to_internal(request)
internal.output_limit = output_limit internal.output_limit = output_limit
self._repair_internal_tool_call_ids(internal) repair_stats = self._repair_internal_tool_call_ids(internal)
if repair_stats["generated_tool_use_ids"] or repair_stats["filled_tool_result_ids"]:
logger.debug(
"[FormatConversionRegistry] repaired internal tool call ids: source={}, target={}, generated_tool_use_ids={}, filled_tool_result_ids={}",
str(source_format).upper(),
str(target_format).upper(),
repair_stats["generated_tool_use_ids"],
repair_stats["filled_tool_result_ids"],
)
return tgt.request_from_internal(internal, target_variant=target_variant) return tgt.request_from_internal(internal, target_variant=target_variant)
except Exception as e: except Exception as e:
raise FormatConversionError(source_format, target_format, str(e)) from e raise FormatConversionError(source_format, target_format, str(e)) from e
@@ -299,7 +316,15 @@ class FormatConversionRegistry:
try: try:
internal = src.request_to_internal(request) internal = src.request_to_internal(request)
internal.output_limit = output_limit internal.output_limit = output_limit
self._repair_internal_tool_call_ids(internal) repair_stats = self._repair_internal_tool_call_ids(internal)
if repair_stats["generated_tool_use_ids"] or repair_stats["filled_tool_result_ids"]:
logger.debug(
"[FormatConversionRegistry] repaired internal tool call ids: source={}, target={}, generated_tool_use_ids={}, filled_tool_result_ids={}",
str(source_format).upper(),
str(target_format).upper(),
repair_stats["generated_tool_use_ids"],
repair_stats["filled_tool_result_ids"],
)
# 异步阶段:解析图片 URL -> base64仅在目标格式需要时 # 异步阶段:解析图片 URL -> base64仅在目标格式需要时
await resolve_image_urls(internal, str(target_format).upper()) await resolve_image_urls(internal, str(target_format).upper())

View File

@@ -393,7 +393,8 @@ def _merge_all_of(obj: dict[str, Any]) -> None:
return return
merged_props: dict[str, Any] = {} merged_props: dict[str, Any] = {}
merged_required: set[str] = set() merged_required: list[str] = []
merged_required_seen: set[str] = set()
other_fields: dict[str, Any] = {} other_fields: dict[str, Any] = {}
for sub in all_of: for sub in all_of:
@@ -407,8 +408,9 @@ def _merge_all_of(obj: dict[str, Any]) -> None:
r = sub.get("required") r = sub.get("required")
if isinstance(r, list): if isinstance(r, list):
for item in r: for item in r:
if isinstance(item, str): if isinstance(item, str) and item not in merged_required_seen:
merged_required.add(item) merged_required_seen.add(item)
merged_required.append(item)
# 合并其余字段 # 合并其余字段
for k, v in sub.items(): for k, v in sub.items():
if k not in ("properties", "required", "allOf") and k not in other_fields: if k not in ("properties", "required", "allOf") and k not in other_fields:

View File

@@ -15,16 +15,19 @@ from typing import Any, cast
from src.core.api_format.conversion.internal import ( from src.core.api_format.conversion.internal import (
ErrorType, ErrorType,
ImageBlock, InternalMessage,
InternalRequest,
Role,
StopReason, StopReason,
TextBlock, TextBlock,
ToolChoice,
ToolChoiceType,
ToolResultBlock, ToolResultBlock,
ToolUseBlock, ToolUseBlock,
UnknownBlock, UnknownBlock,
) )
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer
from src.core.api_format.conversion.stream_events import ( from src.core.api_format.conversion.stream_events import (
ContentBlockStartEvent,
ContentDeltaEvent, ContentDeltaEvent,
MessageStartEvent, MessageStartEvent,
MessageStopEvent, MessageStopEvent,
@@ -389,3 +392,117 @@ def test_claude_system_array_format() -> None:
assert "You are Claude Code" in internal.system assert "You are Claude Code" in internal.system
assert "Extract file paths" in internal.system assert "Extract file paths" in internal.system
assert "\n\n" in internal.system assert "\n\n" in internal.system
def test_claude_request_reuses_raw_tool_choice_and_stabilizes_message_sequence() -> None:
n = ClaudeNormalizer()
internal = n.request_to_internal(
{
"model": "claude-3-sonnet",
"messages": [{"role": "assistant", "content": [{"type": "text", "text": "hi"}]}],
"tool_choice": {"type": "tool", "name": "read_file", "disable_parallel_tool_use": True},
"max_tokens": 16,
}
)
internal.messages = [
InternalMessage(role=Role.ASSISTANT, content=internal.messages[0].content),
InternalMessage(role=Role.ASSISTANT, content=[TextBlock(text="again")]),
]
internal.tool_choice = ToolChoice(
type=ToolChoiceType.TOOL,
tool_name="read_file",
extra={"claude": {"type": "tool", "name": "read_file", "disable_parallel_tool_use": True}},
)
out = n.request_from_internal(internal)
assert out["tool_choice"] == {
"type": "tool",
"name": "read_file",
"disable_parallel_tool_use": True,
}
assert [m["role"] for m in out["messages"]] == ["user", "assistant"]
assert out["messages"][0]["content"] == []
assert isinstance(out["messages"][1]["content"], str)
assert out["messages"][1]["content"] == "hi\nagain"
def test_claude_request_preserves_tool_and_text_order_when_merging_adjacent_roles() -> None:
n = ClaudeNormalizer()
internal = InternalRequest(
model="claude-3-sonnet",
messages=[
InternalMessage(role=Role.USER, content=[TextBlock(text="weather?")]),
InternalMessage(
role=Role.ASSISTANT,
content=[
ToolUseBlock(
tool_id="toolu_1",
tool_name="get_weather",
tool_input={"city": "SF"},
)
],
),
InternalMessage(role=Role.ASSISTANT, content=[TextBlock(text="Use this result.")]),
InternalMessage(
role=Role.USER,
content=[ToolResultBlock(tool_use_id="toolu_1", output={"temp_c": 20})],
),
InternalMessage(role=Role.USER, content=[TextBlock(text="Received.")]),
],
max_tokens=16,
)
out = n.request_from_internal(internal)
out_messages: list[dict[str, Any]] = out["messages"]
assert [m["role"] for m in out_messages] == ["user", "assistant", "user"]
assistant_blocks = cast(list[dict[str, Any]], out_messages[1]["content"])
assert [block["type"] for block in assistant_blocks] == ["tool_use", "text"]
assert assistant_blocks[1]["text"] == "Use this result."
user_blocks = cast(list[dict[str, Any]], out_messages[2]["content"])
assert [block["type"] for block in user_blocks] == ["tool_result", "text"]
assert user_blocks[0]["content"] == {"temp_c": 20}
assert user_blocks[1]["text"] == "Received."
def test_claude_request_preserves_cache_control_system_block_shape() -> None:
n = ClaudeNormalizer()
internal = n.request_to_internal(
{
"model": "claude-3-sonnet",
"messages": [{"role": "user", "content": "hello"}],
"system": [
{"type": "text", "text": "one", "cache_control": {"type": "ephemeral"}},
{"type": "text", "text": "two"},
],
"max_tokens": 16,
}
)
out = n.request_from_internal(internal)
assert isinstance(out["system"], list)
assert out["system"] == [
{"type": "text", "text": "one", "cache_control": {"type": "ephemeral"}},
{"type": "text", "text": "two"},
]
def test_claude_request_string_system_stays_string_without_cache_control() -> None:
n = ClaudeNormalizer()
internal = n.request_to_internal(
{
"model": "claude-3-sonnet",
"messages": [{"role": "user", "content": "hello"}],
"system": [
{"type": "text", "text": "one"},
{"type": "text", "text": "two"},
],
"max_tokens": 16,
}
)
out = n.request_from_internal(internal)
assert out["system"] == "one\n\ntwo"
assert isinstance(out["system"], str)

View File

@@ -275,6 +275,47 @@ def test_openai_chat_prompt_cache_key_preserved_when_convert_to_openai_cli() ->
assert out["prompt_cache_key"] == "cache-key-123" assert out["prompt_cache_key"] == "cache-key-123"
def test_openai_chat_tool_payload_preserves_raw_strings_when_convert_to_openai_cli() -> None:
reg = _make_registry_with_cli()
openai_chat_req = {
"model": "gpt-5",
"messages": [
{"role": "user", "content": "帮我读取 README"},
{
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "read_file",
"arguments": '{"b":2,"a":1}',
},
}
],
},
{
"role": "tool",
"tool_call_id": "call_1",
"content": '{"z":1,"a":2}',
},
],
}
out = reg.convert_request(openai_chat_req, "openai:chat", "openai:cli")
input_items = cast(list[dict[str, Any]], out.get("input") or [])
function_call = next((i for i in input_items if i.get("type") == "function_call"), {})
function_call_output = next(
(i for i in input_items if i.get("type") == "function_call_output"), {}
)
assert function_call["arguments"] == '{"b":2,"a":1}'
assert function_call_output["output"] == '{"z":1,"a":2}'
def test_openai_chat_text_config_maps_to_openai_cli_text_block() -> None: def test_openai_chat_text_config_maps_to_openai_cli_text_block() -> None:
reg = _make_registry_with_cli() reg = _make_registry_with_cli()
@@ -601,6 +642,44 @@ def test_claude_tool_use_to_openai_cli() -> None:
assert fco_items[0]["output"] == "Hello World" assert fco_items[0]["output"] == "Hello World"
def test_openai_cli_request_omits_implicit_empty_defaults_for_standard_responses() -> None:
normalizer = OpenAICliNormalizer()
internal = normalizer.request_to_internal({"model": "gpt-test", "input": []})
out = normalizer.request_from_internal(internal)
assert "instructions" not in out
assert "stream" not in out
assert "store" not in out
def test_openai_cli_request_preserves_explicit_empty_instructions_and_stream_false() -> None:
normalizer = OpenAICliNormalizer()
internal = normalizer.request_to_internal(
{"model": "gpt-test", "input": [], "instructions": "", "stream": False}
)
out = normalizer.request_from_internal(internal)
assert out["instructions"] == ""
assert out["stream"] is False
assert "store" not in out
def test_openai_cli_request_tool_choice_flat_function_roundtrip_preserved() -> None:
normalizer = OpenAICliNormalizer()
request = {
"model": "gpt-test",
"input": [],
"tool_choice": {"type": "function", "name": "read_file"},
}
internal = normalizer.request_to_internal(request)
out = normalizer.request_from_internal(internal)
assert out["tool_choice"] == {"type": "function", "name": "read_file"}
def test_claude_explicit_effort_preserved_in_openai_cli() -> None: def test_claude_explicit_effort_preserved_in_openai_cli() -> None:
reg = _make_registry_with_cli() reg = _make_registry_with_cli()

View File

@@ -23,7 +23,10 @@ from src.core.api_format.conversion.internal import (
ToolUseBlock, ToolUseBlock,
UnknownBlock, UnknownBlock,
) )
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer from src.core.api_format.conversion.normalizers.gemini import (
GeminiNormalizer,
compact_gemini_contents,
)
from src.core.api_format.conversion.stream_events import ( from src.core.api_format.conversion.stream_events import (
ContentDeltaEvent, ContentDeltaEvent,
MessageStartEvent, MessageStartEvent,
@@ -31,6 +34,7 @@ from src.core.api_format.conversion.stream_events import (
ToolCallDeltaEvent, ToolCallDeltaEvent,
) )
from src.core.api_format.conversion.stream_state import StreamState from src.core.api_format.conversion.stream_state import StreamState
from src.core.api_format.schema_utils import clean_gemini_schema
def test_gemini_request_system_and_generation_config_roundtrip() -> None: def test_gemini_request_system_and_generation_config_roundtrip() -> None:
@@ -85,7 +89,7 @@ def test_gemini_request_system_and_generation_config_roundtrip() -> None:
assert out["generation_config"]["max_output_tokens"] == 10 assert out["generation_config"]["max_output_tokens"] == 10
assert out["generation_config"]["stop_sequences"] == ["A", "B"] assert out["generation_config"]["stop_sequences"] == ["A", "B"]
assert out["tools"][0]["function_declarations"][0]["name"] == "get_weather" assert out["tools"][0]["function_declarations"][0]["name"] == "get_weather"
assert out["tool_config"]["function_calling_config"]["mode"] == "ANY" assert out["tool_config"] == {"functionCallingConfig": {"mode": "ANY"}}
def test_gemini_request_parts_image_tool_and_unknown_drop() -> None: def test_gemini_request_parts_image_tool_and_unknown_drop() -> None:
@@ -279,3 +283,81 @@ def test_gemini_error_conversion() -> None:
out = n.error_from_internal(internal) out = n.error_from_internal(internal)
assert out["error"]["status"] == "RESOURCE_EXHAUSTED" assert out["error"]["status"] == "RESOURCE_EXHAUSTED"
assert out["error"]["message"] == "slow down" assert out["error"]["message"] == "slow down"
def test_gemini_compact_contents_drops_invalid_and_merges_same_role() -> None:
contents: list[dict[str, Any]] = [
{"role": "user", "parts": [{"text": "hi"}, {"invalid": True}]},
{"role": "user", "parts": [{"text": "again"}]},
{"role": "model", "parts": [{"bad": 1}]},
{"role": "model", "parts": [{"text": "ok"}]},
{"role": "user", "parts": "bad"},
]
compacted = compact_gemini_contents(contents)
assert compacted == [
{"role": "user", "parts": [{"text": "hi"}, {"text": "again"}]},
{"role": "model", "parts": [{"text": "ok"}]},
]
def test_gemini_request_reuses_raw_tool_config_and_cleans_schema_deterministically() -> None:
n = GeminiNormalizer()
req = {
"model": "gemini-1.5",
"contents": [{"role": "user", "parts": [{"text": "hi"}]}],
"tools": [
{
"functionDeclarations": [
{
"name": "read_file",
"description": "Read file",
"parameters": {
"allOf": [
{
"type": "object",
"properties": {"b": {"type": "string"}},
"required": ["b", "a"],
},
{
"type": "object",
"properties": {"a": {"type": "string"}},
"required": ["a", "c"],
},
]
},
}
]
}
],
"toolConfig": {
"functionCallingConfig": {
"mode": "ANY",
"allowedFunctionNames": ["read_file"],
}
},
}
internal = n.request_to_internal(req)
out = n.request_from_internal(internal)
assert out["tool_config"] == {
"functionCallingConfig": {
"mode": "ANY",
"allowedFunctionNames": ["read_file"],
}
}
params = out["tools"][0]["function_declarations"][0]["parameters"]
assert params["required"] == ["b", "a"]
assert list(params["properties"].keys()) == ["b", "a"]
def test_clean_gemini_schema_allof_required_order_is_deterministic() -> None:
schema = {
"allOf": [
{"type": "object", "properties": {"z": {"type": "string"}}, "required": ["z", "a"]},
{"type": "object", "properties": {"a": {"type": "string"}}, "required": ["a", "m"]},
]
}
clean_gemini_schema(schema)
assert schema["required"] == ["z", "a"]

View File

@@ -15,7 +15,6 @@ import json
from typing import Any, cast from typing import Any, cast
from src.core.api_format.conversion.internal import ( from src.core.api_format.conversion.internal import (
ContentType,
ErrorType, ErrorType,
ImageBlock, ImageBlock,
StopReason, StopReason,
@@ -184,6 +183,60 @@ def test_openai_request_tool_calls_and_tool_role_roundtrip() -> None:
assert json.loads(tool_out["content"]) == {"temp_c": 20, "unit": "C"} assert json.loads(tool_out["content"]) == {"temp_c": 20, "unit": "C"}
def test_openai_request_preserves_empty_string_tool_call_arguments() -> None:
n = OpenAINormalizer()
req = {
"model": "gpt-4o-mini",
"messages": [
{"role": "user", "content": "ping"},
{
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_empty",
"type": "function",
"function": {"name": "noop", "arguments": ""},
}
],
},
],
}
internal = n.request_to_internal(req)
tool_use = next(b for b in internal.messages[1].content if isinstance(b, ToolUseBlock))
assert tool_use.tool_input == {}
assert tool_use.extra["raw"]["arguments"] == ""
out = n.request_from_internal(internal)
assert out["messages"][1]["tool_calls"][0]["function"]["arguments"] == ""
def test_openai_request_preserves_empty_string_legacy_function_call_arguments() -> None:
n = OpenAINormalizer()
req = {
"model": "gpt-4o-mini",
"messages": [
{"role": "user", "content": "ping"},
{
"role": "assistant",
"content": "",
"function_call": {"name": "noop", "arguments": ""},
},
],
}
internal = n.request_to_internal(req)
tool_use = next(b for b in internal.messages[1].content if isinstance(b, ToolUseBlock))
assert tool_use.tool_input == {}
assert tool_use.extra["raw"]["arguments"] == ""
out = n.request_from_internal(internal)
assert out["messages"][1]["tool_calls"][0]["function"]["arguments"] == ""
def test_openai_request_content_image_and_unknown_drop() -> None: def test_openai_request_content_image_and_unknown_drop() -> None:
n = OpenAINormalizer() n = OpenAINormalizer()

View File

@@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
from types import SimpleNamespace from types import SimpleNamespace
from typing import cast
import jwt import jwt
import pytest import pytest
@@ -12,6 +13,7 @@ from src.services.provider.adapters.codex.request_patching import (
maybe_patch_request_for_codex, maybe_patch_request_for_codex,
patch_openai_cli_request_for_codex, patch_openai_cli_request_for_codex,
) )
from src.services.provider.envelope import ProviderEnvelope
def test_patch_openai_cli_request_for_codex_is_passthrough_except_internal_sentinel() -> None: def test_patch_openai_cli_request_for_codex_is_passthrough_except_internal_sentinel() -> None:
@@ -146,6 +148,21 @@ def test_openai_cli_normalizer_request_from_internal_codex_variant_defaults_stor
assert out["store"] is False assert out["store"] is False
def test_openai_cli_normalizer_codex_variant_keeps_instructions_missing_for_default_rule() -> None:
from src.api.handlers.base.request_builder import apply_body_rules
from src.core.api_format.conversion.normalizers.openai_cli import OpenAICliNormalizer
from src.core.api_format.metadata import CODEX_DEFAULT_BODY_RULES
normalizer = OpenAICliNormalizer()
internal = normalizer.request_to_internal({"model": "gpt-test", "input": []})
out = normalizer.request_from_internal(internal, target_variant="codex")
assert "instructions" not in out
patched = apply_body_rules(out, list(CODEX_DEFAULT_BODY_RULES))
assert patched["instructions"] == "You are GPT-5."
def test_codex_envelope_extra_headers_does_not_inject_synthetic_headers() -> None: def test_codex_envelope_extra_headers_does_not_inject_synthetic_headers() -> None:
from src.services.provider.adapters.codex.envelope import codex_oauth_envelope from src.services.provider.adapters.codex.envelope import codex_oauth_envelope
@@ -233,7 +250,7 @@ def test_codex_passthrough_builder_preserves_real_codex_headers() -> None:
endpoint=endpoint, endpoint=endpoint,
key=key, key=key,
pre_computed_auth=("Authorization", "Bearer upstream-token"), pre_computed_auth=("Authorization", "Bearer upstream-token"),
envelope=codex_oauth_envelope, envelope=cast(ProviderEnvelope, codex_oauth_envelope),
) )
assert headers["accept"] == "text/event-stream" assert headers["accept"] == "text/event-stream"

View File

@@ -1,6 +1,9 @@
from typing import Any from typing import Any
from src.api.handlers.base.request_builder import apply_body_rules from src.api.handlers.base.request_builder import (
apply_body_rules,
get_cache_sensitive_protected_body_keys,
)
class TestApplyBodyRulesNestedPaths: class TestApplyBodyRulesNestedPaths:
@@ -47,7 +50,7 @@ class TestApplyBodyRulesNestedPaths:
assert result["extra"]["model"] == "y" # extra 不受保护 assert result["extra"]["model"] == "y" # extra 不受保护
def test_escaped_dot(self) -> None: def test_escaped_dot(self) -> None:
body = {} body: dict[str, Any] = {}
result = apply_body_rules( result = apply_body_rules(
body, body,
[ [
@@ -80,7 +83,7 @@ class TestApplyBodyRulesNestedPaths:
assert result == {"a": {"b": 2}} assert result == {"a": {"b": 2}}
def test_set_complex_value(self) -> None: def test_set_complex_value(self) -> None:
body = {} body: dict[str, Any] = {}
result = apply_body_rules( result = apply_body_rules(
body, body,
[ [
@@ -976,7 +979,7 @@ class TestConditionalBodyRules:
assert result["feature"] is True assert result["feature"] is True
def test_condition_on_missing_nested_path(self) -> None: def test_condition_on_missing_nested_path(self) -> None:
body = {"config": {}} body: dict[str, Any] = {"config": {}}
result = apply_body_rules( result = apply_body_rules(
body, body,
[ [
@@ -1511,3 +1514,117 @@ class TestItemCondition:
original_body=original_body, original_body=original_body,
) )
assert result["matched"] is True assert result["matched"] is True
class TestProtectedBodyKeys:
def test_get_cache_sensitive_protected_body_keys_by_format(self) -> None:
assert get_cache_sensitive_protected_body_keys("openai:chat") == frozenset(
{"model", "stream", "messages", "tools", "tool_choice"}
)
assert get_cache_sensitive_protected_body_keys("openai:cli") == frozenset(
{
"model",
"stream",
"input",
"instructions",
"tools",
"tool_choice",
"prompt_cache_key",
}
)
assert get_cache_sensitive_protected_body_keys("claude:chat") == frozenset(
{"model", "stream", "messages", "system", "tools", "tool_choice"}
)
assert get_cache_sensitive_protected_body_keys("gemini:chat") == frozenset(
{
"model",
"stream",
"contents",
"system_instruction",
"systemInstruction",
"tools",
"tool_config",
"toolConfig",
"generation_config",
"generationConfig",
}
)
def test_gemini_camelcase_alias_prompt_fields_are_protected(self) -> None:
body = {
"contents": [{"role": "user", "parts": [{"text": "hi"}]}],
"systemInstruction": {"parts": [{"text": "system"}]},
"toolConfig": {"functionCallingConfig": {"mode": "AUTO"}},
"generationConfig": {"temperature": 0.1},
"metadata": {"safe": True},
}
protected_keys = get_cache_sensitive_protected_body_keys("gemini:chat")
result = apply_body_rules(
body,
[
{"action": "set", "path": "systemInstruction.parts[0].text", "value": "mutated"},
{"action": "drop", "path": "toolConfig"},
{"action": "rename", "from": "generationConfig", "to": "generation_config"},
{"action": "append", "path": "contents", "value": {"role": "model", "parts": []}},
{
"action": "insert",
"path": "contents",
"index": 0,
"value": {"role": "user", "parts": [{"text": "preface"}]},
},
{"action": "set", "path": "metadata.safe", "value": False},
],
protected_keys=protected_keys,
)
assert result["contents"] == [{"role": "user", "parts": [{"text": "hi"}]}]
assert result["systemInstruction"] == {"parts": [{"text": "system"}]}
assert result["toolConfig"] == {"functionCallingConfig": {"mode": "AUTO"}}
assert result["generationConfig"] == {"temperature": 0.1}
assert "generation_config" not in result
assert result["metadata"]["safe"] is False
def test_protected_prompt_fields_block_all_mutating_actions(self) -> None:
body = {
"messages": [{"role": "user", "content": "hi"}],
"tools": [{"name": "ReadFile"}],
"tool_choice": {"type": "function", "function": {"name": "ReadFile"}},
"metadata": {"safe": True},
}
protected_keys = get_cache_sensitive_protected_body_keys("openai:chat")
result = apply_body_rules(
body,
[
{"action": "set", "path": "messages", "value": []},
{"action": "drop", "path": "tools"},
{"action": "rename", "from": "tool_choice", "to": "choice"},
{
"action": "append",
"path": "messages",
"value": {"role": "assistant", "content": "x"},
},
{
"action": "insert",
"path": "messages",
"index": 0,
"value": {"role": "system", "content": "x"},
},
{
"action": "regex_replace",
"path": "tools[0].name",
"pattern": "Read",
"replacement": "Write",
},
{"action": "name_style", "path": "tools[*].name", "style": "snake_case"},
{"action": "set", "path": "metadata.safe", "value": False},
],
protected_keys=protected_keys,
)
assert result["messages"] == [{"role": "user", "content": "hi"}]
assert result["tools"] == [{"name": "ReadFile"}]
assert result["tool_choice"] == {"type": "function", "function": {"name": "ReadFile"}}
assert "choice" not in result
assert result["metadata"]["safe"] is False

View File

@@ -1,6 +1,10 @@
from typing import Any
import pytest import pytest
from src.api.handlers.base.request_builder import get_cache_sensitive_protected_body_keys
from src.api.handlers.claude.adapter import ClaudeChatAdapter from src.api.handlers.claude.adapter import ClaudeChatAdapter
from src.api.handlers.gemini.adapter import GeminiChatAdapter
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -23,3 +27,117 @@ def test_validate_test_base_url_trims_whitespace() -> None:
assert ClaudeChatAdapter._validate_test_base_url(" https://api.anthropic.com/v1 ") == ( assert ClaudeChatAdapter._validate_test_base_url(" https://api.anthropic.com/v1 ") == (
"https://api.anthropic.com/v1" "https://api.anthropic.com/v1"
) )
@pytest.mark.asyncio
async def test_claude_check_endpoint_passes_cache_sensitive_protected_keys_to_body_rules(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from src.api.handlers.base import endpoint_checker as endpoint_checker_module
from src.api.handlers.base import request_builder as request_builder_module
captured: dict[str, Any] = {}
def fake_apply_body_rules(
body: dict[str, Any],
body_rules: list[dict[str, Any]],
protected_keys: frozenset[str] | None = None,
original_body: dict[str, Any] | None = None,
) -> dict[str, Any]:
captured["body_rules"] = body_rules
captured["protected_keys"] = protected_keys
captured["original_body"] = original_body
return body
async def fake_run_endpoint_check(**kwargs: Any) -> dict[str, Any]:
captured["json_body"] = kwargs["json_body"]
return {"status_code": 200, "json_body": kwargs["json_body"]}
def fake_build_request_body(
cls: type[ClaudeChatAdapter],
request_data: dict[str, Any] | None = None,
*,
base_url: str | None = None,
provider_type: str | None = None,
) -> dict[str, Any]:
del request_data, base_url, provider_type
return {
"messages": [{"role": "user", "content": "hello"}],
"system": "keep",
}
monkeypatch.setattr(request_builder_module, "apply_body_rules", fake_apply_body_rules)
monkeypatch.setattr(endpoint_checker_module, "run_endpoint_check", fake_run_endpoint_check)
monkeypatch.setattr(
ClaudeChatAdapter, "build_request_body", classmethod(fake_build_request_body)
)
result = await ClaudeChatAdapter.check_endpoint(
client=None, # type: ignore[arg-type]
base_url="https://api.anthropic.com/v1",
api_key="test-key",
request_data={"model": "claude-sonnet-4-5-20250929", "stream": False},
body_rules=[{"action": "set", "path": "messages", "value": []}],
)
assert captured["protected_keys"] == get_cache_sensitive_protected_body_keys(
ClaudeChatAdapter.FORMAT_ID
)
assert result["status_code"] == 200
assert result["json_body"] == captured["json_body"]
@pytest.mark.asyncio
async def test_gemini_check_endpoint_passes_alias_aware_protected_keys_to_body_rules(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from src.api.handlers.base import endpoint_checker as endpoint_checker_module
from src.api.handlers.base import request_builder as request_builder_module
captured: dict[str, Any] = {}
def fake_apply_body_rules(
body: dict[str, Any],
body_rules: list[dict[str, Any]],
protected_keys: frozenset[str] | None = None,
original_body: dict[str, Any] | None = None,
) -> dict[str, Any]:
captured["body_rules"] = body_rules
captured["protected_keys"] = protected_keys
captured["original_body"] = original_body
return body
async def fake_run_endpoint_check(**kwargs: Any) -> dict[str, Any]:
captured["json_body"] = kwargs["json_body"]
return {"status_code": 200, "json_body": kwargs["json_body"]}
def fake_build_request_body(
cls: type[GeminiChatAdapter], request_data: dict[str, Any] | None = None
) -> dict[str, Any]:
_ = cls, request_data
return {
"contents": [{"role": "user", "parts": [{"text": "hello"}]}],
"systemInstruction": {"parts": [{"text": "system"}]},
"toolConfig": {"functionCallingConfig": {"mode": "AUTO"}},
"generationConfig": {"temperature": 0.1},
}
monkeypatch.setattr(request_builder_module, "apply_body_rules", fake_apply_body_rules)
monkeypatch.setattr(endpoint_checker_module, "run_endpoint_check", fake_run_endpoint_check)
monkeypatch.setattr(
GeminiChatAdapter, "build_request_body", classmethod(fake_build_request_body)
)
result = await GeminiChatAdapter.check_endpoint(
client=None, # type: ignore[arg-type]
base_url="https://generativelanguage.googleapis.com",
api_key="test-key",
request_data={"model": "gemini-2.5-pro", "stream": False},
body_rules=[{"action": "drop", "path": "toolConfig"}],
)
protected_keys = captured["protected_keys"]
assert protected_keys == get_cache_sensitive_protected_body_keys(GeminiChatAdapter.FORMAT_ID)
assert {"systemInstruction", "toolConfig", "generationConfig"}.issubset(protected_keys)
assert result["status_code"] == 200
assert result["json_body"] == captured["json_body"]

View File

@@ -1,6 +1,9 @@
import json import json
from src.api.handlers.base.request_builder import evaluate_condition from src.api.handlers.base.request_builder import (
evaluate_condition,
summarize_request_payload_shape,
)
from src.core.api_format import ( from src.core.api_format import (
CORE_REDACT_HEADERS, CORE_REDACT_HEADERS,
HeaderBuilder, HeaderBuilder,
@@ -284,3 +287,54 @@ class TestAuthHeaderCasePreservation:
assert "authorization" in headers assert "authorization" in headers
assert "Authorization" not in headers assert "Authorization" not in headers
assert headers["authorization"] == "Bearer provider-token" assert headers["authorization"] == "Bearer provider-token"
class TestRequestPayloadSummary:
def test_prefers_runtime_provider_api_format_over_endpoint_static_format(self) -> None:
payload = {"contents": [{"role": "user", "parts": [{"text": "hi"}]}]}
summary = summarize_request_payload_shape(
payload,
provider_api_format="gemini:chat",
protected_body_keys=frozenset({"contents", "toolConfig"}),
body_rules=[{"action": "drop", "path": "toolConfig"}],
)
assert summary["format"] == "gemini:chat"
assert summary["contents_count"] == 1
assert summary["message_count"] is None
assert summary["protected_body_keys_enabled"] is True
assert summary["protected_body_keys"] == ["contents", "toolConfig"]
assert summary["body_rule_count"] == 1
assert summary["top_level_keys"] == ["contents"]
def test_detects_gemini_camelcase_prompt_bearing_fields(self) -> None:
payload = {
"contents": [{"role": "user", "parts": [{"text": "hi"}]}],
"systemInstruction": {"parts": [{"text": "system"}]},
"toolConfig": {"functionCallingConfig": {"mode": "AUTO"}},
"generationConfig": {"temperature": 0.1},
}
summary = summarize_request_payload_shape(
payload,
provider_api_format="gemini:chat",
protected_body_keys=None,
body_rules=None,
)
assert summary["has_system"] is True
assert summary["has_tool_choice"] is True
assert summary["has_generation_config"] is True
assert summary["function_declaration_count"] == 0
assert summary["tool_count"] is None
assert summary["format"] == "gemini:chat"
assert summary["contents_count"] == 1
assert summary["protected_body_keys_enabled"] is False
assert summary["body_rule_count"] == 0
assert summary["top_level_keys"] == [
"contents",
"generationConfig",
"systemInstruction",
"toolConfig",
]