mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
refactor: 重构格式转换系统为 Hub-and-Spoke Normalizer 架构
- 移除旧的 converters 模块,改用基于 Internal 中间表示的 Normalizer 模式 - 新增 internal.py 定义统一的内部数据结构(InternalRequest/Response/Error) - 新增 normalizer.py 定义 FormatNormalizer 基类接口 - 实现 OpenAI/Claude/Gemini 及其 CLI 格式的 Normalizer - 新增 stream_events.py 和 stream_state.py 支持流式转换 - 重构 registry.py 使用 Normalizer 进行格式转换 - 更新 handlers 适配新的转换架构 - 前端移除 CLI 格式转换按钮的限制 - 添加 golden data 测试确保转换正确性
This commit is contained in:
@@ -36,9 +36,8 @@
|
||||
</Badge>
|
||||
</div>
|
||||
<div class="flex items-center gap-1.5">
|
||||
<!-- 格式转换按钮(非 CLI 格式才显示) -->
|
||||
<!-- 格式转换按钮 -->
|
||||
<Button
|
||||
v-if="!endpoint.api_format.endsWith('_CLI')"
|
||||
variant="ghost"
|
||||
size="icon"
|
||||
class="h-7 w-7 mr-1"
|
||||
|
||||
@@ -22,7 +22,7 @@ Chat Handler Base - Chat API 格式的通用基类
|
||||
import asyncio
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, AsyncGenerator, Callable, Dict, Optional
|
||||
from typing import Any, AsyncGenerator, Callable, Dict, Optional, Union
|
||||
|
||||
import httpx
|
||||
from fastapi import BackgroundTasks, Request
|
||||
@@ -36,7 +36,11 @@ from src.api.handlers.base.response_parser import ResponseParser
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.api.handlers.base.stream_processor import StreamProcessor
|
||||
from src.api.handlers.base.stream_telemetry import StreamTelemetryRecorder
|
||||
from src.api.handlers.base.utils import build_sse_headers, filter_proxy_response_headers
|
||||
from src.api.handlers.base.utils import (
|
||||
build_sse_headers,
|
||||
filter_proxy_response_headers,
|
||||
get_format_converter_registry,
|
||||
)
|
||||
from src.config.settings import config
|
||||
from src.core.error_utils import extract_client_error_message
|
||||
from src.core.exceptions import (
|
||||
@@ -46,6 +50,7 @@ from src.core.exceptions import (
|
||||
ProviderRateLimitException,
|
||||
ProviderTimeoutException,
|
||||
ThinkingSignatureException,
|
||||
UpstreamClientException,
|
||||
)
|
||||
from src.core.logger import logger
|
||||
from src.models.database import (
|
||||
@@ -59,6 +64,91 @@ from src.services.cache.aware_scheduler import ProviderCandidate
|
||||
from src.services.provider.transport import build_provider_url, redact_url_for_log
|
||||
|
||||
|
||||
def _get_error_status_code(e: Exception, default: int = 400) -> int:
|
||||
"""从异常中提取 HTTP 状态码"""
|
||||
code = getattr(e, "status_code", None)
|
||||
return code if isinstance(code, int) and code > 0 else default
|
||||
|
||||
|
||||
def _convert_error_response_best_effort(
|
||||
error_response: Dict[str, Any],
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
将上游错误响应 best-effort 转换为客户端格式。
|
||||
|
||||
说明:错误转换走 Canonical registry。转换失败时构造安全的通用错误响应,
|
||||
避免泄露上游原始错误详情。
|
||||
"""
|
||||
try:
|
||||
registry = get_format_converter_registry()
|
||||
return registry.convert_error_response(error_response, source_format, target_format)
|
||||
except Exception as e:
|
||||
logger.debug(f"错误响应转换失败 ({source_format} -> {target_format}): {e}")
|
||||
# 转换失败时构造安全的通用错误,避免泄露上游详情
|
||||
return _build_client_error_response_best_effort("upstream error", target_format)
|
||||
|
||||
|
||||
def _build_client_error_response_best_effort(
|
||||
message: str,
|
||||
target_format: str,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
当无法解析上游错误 body 时,构造一个目标格式的错误响应(best-effort)。
|
||||
"""
|
||||
try:
|
||||
from src.core.api_format.conversion.internal import ErrorType, InternalError
|
||||
|
||||
registry = get_format_converter_registry()
|
||||
normalizer = registry.get_normalizer(target_format)
|
||||
if normalizer and normalizer.capabilities.supports_error_conversion:
|
||||
return normalizer.error_from_internal(
|
||||
InternalError(type=ErrorType.INVALID_REQUEST, message=message, retryable=False)
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"构建客户端错误响应失败 (target={target_format}): {e}")
|
||||
|
||||
return {"error": {"type": "upstream_client_error", "message": message}}
|
||||
|
||||
|
||||
def _build_error_json_payload(
|
||||
e: Union[ThinkingSignatureException, UpstreamClientException],
|
||||
client_format: str,
|
||||
provider_format: str,
|
||||
needs_conversion: bool = True,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
构建错误 JSON 响应 payload(公共逻辑)。
|
||||
|
||||
从异常中提取上游错误信息,尝试转换为客户端格式。
|
||||
|
||||
Args:
|
||||
e: ThinkingSignatureException 或 UpstreamClientException
|
||||
client_format: 客户端 API 格式
|
||||
provider_format: Provider API 格式
|
||||
needs_conversion: 是否需要格式转换
|
||||
|
||||
Returns:
|
||||
格式化的错误响应字典
|
||||
"""
|
||||
raw = getattr(e, "upstream_error", None)
|
||||
message = getattr(e, "message", str(e))
|
||||
|
||||
if isinstance(raw, str) and raw:
|
||||
try:
|
||||
parsed = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
parsed = None
|
||||
|
||||
if isinstance(parsed, dict):
|
||||
if needs_conversion:
|
||||
return _convert_error_response_best_effort(parsed, provider_format, client_format)
|
||||
return parsed
|
||||
|
||||
return _build_client_error_response_best_effort(message, client_format)
|
||||
|
||||
|
||||
class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
"""
|
||||
Chat Handler 基类
|
||||
@@ -361,7 +451,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
original_headers: Dict[str, Any],
|
||||
original_request_body: Dict[str, Any],
|
||||
query_params: Optional[Dict[str, str]] = None,
|
||||
) -> StreamingResponse:
|
||||
) -> Union[StreamingResponse, JSONResponse]:
|
||||
"""处理流式响应"""
|
||||
logger.debug(f"开始流式响应处理 ({self.FORMAT_ID})")
|
||||
|
||||
@@ -486,12 +576,21 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
background=background_tasks,
|
||||
)
|
||||
|
||||
except ThinkingSignatureException as e:
|
||||
# Thinking 签名错误:orchestrator 层已处理整流重试但仍失败
|
||||
# 记录 original_request_body(客户端原始请求),便于排查问题根因
|
||||
self._log_request_error("流式请求失败(签名错误)", e)
|
||||
except (ThinkingSignatureException, UpstreamClientException) as e:
|
||||
# ThinkingSignatureException: orchestrator 层已处理整流重试但仍失败
|
||||
# UpstreamClientException: 上游客户端错误(HTTP 4xx),不重试,直接返回给客户端
|
||||
error_type = "签名错误" if isinstance(e, ThinkingSignatureException) else "上游客户端错误"
|
||||
self._log_request_error(f"流式请求失败({error_type})", e)
|
||||
await self._record_stream_failure(ctx, e, original_headers, original_request_body)
|
||||
raise
|
||||
client_format = (ctx.client_api_format or "").upper()
|
||||
provider_format = (ctx.provider_api_format or client_format).upper()
|
||||
payload = _build_error_json_payload(
|
||||
e, client_format, provider_format, needs_conversion=ctx.needs_conversion
|
||||
)
|
||||
return JSONResponse(
|
||||
status_code=_get_error_status_code(e),
|
||||
content=payload,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
self._log_request_error("流式请求失败", e)
|
||||
@@ -550,11 +649,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
else:
|
||||
request_body = dict(original_request_body)
|
||||
|
||||
# 跨格式:先做请求体转换(严格模式,失败触发 failover)
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
if needs_conversion:
|
||||
from src.core.api_format import converter_registry
|
||||
|
||||
request_body = converter_registry.convert_request_strict(
|
||||
registry = get_format_converter_registry()
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
str(client_api_format),
|
||||
str(provider_api_format),
|
||||
@@ -726,6 +824,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
status_code = 503
|
||||
if isinstance(error, ThinkingSignatureException):
|
||||
status_code = 400
|
||||
elif isinstance(error, UpstreamClientException):
|
||||
status_code = _get_error_status_code(error)
|
||||
elif isinstance(error, ProviderAuthException):
|
||||
status_code = 503
|
||||
elif isinstance(error, ProviderRateLimitException):
|
||||
@@ -783,6 +883,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
response_headers: Dict[str, str] = {}
|
||||
provider_request_headers: Dict[str, str] = {}
|
||||
provider_request_body: Optional[Dict[str, Any]] = None
|
||||
provider_api_format_for_error: Optional[str] = None
|
||||
client_api_format_for_error: Optional[str] = None
|
||||
needs_conversion_for_error: bool = False
|
||||
provider_id: Optional[str] = None # Provider ID(用于失败记录)
|
||||
endpoint_id: Optional[str] = None # Endpoint ID(用于失败记录)
|
||||
key_id: Optional[str] = None # Key ID(用于失败记录)
|
||||
@@ -796,6 +899,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
) -> Dict[str, Any]:
|
||||
nonlocal provider_name, response_json, status_code, response_headers
|
||||
nonlocal provider_request_headers, provider_request_body, mapped_model_result
|
||||
nonlocal provider_api_format_for_error, client_api_format_for_error, needs_conversion_for_error
|
||||
|
||||
provider_name = str(provider.name)
|
||||
provider_api_format = str(endpoint.api_format or api_format)
|
||||
@@ -804,6 +908,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
api_format.value if hasattr(api_format, "value") else str(api_format)
|
||||
)
|
||||
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
|
||||
provider_api_format_for_error = provider_api_format
|
||||
client_api_format_for_error = client_api_format
|
||||
needs_conversion_for_error = needs_conversion
|
||||
|
||||
# 获取模型映射(优先使用映射匹配到的模型,其次是 Provider 级别的映射)
|
||||
mapped_model = candidate.mapping_matched_model if candidate else None
|
||||
@@ -821,11 +928,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
else:
|
||||
request_body = dict(request_body_ref["body"])
|
||||
|
||||
# 跨格式:先做请求体转换(严格模式,失败触发 failover)
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
if needs_conversion:
|
||||
from src.core.api_format import converter_registry
|
||||
|
||||
request_body = converter_registry.convert_request_strict(
|
||||
registry = get_format_converter_registry()
|
||||
request_body = registry.convert_request(
|
||||
request_body,
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
@@ -898,46 +1004,18 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
status_code = resp.status_code
|
||||
response_headers = dict(resp.headers)
|
||||
|
||||
if resp.status_code == 401:
|
||||
raise ProviderAuthException(str(provider.name))
|
||||
elif resp.status_code == 429:
|
||||
raise ProviderRateLimitException(
|
||||
"请求过于频繁,请稍后重试",
|
||||
provider_name=str(provider.name),
|
||||
response_headers=response_headers,
|
||||
)
|
||||
elif resp.status_code >= 500:
|
||||
# 记录响应体以便调试
|
||||
# 统一使用 HTTPStatusError,让 orchestrator/error_classifier 负责分类(客户端错误/兼容性错误/限流等)
|
||||
try:
|
||||
resp.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
error_body = ""
|
||||
try:
|
||||
error_body = resp.text[:1000]
|
||||
logger.error(
|
||||
f" [{self.request_id}] 上游返回5xx错误: status={resp.status_code}, body={error_body[:500]}"
|
||||
)
|
||||
error_body = resp.text[:4000] if resp.text else ""
|
||||
except Exception:
|
||||
pass
|
||||
raise ProviderNotAvailableException(
|
||||
f"上游服务暂时不可用 (HTTP {resp.status_code})",
|
||||
provider_name=str(provider.name),
|
||||
upstream_status=resp.status_code,
|
||||
upstream_response=error_body,
|
||||
)
|
||||
elif resp.status_code != 200:
|
||||
# 记录非200响应以便调试
|
||||
error_body = ""
|
||||
try:
|
||||
error_body = resp.text[:1000]
|
||||
logger.warning(
|
||||
f" [{self.request_id}] 上游返回非200: status={resp.status_code}, body={error_body[:500]}"
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
raise ProviderNotAvailableException(
|
||||
f"上游服务返回错误 (HTTP {resp.status_code})",
|
||||
provider_name=str(provider.name),
|
||||
upstream_status=resp.status_code,
|
||||
upstream_response=error_body,
|
||||
)
|
||||
error_body = ""
|
||||
# 供 ErrorClassifier 优先读取
|
||||
e.upstream_response = error_body # type: ignore[attr-defined]
|
||||
raise
|
||||
|
||||
# 安全解析 JSON 响应,处理可能的编码错误
|
||||
try:
|
||||
@@ -986,11 +1064,10 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
error_status=parsed.error_type,
|
||||
)
|
||||
|
||||
# 跨格式:响应转换回 client_format(严格模式,失败触发 failover)
|
||||
# 跨格式:响应转换回 client_format(失败触发 failover)
|
||||
if needs_conversion and isinstance(response_json, dict):
|
||||
from src.core.api_format import converter_registry
|
||||
|
||||
response_json = converter_registry.convert_response_strict(
|
||||
registry = get_format_converter_registry()
|
||||
response_json = registry.convert_response(
|
||||
response_json,
|
||||
provider_api_format,
|
||||
client_api_format,
|
||||
@@ -1100,7 +1177,43 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
error_message=str(e),
|
||||
is_stream=False,
|
||||
)
|
||||
raise
|
||||
client_format = (client_api_format_for_error or "").upper()
|
||||
provider_format = (provider_api_format_for_error or client_format).upper()
|
||||
payload = _build_error_json_payload(
|
||||
e, client_format, provider_format, needs_conversion=needs_conversion_for_error
|
||||
)
|
||||
return JSONResponse(
|
||||
status_code=_get_error_status_code(e),
|
||||
content=payload,
|
||||
)
|
||||
|
||||
except UpstreamClientException as e:
|
||||
response_time_ms = self.elapsed_ms()
|
||||
actual_request_body = provider_request_body or original_request_body
|
||||
await self.telemetry.record_failure(
|
||||
provider=provider_name or "unknown",
|
||||
model=model,
|
||||
response_time_ms=response_time_ms,
|
||||
status_code=_get_error_status_code(e),
|
||||
request_headers=original_headers,
|
||||
request_body=actual_request_body,
|
||||
error_message=str(e),
|
||||
is_stream=False,
|
||||
api_format=api_format,
|
||||
provider_request_headers=provider_request_headers,
|
||||
response_headers=response_headers,
|
||||
client_response_headers={"content-type": "application/json"},
|
||||
target_model=mapped_model_result,
|
||||
)
|
||||
client_format = (client_api_format_for_error or "").upper()
|
||||
provider_format = (provider_api_format_for_error or client_format).upper()
|
||||
payload = _build_error_json_payload(
|
||||
e, client_format, provider_format, needs_conversion=needs_conversion_for_error
|
||||
)
|
||||
return JSONResponse(
|
||||
status_code=_get_error_status_code(e),
|
||||
content=payload,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
response_time_ms = self.elapsed_ms()
|
||||
|
||||
@@ -15,12 +15,14 @@ import codecs
|
||||
import json
|
||||
import time
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
Callable,
|
||||
Dict,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
)
|
||||
|
||||
import httpx
|
||||
@@ -28,6 +30,9 @@ from fastapi import BackgroundTasks
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.core.api_format import ApiFormatDefinition
|
||||
|
||||
from src.api.handlers.base.base_handler import (
|
||||
BaseMessageHandler,
|
||||
MessageTelemetry,
|
||||
@@ -46,6 +51,7 @@ from src.api.handlers.base.utils import (
|
||||
check_html_response,
|
||||
check_prefetched_response_error,
|
||||
filter_proxy_response_headers,
|
||||
get_format_converter_registry,
|
||||
)
|
||||
from src.config.constants import StreamDefaults
|
||||
from src.config.settings import config
|
||||
@@ -73,6 +79,108 @@ from src.utils.sse_parser import SSEEventParser
|
||||
from src.utils.timeout import read_first_chunk_with_ttfb_timeout
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# SSE 行解析辅助函数
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
def _parse_sse_data_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
"""
|
||||
解析标准 SSE data 行
|
||||
|
||||
Args:
|
||||
line: 以 "data:" 开头的 SSE 行
|
||||
|
||||
Returns:
|
||||
(parsed_json, status) 元组:
|
||||
- (parsed_dict, "ok") - 解析成功
|
||||
- (None, "empty") - 内容为空
|
||||
- (None, "invalid") - JSON 解析失败,调用方应透传原始行
|
||||
"""
|
||||
data_content = line[5:].strip()
|
||||
if not data_content:
|
||||
return None, "empty"
|
||||
try:
|
||||
return json.loads(data_content), "ok"
|
||||
except json.JSONDecodeError:
|
||||
return None, "invalid"
|
||||
|
||||
|
||||
def _parse_sse_event_data_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
"""
|
||||
解析 event + data 同行格式(如 "event: xxx data: {...}")
|
||||
|
||||
Args:
|
||||
line: 以 "event:" 开头且包含 " data:" 的 SSE 行
|
||||
|
||||
Returns:
|
||||
(parsed_json, status) 元组
|
||||
"""
|
||||
_event_part, data_part = line.split(" data:", 1)
|
||||
data_content = data_part.strip()
|
||||
try:
|
||||
return json.loads(data_content), "ok"
|
||||
except json.JSONDecodeError:
|
||||
return None, "invalid"
|
||||
|
||||
|
||||
def _parse_gemini_json_array_line(line: str) -> Tuple[Optional[Any], str]:
|
||||
"""
|
||||
解析 Gemini JSON-array 格式的裸 JSON 行
|
||||
|
||||
Gemini 流式响应可能是 JSON 数组格式,每行是数组元素。
|
||||
|
||||
Args:
|
||||
line: 原始行(可能是 "[", "]", ",", 或 JSON 对象)
|
||||
|
||||
Returns:
|
||||
(parsed_json, status) 元组
|
||||
"""
|
||||
stripped = line.strip()
|
||||
if stripped in ("", "[", "]", ","):
|
||||
return None, "skip"
|
||||
|
||||
candidate = stripped.lstrip(",").rstrip(",").strip()
|
||||
try:
|
||||
return json.loads(candidate), "ok"
|
||||
except json.JSONDecodeError:
|
||||
logger.debug(f"Gemini JSON-array line skip: {stripped[:50]}")
|
||||
return None, "invalid"
|
||||
|
||||
|
||||
def _format_converted_events_to_sse(
|
||||
converted_events: List[Dict[str, Any]],
|
||||
client_format: str,
|
||||
) -> List[str]:
|
||||
"""
|
||||
将转换后的事件格式化为 SSE 行
|
||||
|
||||
Args:
|
||||
converted_events: 转换后的事件列表
|
||||
client_format: 客户端 API 格式
|
||||
|
||||
Returns:
|
||||
SSE 行列表(每个元素是完整的 SSE 事件,包含尾部空行)
|
||||
"""
|
||||
result: List[str] = []
|
||||
needs_event_line = client_format.upper() in ("CLAUDE", "CLAUDE_CLI")
|
||||
|
||||
for evt in converted_events:
|
||||
payload = json.dumps(evt, ensure_ascii=False)
|
||||
if needs_event_line:
|
||||
evt_type = evt.get("type") if isinstance(evt, dict) else None
|
||||
if isinstance(evt_type, str) and evt_type:
|
||||
# Claude 格式:event + data + 空行
|
||||
result.append(f"event: {evt_type}\ndata: {payload}\n")
|
||||
else:
|
||||
result.append(f"data: {payload}\n")
|
||||
else:
|
||||
# OpenAI 格式:data + 空行
|
||||
result.append(f"data: {payload}\n")
|
||||
|
||||
return result
|
||||
|
||||
|
||||
class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"""
|
||||
CLI 格式消息处理器基类
|
||||
@@ -248,6 +356,111 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
"""
|
||||
return request_body
|
||||
|
||||
@staticmethod
|
||||
def _get_format_metadata(format_id: str) -> Optional["ApiFormatDefinition"]:
|
||||
"""获取格式元数据(解析失败返回 None)"""
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.api_format.metadata import API_FORMAT_DEFINITIONS
|
||||
|
||||
try:
|
||||
fmt = APIFormat(format_id.upper())
|
||||
return API_FORMAT_DEFINITIONS.get(fmt)
|
||||
except (ValueError, KeyError):
|
||||
return None
|
||||
|
||||
def _finalize_converted_request(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
client_api_format: str,
|
||||
provider_api_format: str,
|
||||
mapped_model: Optional[str],
|
||||
fallback_model: str,
|
||||
is_stream: bool,
|
||||
) -> None:
|
||||
"""
|
||||
跨格式转换后统一设置并清理 model/stream 字段(原地修改)
|
||||
|
||||
处理逻辑:
|
||||
1. 根据目标格式决定是否在 body 中设置 model
|
||||
2. 若客户端格式不含 stream 字段但 Provider 需要,则显式设置
|
||||
3. 移除目标格式不允许在 body 中携带的字段(如 Gemini 的 model/stream)
|
||||
|
||||
Args:
|
||||
request_body: 转换后的请求体(会被原地修改)
|
||||
client_api_format: 客户端 API 格式
|
||||
provider_api_format: Provider API 格式
|
||||
mapped_model: 映射后的模型名
|
||||
fallback_model: 备用模型名
|
||||
is_stream: 是否流式请求
|
||||
"""
|
||||
client_meta = self._get_format_metadata(client_api_format)
|
||||
provider_meta = self._get_format_metadata(provider_api_format)
|
||||
|
||||
# 默认:model_in_body=True, stream_in_body=True(如 OpenAI/Claude)
|
||||
client_uses_stream = client_meta.stream_in_body if client_meta else True
|
||||
provider_model_in_body = provider_meta.model_in_body if provider_meta else True
|
||||
provider_stream_in_body = provider_meta.stream_in_body if provider_meta else True
|
||||
|
||||
# 设置 model(仅当 Provider 允许且 body 中需要)
|
||||
if provider_model_in_body:
|
||||
request_body["model"] = mapped_model or fallback_model
|
||||
else:
|
||||
request_body.pop("model", None)
|
||||
|
||||
# 设置 stream(客户端不带但 Provider 需要时显式设置;Provider 不需要时移除)
|
||||
if provider_stream_in_body:
|
||||
if not client_uses_stream:
|
||||
request_body["stream"] = is_stream
|
||||
else:
|
||||
request_body.pop("stream", None)
|
||||
|
||||
def _convert_request_for_cross_format(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
client_api_format: str,
|
||||
provider_api_format: str,
|
||||
mapped_model: Optional[str],
|
||||
fallback_model: str,
|
||||
is_stream: bool,
|
||||
) -> Tuple[Dict[str, Any], str]:
|
||||
"""
|
||||
跨格式请求转换的公共逻辑
|
||||
|
||||
将客户端格式的请求体转换为 Provider 格式,并处理 model/stream 字段的补齐和清理。
|
||||
|
||||
Args:
|
||||
request_body: 原始请求体(会被修改)
|
||||
client_api_format: 客户端 API 格式
|
||||
provider_api_format: Provider API 格式
|
||||
mapped_model: 映射后的模型名
|
||||
fallback_model: 备用模型名(通常是原始请求的 model)
|
||||
is_stream: 是否流式请求
|
||||
|
||||
Returns:
|
||||
(转换后的请求体, 用于 URL 的模型名)
|
||||
"""
|
||||
registry = get_format_converter_registry()
|
||||
converted_body = registry.convert_request(
|
||||
request_body,
|
||||
str(client_api_format),
|
||||
str(provider_api_format),
|
||||
)
|
||||
|
||||
# 先计算 URL 模型(在清理 body 中的 model 字段之前)
|
||||
url_model = self.get_model_for_url(converted_body, mapped_model) or mapped_model or fallback_model
|
||||
|
||||
# 统一设置并清理 model/stream 字段
|
||||
self._finalize_converted_request(
|
||||
converted_body,
|
||||
str(client_api_format),
|
||||
str(provider_api_format),
|
||||
mapped_model,
|
||||
fallback_model,
|
||||
is_stream,
|
||||
)
|
||||
|
||||
return converted_body, url_model
|
||||
|
||||
def get_model_for_url(
|
||||
self,
|
||||
request_body: Dict[str, Any],
|
||||
@@ -469,8 +682,29 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
else:
|
||||
request_body = original_request_body
|
||||
|
||||
# 准备发送给 Provider 的请求体(子类可覆盖以移除不需要的字段)
|
||||
request_body = self.prepare_provider_request_body(request_body)
|
||||
client_api_format = (
|
||||
ctx.client_api_format.value
|
||||
if hasattr(ctx.client_api_format, "value")
|
||||
else str(ctx.client_api_format)
|
||||
)
|
||||
provider_api_format = str(ctx.provider_api_format or "")
|
||||
needs_conversion = bool(getattr(candidate, "needs_conversion", False)) if candidate else False
|
||||
ctx.needs_conversion = needs_conversion
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
if needs_conversion and provider_api_format:
|
||||
request_body, url_model = self._convert_request_for_cross_format(
|
||||
request_body,
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
mapped_model,
|
||||
ctx.model,
|
||||
is_stream=True,
|
||||
)
|
||||
else:
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖)
|
||||
request_body = self.prepare_provider_request_body(request_body)
|
||||
url_model = self.get_model_for_url(request_body, mapped_model) or mapped_model or ctx.model
|
||||
|
||||
# 使用 RequestBuilder 构建请求体和请求头
|
||||
# 注意:mapped_model 已经应用到 request_body,这里不再传递
|
||||
@@ -487,10 +721,6 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx.provider_request_headers = provider_headers
|
||||
ctx.provider_request_body = provider_payload
|
||||
|
||||
# 获取用于 URL 的模型名(子类可覆盖此方法,如 Gemini 需要特殊处理)
|
||||
# 使用 ctx.model 作为 fallback(它已从 path_params 获取)
|
||||
url_model = self.get_model_for_url(request_body, mapped_model) or ctx.model
|
||||
|
||||
url = build_provider_url(
|
||||
endpoint,
|
||||
query_params=query_params,
|
||||
@@ -658,6 +888,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx,
|
||||
event.get("event"),
|
||||
event.get("data") or "",
|
||||
record_chunk=not needs_conversion,
|
||||
)
|
||||
self._mark_first_output(ctx, output_state)
|
||||
yield b"\n"
|
||||
@@ -689,7 +920,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
# 格式转换或直接透传
|
||||
if needs_conversion:
|
||||
converted_lines = self._convert_sse_line(ctx, line, events)
|
||||
converted_lines, converted_events = self._convert_sse_line(ctx, line, events)
|
||||
# 记录转换后的数据到 parsed_chunks
|
||||
self._record_converted_chunks(ctx, converted_events)
|
||||
for converted_line in converted_lines:
|
||||
if converted_line:
|
||||
self._mark_first_output(ctx, output_state)
|
||||
@@ -703,6 +936,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx,
|
||||
event.get("event"),
|
||||
event.get("data") or "",
|
||||
record_chunk=not needs_conversion,
|
||||
)
|
||||
|
||||
if ctx.data_count > 0:
|
||||
@@ -714,6 +948,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx,
|
||||
event.get("event"),
|
||||
event.get("data") or "",
|
||||
record_chunk=not needs_conversion,
|
||||
)
|
||||
|
||||
# 检查是否收到数据
|
||||
@@ -735,6 +970,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
else:
|
||||
logger.debug("流式数据转发完成")
|
||||
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
|
||||
client_fmt = (ctx.client_api_format or "").upper()
|
||||
if needs_conversion and client_fmt == "OPENAI":
|
||||
yield b"data: [DONE]\n\n"
|
||||
|
||||
except GeneratorExit:
|
||||
raise
|
||||
@@ -1011,6 +1250,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx,
|
||||
event.get("event"),
|
||||
event.get("data") or "",
|
||||
record_chunk=not needs_conversion,
|
||||
)
|
||||
self._mark_first_output(ctx, output_state)
|
||||
yield b"\n"
|
||||
@@ -1020,7 +1260,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
# 格式转换或直接透传
|
||||
if needs_conversion:
|
||||
converted_lines = self._convert_sse_line(ctx, line, events)
|
||||
converted_lines, converted_events = self._convert_sse_line(ctx, line, events)
|
||||
# 记录转换后的数据到 parsed_chunks
|
||||
self._record_converted_chunks(ctx, converted_events)
|
||||
for converted_line in converted_lines:
|
||||
if converted_line:
|
||||
self._mark_first_output(ctx, output_state)
|
||||
@@ -1034,6 +1276,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx,
|
||||
event.get("event"),
|
||||
event.get("data") or "",
|
||||
record_chunk=not needs_conversion,
|
||||
)
|
||||
|
||||
if ctx.data_count > 0:
|
||||
@@ -1064,6 +1307,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx,
|
||||
event.get("event"),
|
||||
event.get("data") or "",
|
||||
record_chunk=not needs_conversion,
|
||||
)
|
||||
self._mark_first_output(ctx, output_state)
|
||||
yield b"\n"
|
||||
@@ -1095,7 +1339,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
# 格式转换或直接透传
|
||||
if needs_conversion:
|
||||
converted_lines = self._convert_sse_line(ctx, line, events)
|
||||
converted_lines, converted_events = self._convert_sse_line(ctx, line, events)
|
||||
# 记录转换后的数据到 parsed_chunks
|
||||
self._record_converted_chunks(ctx, converted_events)
|
||||
for converted_line in converted_lines:
|
||||
if converted_line:
|
||||
self._mark_first_output(ctx, output_state)
|
||||
@@ -1109,6 +1355,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx,
|
||||
event.get("event"),
|
||||
event.get("data") or "",
|
||||
record_chunk=not needs_conversion,
|
||||
)
|
||||
|
||||
if ctx.data_count > 0:
|
||||
@@ -1121,6 +1368,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx,
|
||||
event.get("event"),
|
||||
event.get("data") or "",
|
||||
record_chunk=not needs_conversion,
|
||||
)
|
||||
|
||||
# 检查是否收到数据
|
||||
@@ -1144,6 +1392,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
yield f"event: error\ndata: {json.dumps(error_event)}\n\n".encode("utf-8")
|
||||
else:
|
||||
logger.debug("流式数据转发完成")
|
||||
# 为 OpenAI 客户端补齐 [DONE] 标记(非 CLI 格式)
|
||||
client_fmt = (ctx.client_api_format or "").upper()
|
||||
if needs_conversion and client_fmt == "OPENAI":
|
||||
yield b"data: [DONE]\n\n"
|
||||
|
||||
except GeneratorExit:
|
||||
raise
|
||||
@@ -1189,12 +1441,21 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx: StreamContext,
|
||||
event_name: Optional[str],
|
||||
data_str: str,
|
||||
record_chunk: bool = False,
|
||||
) -> None:
|
||||
"""
|
||||
处理 SSE 事件
|
||||
|
||||
通用框架:解析 JSON、更新计数器
|
||||
子类可覆盖 _process_event_data() 实现格式特定逻辑
|
||||
|
||||
Args:
|
||||
ctx: 流上下文
|
||||
event_name: 事件名称(如 message_start, content_block_delta 等)
|
||||
data_str: 事件数据字符串(JSON 格式)
|
||||
record_chunk: 是否记录到 parsed_chunks(不需要格式转换时应为 True)
|
||||
当为 True 时,同时更新 data_count;
|
||||
当为 False 时,data_count 由 _record_converted_chunks 更新
|
||||
"""
|
||||
if not data_str:
|
||||
return
|
||||
@@ -1208,8 +1469,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
except json.JSONDecodeError:
|
||||
return
|
||||
|
||||
ctx.data_count += 1
|
||||
ctx.parsed_chunks.append(data)
|
||||
# 当不需要格式转换时,记录原始数据到 parsed_chunks 并更新 data_count
|
||||
# 当需要格式转换时(record_chunk=False),data_count 由 _record_converted_chunks 更新
|
||||
if record_chunk and isinstance(data, dict):
|
||||
ctx.parsed_chunks.append(data)
|
||||
ctx.data_count += 1
|
||||
|
||||
if not isinstance(data, dict):
|
||||
return
|
||||
@@ -1239,12 +1503,28 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx.response_id = data["id"]
|
||||
|
||||
# 使用解析器提取 usage
|
||||
# Claude/CLI 流式响应的 usage 可能在首个 chunk 或最后一个 chunk 中
|
||||
# 首个 chunk 可能部分为 0,最后一个 chunk 包含完整值,因此取最大值确保正确计费
|
||||
usage = self.parser.extract_usage_from_response(data)
|
||||
if usage and not ctx.final_usage:
|
||||
ctx.input_tokens = usage.get("input_tokens", 0)
|
||||
ctx.output_tokens = usage.get("output_tokens", 0)
|
||||
ctx.cached_tokens = usage.get("cache_read_tokens", 0)
|
||||
ctx.final_usage = usage
|
||||
if usage:
|
||||
new_input = usage.get("input_tokens", 0)
|
||||
new_output = usage.get("output_tokens", 0)
|
||||
new_cached = usage.get("cache_read_tokens", 0)
|
||||
new_cache_creation = usage.get("cache_creation_tokens", 0)
|
||||
|
||||
# 取最大值更新
|
||||
if new_input > ctx.input_tokens:
|
||||
ctx.input_tokens = new_input
|
||||
if new_output > ctx.output_tokens:
|
||||
ctx.output_tokens = new_output
|
||||
if new_cached > ctx.cached_tokens:
|
||||
ctx.cached_tokens = new_cached
|
||||
if new_cache_creation > ctx.cache_creation_tokens:
|
||||
ctx.cache_creation_tokens = new_cache_creation
|
||||
|
||||
# 保存最后一个非空 usage 作为 final_usage
|
||||
if any([new_input, new_output, new_cached, new_cache_creation]):
|
||||
ctx.final_usage = usage
|
||||
|
||||
# 提取文本内容
|
||||
text = self.parser.extract_text_content(data)
|
||||
@@ -1258,6 +1538,41 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
if isinstance(response_obj, dict):
|
||||
ctx.final_response = response_obj
|
||||
|
||||
def _record_converted_chunks(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
converted_events: List[Dict[str, Any]],
|
||||
) -> None:
|
||||
"""
|
||||
记录转换后的 chunk 数据到 parsed_chunks,并更新统计信息
|
||||
|
||||
当需要格式转换时,记录的是转换后的数据(客户端实际收到的格式);
|
||||
同时更新 data_count、has_completion 等统计信息。
|
||||
|
||||
Args:
|
||||
ctx: 流上下文
|
||||
converted_events: 转换后的事件列表
|
||||
"""
|
||||
for evt in converted_events:
|
||||
if isinstance(evt, dict):
|
||||
ctx.parsed_chunks.append(evt)
|
||||
ctx.data_count += 1
|
||||
|
||||
# 检测完成事件(根据客户端格式判断)
|
||||
# OpenAI 格式: choices[].finish_reason
|
||||
# Claude 格式: type == "message_stop" 或 stop_reason
|
||||
event_type = evt.get("type", "")
|
||||
if event_type == "message_stop":
|
||||
ctx.has_completion = True
|
||||
elif event_type == "response.completed":
|
||||
ctx.has_completion = True
|
||||
elif "choices" in evt:
|
||||
choices = evt.get("choices", [])
|
||||
for choice in choices:
|
||||
if isinstance(choice, dict) and choice.get("finish_reason"):
|
||||
ctx.has_completion = True
|
||||
break
|
||||
|
||||
def _finalize_stream_metadata(self, ctx: StreamContext) -> None:
|
||||
"""
|
||||
在记录统计前从 parsed_chunks 中提取额外的元数据 - 子类可覆盖
|
||||
@@ -1634,8 +1949,25 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
else:
|
||||
request_body = dict(request_body_ref["body"])
|
||||
|
||||
# 准备发送给 Provider 的请求体(子类可覆盖以移除不需要的字段)
|
||||
request_body = self.prepare_provider_request_body(request_body)
|
||||
client_api_format = (
|
||||
api_format.value if hasattr(api_format, "value") else str(api_format)
|
||||
)
|
||||
needs_conversion = bool(getattr(candidate, "needs_conversion", False))
|
||||
|
||||
# 跨格式:先做请求体转换(失败触发 failover)
|
||||
if needs_conversion and provider_api_format:
|
||||
request_body, url_model = self._convert_request_for_cross_format(
|
||||
request_body,
|
||||
client_api_format,
|
||||
provider_api_format,
|
||||
mapped_model,
|
||||
model,
|
||||
is_stream=False,
|
||||
)
|
||||
else:
|
||||
# 同格式:按原逻辑做轻量清理(子类可覆盖)
|
||||
request_body = self.prepare_provider_request_body(request_body)
|
||||
url_model = self.get_model_for_url(request_body, mapped_model) or mapped_model or model
|
||||
|
||||
# 使用 RequestBuilder 构建请求体和请求头
|
||||
# 注意:mapped_model 已经应用到 request_body,这里不再传递
|
||||
@@ -1652,9 +1984,6 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
provider_request_headers = provider_headers
|
||||
provider_request_body = provider_payload
|
||||
|
||||
# 获取用于 URL 的模型名(子类可覆盖此方法,如 Gemini 需要特殊处理)
|
||||
url_model = self.get_model_for_url(request_body, mapped_model)
|
||||
|
||||
url = build_provider_url(
|
||||
endpoint,
|
||||
query_params=query_params,
|
||||
@@ -1795,16 +2124,15 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
if response_json is None:
|
||||
response_json = {}
|
||||
|
||||
# 检查是否需要格式转换
|
||||
if (
|
||||
provider_api_format
|
||||
and api_format
|
||||
and provider_api_format.upper() != api_format.upper()
|
||||
):
|
||||
from src.core.api_format import converter_registry
|
||||
# 检查是否需要格式转换(同族格式无需转换,如 CLAUDE 和 CLAUDE_CLI)
|
||||
from src.core.api_format.utils import get_base_format
|
||||
|
||||
provider_base = get_base_format(provider_api_format) if provider_api_format else None
|
||||
client_base = get_base_format(api_format) if api_format else None
|
||||
if provider_base and client_base and provider_base != client_base:
|
||||
try:
|
||||
response_json = converter_registry.convert_response(
|
||||
registry = get_format_converter_registry()
|
||||
response_json = registry.convert_response(
|
||||
response_json,
|
||||
provider_api_format,
|
||||
api_format,
|
||||
@@ -1954,10 +2282,30 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
当 Provider 的 API 格式与客户端请求的 API 格式不同时,需要转换响应。
|
||||
例如:客户端请求 Claude 格式,但 Provider 返回 OpenAI 格式。
|
||||
|
||||
注意:CLAUDE 和 CLAUDE_CLI、OPENAI 和 OPENAI_CLI 等同族格式在响应层面是兼容的,
|
||||
因此使用 get_base_format 比较基础格式(去除 _CLI 后缀)。
|
||||
"""
|
||||
from src.core.api_format.utils import get_base_format
|
||||
|
||||
if not ctx.provider_api_format or not ctx.client_api_format:
|
||||
logger.debug(
|
||||
f"[{getattr(ctx, 'request_id', 'unknown')}] _needs_format_conversion: "
|
||||
f"provider_api_format={ctx.provider_api_format!r}, client_api_format={ctx.client_api_format!r} -> False (missing)"
|
||||
)
|
||||
return False
|
||||
return ctx.provider_api_format.upper() != ctx.client_api_format.upper()
|
||||
|
||||
# 比较基础格式(CLAUDE_CLI -> CLAUDE, OPENAI_CLI -> OPENAI)
|
||||
provider_base = get_base_format(ctx.provider_api_format)
|
||||
client_base = get_base_format(ctx.client_api_format)
|
||||
result = provider_base != client_base
|
||||
|
||||
logger.debug(
|
||||
f"[{getattr(ctx, 'request_id', 'unknown')}] _needs_format_conversion: "
|
||||
f"provider={ctx.provider_api_format}(base={provider_base}), "
|
||||
f"client={ctx.client_api_format}(base={client_base}) -> {result}"
|
||||
)
|
||||
return result
|
||||
|
||||
def _mark_first_output(self, ctx: StreamContext, state: Dict[str, bool]) -> None:
|
||||
"""
|
||||
@@ -1983,7 +2331,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx: StreamContext,
|
||||
line: str,
|
||||
events: list, # noqa: ARG002 - 预留给上下文感知转换
|
||||
) -> List[str]:
|
||||
) -> Tuple[List[str], List[Dict[str, Any]]]:
|
||||
"""
|
||||
将 SSE 行从 Provider 格式转换为客户端格式
|
||||
|
||||
@@ -1993,76 +2341,118 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
events: 当前累积的事件列表(预留参数,用于未来上下文感知转换如合并相邻事件)
|
||||
|
||||
Returns:
|
||||
转换后的 SSE 行列表(一入多出),空列表表示跳过该行
|
||||
(sse_lines, converted_events) 元组:
|
||||
- sse_lines: 转换后的 SSE 行列表(一入多出),空列表表示跳过该行
|
||||
- converted_events: 转换后的事件对象列表(用于记录到 parsed_chunks)
|
||||
"""
|
||||
from src.core.api_format import (
|
||||
GeminiStreamConversionState,
|
||||
StreamConversionState,
|
||||
converter_registry,
|
||||
)
|
||||
# 空行直接返回
|
||||
if not line or line.strip() == "":
|
||||
return ([line] if line else [], [])
|
||||
|
||||
# 如果是空行或特殊控制行,直接返回
|
||||
if not line or line.strip() == "" or line == "data: [DONE]":
|
||||
return [line] if line else []
|
||||
|
||||
provider_format = (ctx.provider_api_format or "").upper()
|
||||
client_format = (ctx.client_api_format or "").upper()
|
||||
|
||||
# 兼容 Gemini 上游 JSON-array/chunks:可能是“裸 JSON 行”而非 `data: {...}`
|
||||
if not line.startswith("data:"):
|
||||
if provider_format != "GEMINI":
|
||||
return [line]
|
||||
|
||||
stripped = line.strip()
|
||||
if stripped in ("", "[", "]", ","):
|
||||
return []
|
||||
|
||||
candidate = stripped.lstrip(",").rstrip(",").strip()
|
||||
try:
|
||||
data_obj = json.loads(candidate)
|
||||
except json.JSONDecodeError:
|
||||
# 不是可解析的 JSON 对象:不透传(避免把 Gemini 原始分块泄漏到目标 SSE)
|
||||
logger.debug(f"Gemini JSON-array line skip: {stripped[:50]}")
|
||||
return []
|
||||
else:
|
||||
# 提取 data 内容
|
||||
data_content = line[5:].strip() # 去掉 "data:" 前缀
|
||||
|
||||
# 尝试解析 JSON
|
||||
try:
|
||||
data_obj = json.loads(data_content)
|
||||
except json.JSONDecodeError:
|
||||
# 无法解析,直接透传
|
||||
return [line]
|
||||
|
||||
# 初始化流式转换状态(首次调用时,根据 Provider 格式选择状态类)
|
||||
if ctx.stream_conversion_state is None:
|
||||
if provider_format == "GEMINI":
|
||||
ctx.stream_conversion_state = GeminiStreamConversionState(
|
||||
model=ctx.mapped_model or ctx.model or "",
|
||||
message_id=ctx.response_id or ctx.request_id or "",
|
||||
)
|
||||
# [DONE] 标记处理:只有 OpenAI 客户端需要,Claude 客户端不需要
|
||||
if line == "data: [DONE]":
|
||||
if client_format.startswith("OPENAI"):
|
||||
return [line], []
|
||||
else:
|
||||
ctx.stream_conversion_state = StreamConversionState(
|
||||
model=ctx.mapped_model or ctx.model or "",
|
||||
message_id=ctx.response_id or ctx.request_id or "",
|
||||
)
|
||||
# Claude/Gemini 客户端不需要 [DONE] 标记
|
||||
return [], []
|
||||
|
||||
# 使用注册表进行格式转换(严格模式,返回 List[Dict])
|
||||
provider_format = (ctx.provider_api_format or "").upper()
|
||||
|
||||
# 过滤上游控制行(id/retry),避免与目标格式混淆
|
||||
if line.startswith(("id:", "retry:")):
|
||||
return [], []
|
||||
|
||||
# 解析 SSE 行为 JSON 对象
|
||||
data_obj, status = self._parse_sse_line_to_json(line, provider_format)
|
||||
|
||||
# 根据解析状态决定行为
|
||||
if status == "empty" or status == "skip":
|
||||
return [], []
|
||||
if status == "invalid" or status == "passthrough":
|
||||
return [line], []
|
||||
|
||||
# 初始化流式转换状态
|
||||
if ctx.stream_conversion_state is None:
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
# 使用客户端请求的模型(ctx.model),而非映射后的上游模型(ctx.mapped_model)
|
||||
init_model = ctx.model or ""
|
||||
logger.debug(
|
||||
f"[{ctx.request_id}] StreamState init: ctx.model={ctx.model!r}, "
|
||||
f"mapped_model={ctx.mapped_model!r}, using={init_model!r}"
|
||||
)
|
||||
ctx.stream_conversion_state = StreamState(
|
||||
model=init_model,
|
||||
message_id=ctx.response_id or ctx.request_id or "",
|
||||
)
|
||||
|
||||
# 执行格式转换
|
||||
try:
|
||||
converted_events = converter_registry.convert_stream_chunk_strict(
|
||||
registry = get_format_converter_registry()
|
||||
# status == "ok" 时 data_obj 必定是有效的 dict(防御性检查)
|
||||
if data_obj is None:
|
||||
return [], []
|
||||
converted_events = registry.convert_stream_chunk(
|
||||
data_obj,
|
||||
provider_format,
|
||||
client_format,
|
||||
state=ctx.stream_conversion_state,
|
||||
)
|
||||
|
||||
# 转换为 SSE 行列表
|
||||
result = []
|
||||
for evt in converted_events:
|
||||
result.append(f"data: {json.dumps(evt, ensure_ascii=False)}")
|
||||
return result
|
||||
result = _format_converted_events_to_sse(converted_events, client_format)
|
||||
if result:
|
||||
logger.debug(
|
||||
f"[{getattr(ctx, 'request_id', 'unknown')}] 流式转换: "
|
||||
f"{provider_format}->{client_format}, events={len(converted_events)}, "
|
||||
f"first_output={result[0][:100] if result else 'empty'}..."
|
||||
)
|
||||
return result, converted_events
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"格式转换失败,透传原始数据: {e}")
|
||||
return [line]
|
||||
return [line], []
|
||||
|
||||
def _parse_sse_line_to_json(
|
||||
self, line: str, provider_format: str
|
||||
) -> Tuple[Optional[Any], str]:
|
||||
"""
|
||||
解析 SSE 行为 JSON 对象
|
||||
|
||||
支持多种格式:
|
||||
- 标准 SSE: "data: {...}"
|
||||
- event+data 同行: "event: xxx data: {...}"
|
||||
- Gemini JSON-array: 裸 JSON 行
|
||||
|
||||
Args:
|
||||
line: 原始 SSE 行
|
||||
provider_format: Provider API 格式
|
||||
|
||||
Returns:
|
||||
(parsed_json, status) 元组:
|
||||
- (obj, "ok") - 解析成功
|
||||
- (None, "empty") - 内容为空,应跳过
|
||||
- (None, "invalid") - JSON 解析失败,应透传原始行
|
||||
- (None, "skip") - 应跳过(如纯 event 行)
|
||||
- (None, "passthrough") - 无法识别,应透传原始行
|
||||
"""
|
||||
# 标准 SSE: data: {...}
|
||||
if line.startswith("data:"):
|
||||
return _parse_sse_data_line(line)
|
||||
|
||||
# event + data 同行: event: xxx data: {...}
|
||||
if line.startswith("event:") and " data:" in line:
|
||||
return _parse_sse_event_data_line(line)
|
||||
|
||||
# 纯 event 行不参与转换
|
||||
if line.startswith("event:"):
|
||||
return None, "skip"
|
||||
|
||||
# Gemini JSON-array 格式
|
||||
if provider_format == "GEMINI":
|
||||
return _parse_gemini_json_array_line(line)
|
||||
|
||||
# 其他格式:无法识别,透传
|
||||
return None, "passthrough"
|
||||
|
||||
|
||||
@@ -169,14 +169,17 @@ class OpenAIResponseParser(ResponseParser):
|
||||
|
||||
# 提取 usage 信息(某些 OpenAI 兼容 API 如豆包会在最后一个 chunk 中发送 usage)
|
||||
# 这个 chunk 通常 choices 为空数组,但包含完整的 usage 信息
|
||||
# 使用取最大值策略确保正确统计
|
||||
usage = parsed.get("usage")
|
||||
if usage and isinstance(usage, dict):
|
||||
chunk.input_tokens = usage.get("prompt_tokens", 0)
|
||||
chunk.output_tokens = usage.get("completion_tokens", 0)
|
||||
|
||||
# 更新 stats
|
||||
stats.input_tokens = chunk.input_tokens
|
||||
stats.output_tokens = chunk.output_tokens
|
||||
# 取最大值更新 stats
|
||||
if chunk.input_tokens > stats.input_tokens:
|
||||
stats.input_tokens = chunk.input_tokens
|
||||
if chunk.output_tokens > stats.output_tokens:
|
||||
stats.output_tokens = chunk.output_tokens
|
||||
|
||||
stats.chunk_count += 1
|
||||
stats.data_count += 1
|
||||
@@ -287,6 +290,9 @@ class ClaudeResponseParser(ResponseParser):
|
||||
stats.has_completion = True
|
||||
|
||||
# 提取 usage
|
||||
# Claude 流式响应的 usage 可能在首个 chunk(message_start)或最后一个 chunk(message_delta)中
|
||||
# 首个 chunk 通常包含 input_tokens,最后一个 chunk 包含 output_tokens
|
||||
# 使用取最大值策略确保正确统计
|
||||
usage = self._parser.extract_usage(parsed)
|
||||
if usage:
|
||||
chunk.input_tokens = usage.get("input_tokens", 0)
|
||||
@@ -294,10 +300,15 @@ class ClaudeResponseParser(ResponseParser):
|
||||
chunk.cache_creation_tokens = usage.get("cache_creation_tokens", 0)
|
||||
chunk.cache_read_tokens = usage.get("cache_read_tokens", 0)
|
||||
|
||||
stats.input_tokens = chunk.input_tokens
|
||||
stats.output_tokens = chunk.output_tokens
|
||||
stats.cache_creation_tokens = chunk.cache_creation_tokens
|
||||
stats.cache_read_tokens = chunk.cache_read_tokens
|
||||
# 取最大值更新 stats
|
||||
if chunk.input_tokens > stats.input_tokens:
|
||||
stats.input_tokens = chunk.input_tokens
|
||||
if chunk.output_tokens > stats.output_tokens:
|
||||
stats.output_tokens = chunk.output_tokens
|
||||
if chunk.cache_creation_tokens > stats.cache_creation_tokens:
|
||||
stats.cache_creation_tokens = chunk.cache_creation_tokens
|
||||
if chunk.cache_read_tokens > stats.cache_read_tokens:
|
||||
stats.cache_read_tokens = chunk.cache_read_tokens
|
||||
|
||||
# 检查错误
|
||||
if self._parser.is_error_event(parsed):
|
||||
@@ -432,15 +443,21 @@ class GeminiResponseParser(ResponseParser):
|
||||
stats.has_completion = True
|
||||
|
||||
# 提取 usage
|
||||
# Gemini 流式响应的 usage 可能出现在多个 chunk 中
|
||||
# 使用取最大值策略确保正确统计
|
||||
usage = self._parser.extract_usage(parsed)
|
||||
if usage:
|
||||
chunk.input_tokens = usage.get("input_tokens", 0)
|
||||
chunk.output_tokens = usage.get("output_tokens", 0)
|
||||
chunk.cache_read_tokens = usage.get("cached_tokens", 0)
|
||||
|
||||
stats.input_tokens = chunk.input_tokens
|
||||
stats.output_tokens = chunk.output_tokens
|
||||
stats.cache_read_tokens = chunk.cache_read_tokens
|
||||
# 取最大值更新 stats
|
||||
if chunk.input_tokens > stats.input_tokens:
|
||||
stats.input_tokens = chunk.input_tokens
|
||||
if chunk.output_tokens > stats.output_tokens:
|
||||
stats.output_tokens = chunk.output_tokens
|
||||
if chunk.cache_read_tokens > stats.cache_read_tokens:
|
||||
stats.cache_read_tokens = chunk.cache_read_tokens
|
||||
|
||||
# 检查错误
|
||||
if self._parser.is_error_event(parsed):
|
||||
|
||||
@@ -65,7 +65,7 @@ def build_test_request_body(
|
||||
) -> Dict[str, Any]:
|
||||
"""构建测试请求体,自动处理格式转换
|
||||
|
||||
使用 converter_registry 将 OpenAI 格式的测试请求转换为目标格式。
|
||||
使用格式转换注册表将 OpenAI 格式的测试请求转换为目标格式。
|
||||
|
||||
Args:
|
||||
format_id: 目标 API 格式 ID(如 "CLAUDE", "GEMINI", "OPENAI_CLI")
|
||||
@@ -74,18 +74,22 @@ def build_test_request_body(
|
||||
Returns:
|
||||
转换为目标 API 格式的请求体
|
||||
"""
|
||||
from src.core.api_format.conversion import converter_registry
|
||||
from src.core.api_format.conversion import (
|
||||
format_conversion_registry,
|
||||
register_default_normalizers,
|
||||
)
|
||||
from src.core.api_format.utils import get_base_format
|
||||
|
||||
register_default_normalizers()
|
||||
|
||||
# 获取测试请求数据(OpenAI 格式)
|
||||
source_data = get_test_request_data(request_data)
|
||||
|
||||
# CLI 格式使用基础格式进行转换(CLAUDE_CLI -> CLAUDE)
|
||||
# 因为 converter_registry 只注册了基础格式之间的转换器
|
||||
target_format = get_base_format(format_id) or format_id
|
||||
|
||||
# 使用注册表进行格式转换 (OPENAI -> 目标基础格式)
|
||||
return converter_registry.convert_request(source_data, "OPENAI", target_format)
|
||||
return format_conversion_registry.convert_request(source_data, "OPENAI", target_format)
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
|
||||
@@ -10,15 +10,10 @@
|
||||
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.core.api_format import (
|
||||
ClaudeStreamConversionState,
|
||||
GeminiStreamConversionState,
|
||||
OpenAIStreamConversionState,
|
||||
StreamConversionState,
|
||||
)
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -95,14 +90,7 @@ class StreamContext:
|
||||
parsed_chunks: List[Dict[str, Any]] = field(default_factory=list)
|
||||
|
||||
# 流式格式转换状态(跨 chunk 追踪)
|
||||
stream_conversion_state: Optional[
|
||||
Union[
|
||||
"StreamConversionState",
|
||||
"GeminiStreamConversionState",
|
||||
"ClaudeStreamConversionState",
|
||||
"OpenAIStreamConversionState",
|
||||
]
|
||||
] = None
|
||||
stream_conversion_state: Optional["StreamState"] = None
|
||||
|
||||
def reset_for_retry(self) -> None:
|
||||
"""
|
||||
|
||||
@@ -28,10 +28,11 @@ from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.api.handlers.base.utils import (
|
||||
check_html_response,
|
||||
check_prefetched_response_error,
|
||||
get_format_converter_registry,
|
||||
)
|
||||
from src.config.constants import StreamDefaults
|
||||
from src.config.settings import config
|
||||
from src.core.api_format import FormatConversionError, converter_registry
|
||||
from src.core.api_format.conversion.exceptions import FormatConversionError
|
||||
from src.core.exceptions import (
|
||||
EmbeddedErrorException,
|
||||
ProviderNotAvailableException,
|
||||
@@ -306,7 +307,8 @@ class StreamProcessor:
|
||||
try:
|
||||
# 试转换:传 state=None,不保留状态
|
||||
# 如果失败触发 failover,下一个候选会使用干净的 state
|
||||
converter_registry.convert_stream_chunk_strict(
|
||||
registry = get_format_converter_registry()
|
||||
registry.convert_stream_chunk(
|
||||
data,
|
||||
provider_format,
|
||||
client_format,
|
||||
@@ -451,38 +453,17 @@ class StreamProcessor:
|
||||
|
||||
# 处理预读数据
|
||||
if needs_conversion:
|
||||
# 延迟导入:仅在需要转换时加载转换器模块
|
||||
from src.core.api_format import (
|
||||
ClaudeStreamConversionState,
|
||||
GeminiStreamConversionState,
|
||||
OpenAIStreamConversionState,
|
||||
StreamConversionState,
|
||||
converter_registry,
|
||||
)
|
||||
registry = get_format_converter_registry()
|
||||
|
||||
# 初始化流式转换状态(首次使用时,根据 Provider 格式选择状态类)
|
||||
# 状态类对应 Provider 的响应格式,用于正确解析和累积流式数据
|
||||
# 初始化流式转换状态(Canonical)
|
||||
if ctx.stream_conversion_state is None:
|
||||
if provider_format == "GEMINI":
|
||||
ctx.stream_conversion_state = GeminiStreamConversionState(
|
||||
model=ctx.mapped_model or ctx.model or "",
|
||||
message_id=ctx.response_id or ctx.request_id or "",
|
||||
)
|
||||
elif provider_format == "OPENAI":
|
||||
ctx.stream_conversion_state = OpenAIStreamConversionState(
|
||||
model=ctx.mapped_model or ctx.model or "",
|
||||
)
|
||||
elif provider_format == "CLAUDE":
|
||||
ctx.stream_conversion_state = ClaudeStreamConversionState(
|
||||
model=ctx.mapped_model or ctx.model or "",
|
||||
message_id=ctx.response_id or ctx.request_id or "",
|
||||
)
|
||||
else:
|
||||
# 兜底:使用通用状态
|
||||
ctx.stream_conversion_state = StreamConversionState(
|
||||
model=ctx.mapped_model or ctx.model or "",
|
||||
message_id=ctx.response_id or ctx.request_id or "",
|
||||
)
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
# 使用客户端请求的模型(ctx.model),而非映射后的上游模型(ctx.mapped_model)
|
||||
ctx.stream_conversion_state = StreamState(
|
||||
model=ctx.model or "",
|
||||
message_id=ctx.response_id or ctx.request_id or "",
|
||||
)
|
||||
|
||||
skip_next_blank_line = False
|
||||
empty_yield_count = 0 # 空转计数(防护异常情况)
|
||||
@@ -545,7 +526,7 @@ class StreamProcessor:
|
||||
return []
|
||||
|
||||
try:
|
||||
converted_events = converter_registry.convert_stream_chunk_strict(
|
||||
converted_events = registry.convert_stream_chunk(
|
||||
data_obj,
|
||||
provider_format,
|
||||
client_format,
|
||||
|
||||
@@ -2,13 +2,34 @@
|
||||
Handler 基础工具函数
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional
|
||||
|
||||
from src.core.exceptions import EmbeddedErrorException, ProviderNotAvailableException
|
||||
from src.core.api_format import filter_response_headers
|
||||
from src.core.logger import logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
||||
|
||||
|
||||
def get_format_converter_registry() -> "FormatConversionRegistry":
|
||||
"""
|
||||
获取格式转换注册表(线程安全)
|
||||
|
||||
该函数确保 normalizers 已注册后再返回全局注册表实例。
|
||||
register_default_normalizers 内部已有双重检查锁,可安全多次调用。
|
||||
"""
|
||||
from src.core.api_format.conversion.registry import (
|
||||
format_conversion_registry,
|
||||
register_default_normalizers,
|
||||
)
|
||||
|
||||
register_default_normalizers()
|
||||
return format_conversion_registry
|
||||
|
||||
|
||||
def extract_cache_creation_tokens(usage: Dict[str, Any]) -> int:
|
||||
"""
|
||||
|
||||
@@ -198,7 +198,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
|
||||
else:
|
||||
return f"{base_url}/v1/messages"
|
||||
|
||||
# build_request_body 使用基类实现,通过 converter_registry 自动转换 OPENAI -> CLAUDE
|
||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE
|
||||
|
||||
|
||||
def build_claude_adapter(x_app_header: Optional[str]):
|
||||
|
||||
@@ -74,18 +74,23 @@ class ClaudeChatHandler(ChatHandlerBase):
|
||||
Returns:
|
||||
ClaudeMessagesRequest 对象
|
||||
"""
|
||||
from src.core.api_format import OpenAIToClaudeConverter
|
||||
from src.core.api_format.conversion.registry import (
|
||||
format_conversion_registry,
|
||||
register_default_normalizers,
|
||||
)
|
||||
from src.models.claude import ClaudeMessagesRequest
|
||||
from src.models.openai import OpenAIRequest
|
||||
|
||||
register_default_normalizers()
|
||||
|
||||
# 如果已经是 Claude 格式,直接返回
|
||||
if isinstance(request, ClaudeMessagesRequest):
|
||||
return request
|
||||
|
||||
# 如果是 OpenAI 格式,转换为 Claude 格式
|
||||
if isinstance(request, OpenAIRequest):
|
||||
converter = OpenAIToClaudeConverter()
|
||||
claude_dict = converter.convert_request(request.dict())
|
||||
req_dict = request.model_dump() if hasattr(request, "model_dump") else request.dict()
|
||||
claude_dict = format_conversion_registry.convert_request(req_dict, "OPENAI", "CLAUDE")
|
||||
return ClaudeMessagesRequest(**claude_dict)
|
||||
|
||||
# 如果是字典,根据内容判断格式
|
||||
@@ -94,8 +99,7 @@ class ClaudeChatHandler(ChatHandlerBase):
|
||||
first_msg = request["messages"][0]
|
||||
if "role" in first_msg and "content" in first_msg:
|
||||
# 可能是 OpenAI 格式
|
||||
converter = OpenAIToClaudeConverter()
|
||||
claude_dict = converter.convert_request(request)
|
||||
claude_dict = format_conversion_registry.convert_request(request, "OPENAI", "CLAUDE")
|
||||
return ClaudeMessagesRequest(**claude_dict)
|
||||
|
||||
# 否则假设已经是 Claude 格式
|
||||
|
||||
@@ -128,7 +128,7 @@ class ClaudeCliAdapter(CliAdapterBase):
|
||||
else:
|
||||
return f"{base_url}/v1/messages"
|
||||
|
||||
# build_request_body 使用基类实现,通过 converter_registry 自动转换 OPENAI -> CLAUDE_CLI
|
||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> CLAUDE_CLI
|
||||
|
||||
@classmethod
|
||||
def get_cli_user_agent(cls) -> Optional[str]:
|
||||
|
||||
@@ -7,20 +7,10 @@ Gemini API Handler 模块
|
||||
from src.api.handlers.gemini.adapter import GeminiChatAdapter, build_gemini_adapter
|
||||
from src.api.handlers.gemini.handler import GeminiChatHandler
|
||||
from src.api.handlers.gemini.stream_parser import GeminiStreamParser
|
||||
from src.core.api_format import (
|
||||
ClaudeToGeminiConverter,
|
||||
GeminiToClaudeConverter,
|
||||
GeminiToOpenAIConverter,
|
||||
OpenAIToGeminiConverter,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"GeminiChatAdapter",
|
||||
"GeminiChatHandler",
|
||||
"GeminiStreamParser",
|
||||
"ClaudeToGeminiConverter",
|
||||
"GeminiToClaudeConverter",
|
||||
"OpenAIToGeminiConverter",
|
||||
"GeminiToOpenAIConverter",
|
||||
"build_gemini_adapter",
|
||||
]
|
||||
|
||||
@@ -223,7 +223,7 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
else:
|
||||
return f"{base_url}/v1beta"
|
||||
|
||||
# build_request_body 使用基类实现,通过 converter_registry 自动转换 OPENAI -> GEMINI
|
||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> GEMINI
|
||||
|
||||
@classmethod
|
||||
async def check_endpoint(
|
||||
|
||||
@@ -62,28 +62,30 @@ class GeminiChatHandler(ChatHandlerBase):
|
||||
Returns:
|
||||
GeminiRequest 对象
|
||||
"""
|
||||
from src.core.api_format import (
|
||||
ClaudeToGeminiConverter,
|
||||
OpenAIToGeminiConverter,
|
||||
from src.core.api_format.conversion.registry import (
|
||||
format_conversion_registry,
|
||||
register_default_normalizers,
|
||||
)
|
||||
from src.models.claude import ClaudeMessagesRequest
|
||||
from src.models.gemini import GeminiRequest
|
||||
from src.models.openai import OpenAIRequest
|
||||
|
||||
register_default_normalizers()
|
||||
|
||||
# 如果已经是 Gemini 格式,直接返回
|
||||
if isinstance(request, GeminiRequest):
|
||||
return request
|
||||
|
||||
# 如果是 Claude 格式,转换为 Gemini 格式
|
||||
if isinstance(request, ClaudeMessagesRequest):
|
||||
converter = ClaudeToGeminiConverter()
|
||||
gemini_dict = converter.convert_request(request.model_dump())
|
||||
req_dict = request.model_dump() if hasattr(request, "model_dump") else request.dict()
|
||||
gemini_dict = format_conversion_registry.convert_request(req_dict, "CLAUDE", "GEMINI")
|
||||
return GeminiRequest(**gemini_dict)
|
||||
|
||||
# 如果是 OpenAI 格式,转换为 Gemini 格式
|
||||
if isinstance(request, OpenAIRequest):
|
||||
converter = OpenAIToGeminiConverter()
|
||||
gemini_dict = converter.convert_request(request.model_dump())
|
||||
req_dict = request.model_dump() if hasattr(request, "model_dump") else request.dict()
|
||||
gemini_dict = format_conversion_registry.convert_request(req_dict, "OPENAI", "GEMINI")
|
||||
return GeminiRequest(**gemini_dict)
|
||||
|
||||
# 如果是字典,根据内容判断格式并转换
|
||||
@@ -100,13 +102,11 @@ class GeminiChatHandler(ChatHandlerBase):
|
||||
messages = request.get("messages", [])
|
||||
if messages and isinstance(messages[0].get("content"), list):
|
||||
# 可能是 Claude 格式
|
||||
converter = ClaudeToGeminiConverter()
|
||||
gemini_dict = converter.convert_request(request)
|
||||
gemini_dict = format_conversion_registry.convert_request(request, "CLAUDE", "GEMINI")
|
||||
return GeminiRequest(**gemini_dict)
|
||||
else:
|
||||
# 可能是 OpenAI 格式
|
||||
converter = OpenAIToGeminiConverter()
|
||||
gemini_dict = converter.convert_request(request)
|
||||
gemini_dict = format_conversion_registry.convert_request(request, "OPENAI", "GEMINI")
|
||||
return GeminiRequest(**gemini_dict)
|
||||
|
||||
# 默认尝试作为 Gemini 格式
|
||||
|
||||
@@ -149,7 +149,7 @@ class GeminiCliAdapter(CliAdapterBase):
|
||||
prefix = f"{base_url}/v1beta"
|
||||
return f"{prefix}/models/{effective_model_name}:generateContent"
|
||||
|
||||
# build_request_body 使用基类实现,通过 converter_registry 自动转换 OPENAI -> GEMINI_CLI
|
||||
# build_request_body 使用基类实现,通过 format_conversion_registry 自动转换 OPENAI -> GEMINI_CLI
|
||||
|
||||
@classmethod
|
||||
def get_cli_user_agent(cls) -> Optional[str]:
|
||||
|
||||
@@ -73,18 +73,23 @@ class OpenAIChatHandler(ChatHandlerBase):
|
||||
Returns:
|
||||
OpenAIRequest 对象
|
||||
"""
|
||||
from src.core.api_format import ClaudeToOpenAIConverter
|
||||
from src.core.api_format.conversion.registry import (
|
||||
format_conversion_registry,
|
||||
register_default_normalizers,
|
||||
)
|
||||
from src.models.claude import ClaudeMessagesRequest
|
||||
from src.models.openai import OpenAIRequest
|
||||
|
||||
register_default_normalizers()
|
||||
|
||||
# 如果已经是 OpenAI 格式,直接返回
|
||||
if isinstance(request, OpenAIRequest):
|
||||
return request
|
||||
|
||||
# 如果是 Claude 格式,转换为 OpenAI 格式
|
||||
if isinstance(request, ClaudeMessagesRequest):
|
||||
converter = ClaudeToOpenAIConverter()
|
||||
openai_dict = converter.convert_request(request.dict())
|
||||
req_dict = request.model_dump() if hasattr(request, "model_dump") else request.dict()
|
||||
openai_dict = format_conversion_registry.convert_request(req_dict, "CLAUDE", "OPENAI")
|
||||
return OpenAIRequest(**openai_dict)
|
||||
|
||||
# 如果是字典,尝试判断格式
|
||||
@@ -93,8 +98,7 @@ class OpenAIChatHandler(ChatHandlerBase):
|
||||
return OpenAIRequest(**request)
|
||||
except Exception:
|
||||
try:
|
||||
converter = ClaudeToOpenAIConverter()
|
||||
openai_dict = converter.convert_request(request)
|
||||
openai_dict = format_conversion_registry.convert_request(request, "CLAUDE", "OPENAI")
|
||||
return OpenAIRequest(**openai_dict)
|
||||
except Exception:
|
||||
return OpenAIRequest(**request)
|
||||
|
||||
@@ -27,8 +27,10 @@ from src.core.api_format import (
|
||||
ApiFormatDefinition,
|
||||
detect_format_and_key_from_starlette,
|
||||
)
|
||||
from src.core.api_format.conversion import converter_registry
|
||||
from src.core.api_format.utils import is_cli_format
|
||||
from src.core.api_format.conversion import (
|
||||
format_conversion_registry,
|
||||
register_default_normalizers,
|
||||
)
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.database import ApiKey, User
|
||||
@@ -114,11 +116,8 @@ def _get_convertible_formats(client_format: str, global_conversion_enabled: bool
|
||||
|
||||
client_format_upper = client_format.upper()
|
||||
|
||||
# CLI 格式不支持转换
|
||||
if is_cli_format(client_format_upper):
|
||||
return _get_formats_for_api(client_format)
|
||||
|
||||
# 收集所有可转换的格式
|
||||
register_default_normalizers()
|
||||
convertible_formats = []
|
||||
for target_format in _ALL_CHAT_FORMATS:
|
||||
# 相同格式始终可用
|
||||
@@ -127,7 +126,11 @@ def _get_convertible_formats(client_format: str, global_conversion_enabled: bool
|
||||
continue
|
||||
|
||||
# 检查是否有双向转换器
|
||||
if converter_registry.can_convert_full(client_format_upper, target_format, require_stream=False):
|
||||
if format_conversion_registry.can_convert_full(
|
||||
client_format_upper,
|
||||
target_format,
|
||||
require_stream=False,
|
||||
):
|
||||
convertible_formats.append(target_format)
|
||||
|
||||
return convertible_formats if convertible_formats else _get_formats_for_api(client_format)
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""
|
||||
API 格式核心模块
|
||||
|
||||
统一管理 API 格式相关的枚举、元数据、工具函数和格式转换功能。
|
||||
统一管理 API 格式相关的枚举、元数据、工具函数等。
|
||||
|
||||
模块组成:
|
||||
- enums.py: APIFormat 枚举定义
|
||||
@@ -9,29 +9,7 @@ API 格式核心模块
|
||||
- headers.py: 请求头处理(构建、过滤、脱敏)
|
||||
- utils.py: 工具函数(is_cli_format, get_base_format 等)
|
||||
- detection.py: 格式检测(从请求头、响应内容检测格式)
|
||||
- conversion/: 格式转换子模块
|
||||
"""
|
||||
|
||||
from src.core.api_format.conversion import (
|
||||
ClaudeStreamConversionState,
|
||||
ClaudeToGeminiConverter,
|
||||
ClaudeToOpenAIConverter,
|
||||
FormatConversionError,
|
||||
FormatConverterRegistry,
|
||||
GeminiStreamConversionState,
|
||||
GeminiToClaudeConverter,
|
||||
GeminiToOpenAIConverter,
|
||||
OpenAIStreamConversionState,
|
||||
OpenAIToClaudeConverter,
|
||||
OpenAIToGeminiConverter,
|
||||
RequestConverter,
|
||||
ResponseConverter,
|
||||
StreamChunkConverter,
|
||||
StreamConversionState,
|
||||
converter_registry,
|
||||
is_format_compatible,
|
||||
register_all_converters,
|
||||
)
|
||||
from src.core.api_format.detection import (
|
||||
detect_cli_format_from_path,
|
||||
detect_format_and_key_from_starlette,
|
||||
@@ -128,33 +106,9 @@ __all__ = [
|
||||
"get_adapter_protected_keys",
|
||||
"extract_set_headers_from_rules",
|
||||
"get_extra_headers_from_endpoint",
|
||||
# Registry
|
||||
"FormatConverterRegistry",
|
||||
"converter_registry",
|
||||
"register_all_converters",
|
||||
# Protocols
|
||||
"RequestConverter",
|
||||
"ResponseConverter",
|
||||
"StreamChunkConverter",
|
||||
# State
|
||||
"StreamConversionState",
|
||||
"GeminiStreamConversionState",
|
||||
"ClaudeStreamConversionState",
|
||||
"OpenAIStreamConversionState",
|
||||
# Exceptions
|
||||
"FormatConversionError",
|
||||
# Compatibility
|
||||
"is_format_compatible",
|
||||
# Detection
|
||||
"detect_format_from_request",
|
||||
"detect_format_and_key_from_starlette",
|
||||
"detect_format_from_response",
|
||||
"detect_cli_format_from_path",
|
||||
# Converters
|
||||
"OpenAIToClaudeConverter",
|
||||
"ClaudeToOpenAIConverter",
|
||||
"ClaudeToGeminiConverter",
|
||||
"GeminiToClaudeConverter",
|
||||
"OpenAIToGeminiConverter",
|
||||
"GeminiToOpenAIConverter",
|
||||
]
|
||||
|
||||
@@ -1,90 +1,30 @@
|
||||
"""
|
||||
格式转换核心模块
|
||||
API 格式转换子模块(Canonical)
|
||||
|
||||
该目录用于承载与「API 格式转换」相关的核心能力(不依赖 FastAPI/Handler 层),
|
||||
以便 services/core 可复用并避免出现 services -> api 的反向依赖。
|
||||
|
||||
模块组成:
|
||||
- registry.py: 转换器注册表,管理转换器实例和能力查询
|
||||
- protocols.py: 转换器协议定义(Protocol)
|
||||
- state.py: 流式转换状态类
|
||||
- exceptions.py: 转换异常定义
|
||||
- compatibility.py: 格式兼容性检查函数
|
||||
- converters/: 内置格式转换器实现
|
||||
对外提供:
|
||||
- `format_conversion_registry`: 全局转换注册表(Hub-and-Spoke)
|
||||
- `register_default_normalizers()`: 注册默认 Normalizers(OPENAI/CLAUDE/GEMINI)
|
||||
- `StreamState`: 统一流式状态容器
|
||||
"""
|
||||
|
||||
from src.core.api_format.conversion.compatibility import is_format_compatible
|
||||
from src.core.api_format.conversion.converters import (
|
||||
ClaudeToGeminiConverter,
|
||||
ClaudeToOpenAIConverter,
|
||||
GeminiToClaudeConverter,
|
||||
GeminiToOpenAIConverter,
|
||||
OpenAIToClaudeConverter,
|
||||
OpenAIToGeminiConverter,
|
||||
)
|
||||
from src.core.api_format.conversion.exceptions import FormatConversionError
|
||||
from src.core.api_format.conversion.protocols import (
|
||||
RequestConverter,
|
||||
ResponseConverter,
|
||||
StreamChunkConverter,
|
||||
)
|
||||
from src.core.api_format.conversion.registry import (
|
||||
FormatConverterRegistry,
|
||||
converter_registry,
|
||||
FormatConversionRegistry,
|
||||
format_conversion_registry,
|
||||
register_default_normalizers,
|
||||
)
|
||||
from src.core.api_format.conversion.state import (
|
||||
ClaudeStreamConversionState,
|
||||
GeminiStreamConversionState,
|
||||
OpenAIStreamConversionState,
|
||||
StreamConversionState,
|
||||
)
|
||||
from src.core.logger import logger
|
||||
|
||||
|
||||
def register_all_converters() -> None:
|
||||
"""
|
||||
注册所有内置的格式转换器
|
||||
|
||||
在应用启动时调用此函数
|
||||
"""
|
||||
# Claude <-> OpenAI
|
||||
converter_registry.register("OPENAI", "CLAUDE", OpenAIToClaudeConverter())
|
||||
converter_registry.register("CLAUDE", "OPENAI", ClaudeToOpenAIConverter())
|
||||
|
||||
# Claude <-> Gemini
|
||||
converter_registry.register("CLAUDE", "GEMINI", ClaudeToGeminiConverter())
|
||||
converter_registry.register("GEMINI", "CLAUDE", GeminiToClaudeConverter())
|
||||
|
||||
# OpenAI <-> Gemini
|
||||
converter_registry.register("OPENAI", "GEMINI", OpenAIToGeminiConverter())
|
||||
converter_registry.register("GEMINI", "OPENAI", GeminiToOpenAIConverter())
|
||||
|
||||
logger.info(f"[ConverterRegistry] 已注册 {len(converter_registry.list_converters())} 个格式转换器")
|
||||
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
__all__ = [
|
||||
# Registry
|
||||
"FormatConverterRegistry",
|
||||
"converter_registry",
|
||||
"register_all_converters",
|
||||
# Protocols
|
||||
"RequestConverter",
|
||||
"ResponseConverter",
|
||||
"StreamChunkConverter",
|
||||
# State
|
||||
"StreamConversionState",
|
||||
"GeminiStreamConversionState",
|
||||
"ClaudeStreamConversionState",
|
||||
"OpenAIStreamConversionState",
|
||||
"FormatConversionRegistry",
|
||||
"format_conversion_registry",
|
||||
"register_default_normalizers",
|
||||
# Stream state
|
||||
"StreamState",
|
||||
# Exceptions
|
||||
"FormatConversionError",
|
||||
# Compatibility
|
||||
"is_format_compatible",
|
||||
# Converters
|
||||
"OpenAIToClaudeConverter",
|
||||
"ClaudeToOpenAIConverter",
|
||||
"ClaudeToGeminiConverter",
|
||||
"GeminiToClaudeConverter",
|
||||
"OpenAIToGeminiConverter",
|
||||
"GeminiToOpenAIConverter",
|
||||
]
|
||||
|
||||
@@ -9,10 +9,10 @@ from __future__ import annotations
|
||||
import logging
|
||||
from typing import TYPE_CHECKING, Optional, Tuple
|
||||
|
||||
from src.core.api_format.utils import is_cli_format
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.core.api_format.conversion.registry import FormatConverterRegistry
|
||||
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
||||
|
||||
from src.core.api_format.utils import get_base_format
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -23,7 +23,7 @@ def is_format_compatible(
|
||||
endpoint_format_acceptance_config: Optional[dict],
|
||||
is_stream: bool,
|
||||
global_conversion_enabled: bool,
|
||||
registry: Optional["FormatConverterRegistry"] = None,
|
||||
registry: Optional["FormatConversionRegistry"] = None,
|
||||
) -> Tuple[bool, bool, Optional[str]]:
|
||||
"""
|
||||
检查端点是否兼容客户端格式
|
||||
@@ -44,9 +44,13 @@ def is_format_compatible(
|
||||
"""
|
||||
# 延迟导入避免循环依赖
|
||||
if registry is None:
|
||||
from src.core.api_format.conversion.registry import converter_registry
|
||||
from src.core.api_format.conversion.registry import (
|
||||
format_conversion_registry,
|
||||
register_default_normalizers,
|
||||
)
|
||||
|
||||
registry = converter_registry
|
||||
register_default_normalizers()
|
||||
registry = format_conversion_registry
|
||||
|
||||
provider_format = endpoint_api_format.upper()
|
||||
client_format_upper = client_format.upper()
|
||||
@@ -55,11 +59,12 @@ def is_format_compatible(
|
||||
if provider_format == client_format_upper:
|
||||
return True, False, None
|
||||
|
||||
# 2. CLI 格式不参与转换
|
||||
if is_cli_format(client_format_upper):
|
||||
return False, False, "CLI 格式不支持转换"
|
||||
if is_cli_format(provider_format):
|
||||
return False, False, "Provider 为 CLI 格式,不支持转换"
|
||||
# 2. 同族格式匹配(CLAUDE 和 CLAUDE_CLI、OPENAI 和 OPENAI_CLI 等)
|
||||
# 这些格式在响应层面是兼容的,只是认证方式不同
|
||||
provider_base = get_base_format(provider_format)
|
||||
client_base = get_base_format(client_format_upper)
|
||||
if provider_base == client_base:
|
||||
return True, False, None
|
||||
|
||||
# 3. 检查全局开关
|
||||
if not global_conversion_enabled:
|
||||
|
||||
@@ -1,26 +0,0 @@
|
||||
"""
|
||||
格式转换器集合
|
||||
|
||||
将 Claude/OpenAI/Gemini 格式之间的转换器统一导出。
|
||||
"""
|
||||
|
||||
from .claude_to_openai import ClaudeToOpenAIConverter
|
||||
from .gemini import (
|
||||
ClaudeToGeminiConverter,
|
||||
GeminiToClaudeConverter,
|
||||
GeminiToOpenAIConverter,
|
||||
OpenAIToGeminiConverter,
|
||||
)
|
||||
from .openai_to_claude import OpenAIToClaudeConverter
|
||||
|
||||
__all__ = [
|
||||
# OpenAI <-> Claude
|
||||
"OpenAIToClaudeConverter",
|
||||
"ClaudeToOpenAIConverter",
|
||||
# Claude <-> Gemini
|
||||
"ClaudeToGeminiConverter",
|
||||
"GeminiToClaudeConverter",
|
||||
# OpenAI <-> Gemini
|
||||
"OpenAIToGeminiConverter",
|
||||
"GeminiToOpenAIConverter",
|
||||
]
|
||||
@@ -1,450 +0,0 @@
|
||||
"""
|
||||
Claude -> OpenAI 格式转换器
|
||||
|
||||
将 Claude Messages API 格式转换为 OpenAI Chat Completions API 格式。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.core.api_format.conversion.state import StreamConversionState
|
||||
|
||||
|
||||
class ClaudeToOpenAIConverter:
|
||||
"""
|
||||
Claude -> OpenAI 格式转换器
|
||||
|
||||
支持:
|
||||
- 请求转换:Claude Request -> OpenAI Chat Request
|
||||
- 响应转换:Claude Response -> OpenAI Chat Response
|
||||
- 流式转换:Claude SSE -> OpenAI SSE
|
||||
"""
|
||||
|
||||
# 内容类型常量
|
||||
CONTENT_TYPE_TEXT = "text"
|
||||
CONTENT_TYPE_IMAGE = "image"
|
||||
CONTENT_TYPE_TOOL_USE = "tool_use"
|
||||
CONTENT_TYPE_TOOL_RESULT = "tool_result"
|
||||
|
||||
# 停止原因映射
|
||||
STOP_REASON_MAP = {
|
||||
"end_turn": "stop",
|
||||
"max_tokens": "length",
|
||||
"stop_sequence": "stop",
|
||||
"tool_use": "tool_calls",
|
||||
}
|
||||
|
||||
def __init__(self, model_mapping: Optional[Dict[str, str]] = None):
|
||||
"""
|
||||
Args:
|
||||
model_mapping: Claude 模型到 OpenAI 模型的映射
|
||||
"""
|
||||
self._model_mapping = model_mapping or {}
|
||||
|
||||
# ==================== 请求转换 ====================
|
||||
|
||||
def convert_request(self, request: Union[Dict[str, Any], Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
将 Claude 请求转换为 OpenAI 格式
|
||||
|
||||
Args:
|
||||
request: Claude 请求(Dict 或 Pydantic 模型)
|
||||
|
||||
Returns:
|
||||
OpenAI 格式的请求字典
|
||||
"""
|
||||
if hasattr(request, "model_dump"):
|
||||
data = request.model_dump(exclude_none=True)
|
||||
else:
|
||||
data = dict(request)
|
||||
|
||||
# 模型映射
|
||||
model = data.get("model", "")
|
||||
openai_model = self._model_mapping.get(model, model)
|
||||
|
||||
# 构建消息列表
|
||||
messages: List[Dict[str, Any]] = []
|
||||
|
||||
# 处理 system 消息
|
||||
system_content = self._extract_text_content(data.get("system"))
|
||||
if system_content:
|
||||
messages.append({"role": "system", "content": system_content})
|
||||
|
||||
# 处理对话消息
|
||||
for message in data.get("messages", []):
|
||||
converted = self._convert_message(message)
|
||||
if converted:
|
||||
messages.append(converted)
|
||||
|
||||
# 构建 OpenAI 请求
|
||||
result: Dict[str, Any] = {
|
||||
"model": openai_model,
|
||||
"messages": messages,
|
||||
}
|
||||
|
||||
# 可选参数
|
||||
if data.get("max_tokens"):
|
||||
result["max_tokens"] = data["max_tokens"]
|
||||
if data.get("temperature") is not None:
|
||||
result["temperature"] = data["temperature"]
|
||||
if data.get("top_p") is not None:
|
||||
result["top_p"] = data["top_p"]
|
||||
if data.get("stream"):
|
||||
result["stream"] = data["stream"]
|
||||
if data.get("stop_sequences"):
|
||||
result["stop"] = data["stop_sequences"]
|
||||
|
||||
# 工具转换
|
||||
tools = self._convert_tools(data.get("tools"))
|
||||
if tools:
|
||||
result["tools"] = tools
|
||||
|
||||
tool_choice = self._convert_tool_choice(data.get("tool_choice"))
|
||||
if tool_choice:
|
||||
result["tool_choice"] = tool_choice
|
||||
|
||||
return result
|
||||
|
||||
def _convert_message(self, message: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
"""转换单条消息"""
|
||||
role = message.get("role")
|
||||
|
||||
if role == "user":
|
||||
return self._convert_user_message(message)
|
||||
if role == "assistant":
|
||||
return self._convert_assistant_message(message)
|
||||
|
||||
return None
|
||||
|
||||
def _convert_user_message(self, message: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""转换用户消息"""
|
||||
content = message.get("content")
|
||||
|
||||
if isinstance(content, str):
|
||||
return {"role": "user", "content": content}
|
||||
|
||||
openai_content: List[Dict[str, Any]] = []
|
||||
for block in content or []:
|
||||
block_type = block.get("type")
|
||||
|
||||
if block_type == self.CONTENT_TYPE_TEXT:
|
||||
openai_content.append({"type": "text", "text": block.get("text", "")})
|
||||
elif block_type == self.CONTENT_TYPE_IMAGE:
|
||||
source = block.get("source", {})
|
||||
media_type = source.get("media_type", "image/jpeg")
|
||||
data = source.get("data", "")
|
||||
openai_content.append(
|
||||
{"type": "image_url", "image_url": {"url": f"data:{media_type};base64,{data}"}}
|
||||
)
|
||||
elif block_type == self.CONTENT_TYPE_TOOL_RESULT:
|
||||
tool_content = block.get("content", "")
|
||||
rendered = self._render_tool_content(tool_content)
|
||||
openai_content.append({"type": "text", "text": f"Tool result: {rendered}"})
|
||||
|
||||
# 简化单文本内容
|
||||
if len(openai_content) == 1 and openai_content[0]["type"] == "text":
|
||||
return {"role": "user", "content": openai_content[0]["text"]}
|
||||
|
||||
return {"role": "user", "content": openai_content or ""}
|
||||
|
||||
def _convert_assistant_message(self, message: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""转换助手消息"""
|
||||
content = message.get("content")
|
||||
text_parts: List[str] = []
|
||||
tool_calls: List[Dict[str, Any]] = []
|
||||
|
||||
if isinstance(content, str):
|
||||
text_parts.append(content)
|
||||
else:
|
||||
for idx, block in enumerate(content or []):
|
||||
block_type = block.get("type")
|
||||
|
||||
if block_type == self.CONTENT_TYPE_TEXT:
|
||||
text_parts.append(block.get("text", ""))
|
||||
elif block_type == self.CONTENT_TYPE_TOOL_USE:
|
||||
tool_calls.append(
|
||||
{
|
||||
"id": block.get("id", f"call_{idx}"),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": block.get("name", ""),
|
||||
"arguments": json.dumps(block.get("input", {}), ensure_ascii=False),
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
result: Dict[str, Any] = {"role": "assistant"}
|
||||
|
||||
message_content = "\n".join([p for p in text_parts if p]) or None
|
||||
if message_content:
|
||||
result["content"] = message_content
|
||||
|
||||
if tool_calls:
|
||||
result["tool_calls"] = tool_calls
|
||||
|
||||
return result
|
||||
|
||||
def _convert_tools(
|
||||
self, tools: Optional[List[Dict[str, Any]]]
|
||||
) -> Optional[List[Dict[str, Any]]]:
|
||||
"""转换工具定义"""
|
||||
if not tools:
|
||||
return None
|
||||
|
||||
result: List[Dict[str, Any]] = []
|
||||
for tool in tools:
|
||||
result.append(
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tool.get("name", ""),
|
||||
"description": tool.get("description"),
|
||||
"parameters": tool.get("input_schema", {}),
|
||||
},
|
||||
}
|
||||
)
|
||||
return result
|
||||
|
||||
def _convert_tool_choice(
|
||||
self, tool_choice: Optional[Dict[str, Any]]
|
||||
) -> Optional[Union[str, Dict[str, Any]]]:
|
||||
"""转换工具选择"""
|
||||
if tool_choice is None:
|
||||
return None
|
||||
|
||||
choice_type = tool_choice.get("type")
|
||||
if choice_type in ("tool", "tool_use"):
|
||||
return {"type": "function", "function": {"name": tool_choice.get("name", "")}}
|
||||
if choice_type == "any":
|
||||
return "required"
|
||||
if choice_type == "auto":
|
||||
return "auto"
|
||||
|
||||
return tool_choice
|
||||
|
||||
# ==================== 响应转换 ====================
|
||||
|
||||
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
将 Claude 响应转换为 OpenAI 格式
|
||||
|
||||
Args:
|
||||
response: Claude 响应字典
|
||||
|
||||
Returns:
|
||||
OpenAI 格式的响应字典
|
||||
"""
|
||||
# 提取内容
|
||||
content_parts: List[str] = []
|
||||
tool_calls: List[Dict[str, Any]] = []
|
||||
|
||||
for idx, block in enumerate(response.get("content", [])):
|
||||
block_type = block.get("type")
|
||||
|
||||
if block_type == self.CONTENT_TYPE_TEXT:
|
||||
content_parts.append(block.get("text", ""))
|
||||
elif block_type == self.CONTENT_TYPE_TOOL_USE:
|
||||
tool_calls.append(
|
||||
{
|
||||
"id": block.get("id", f"call_{idx}"),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": block.get("name", ""),
|
||||
"arguments": json.dumps(block.get("input", {}), ensure_ascii=False),
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
# 构建消息
|
||||
message: Dict[str, Any] = {"role": "assistant"}
|
||||
text_content = "\n".join([p for p in content_parts if p]) or None
|
||||
if text_content:
|
||||
message["content"] = text_content
|
||||
if tool_calls:
|
||||
message["tool_calls"] = tool_calls
|
||||
|
||||
# 转换停止原因
|
||||
stop_reason = response.get("stop_reason")
|
||||
finish_reason = self.STOP_REASON_MAP.get(stop_reason, stop_reason) if stop_reason else None
|
||||
|
||||
# 转换 usage
|
||||
usage = response.get("usage", {})
|
||||
openai_usage = {
|
||||
"prompt_tokens": usage.get("input_tokens", 0),
|
||||
"completion_tokens": usage.get("output_tokens", 0),
|
||||
"total_tokens": (usage.get("input_tokens", 0) + usage.get("output_tokens", 0)),
|
||||
}
|
||||
|
||||
return {
|
||||
"id": f"chatcmpl-{response.get('id', uuid.uuid4().hex[:8])}",
|
||||
"object": "chat.completion",
|
||||
"created": int(time.time()),
|
||||
"model": response.get("model", ""),
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": message,
|
||||
"finish_reason": finish_reason,
|
||||
}
|
||||
],
|
||||
"usage": openai_usage,
|
||||
}
|
||||
|
||||
# ==================== 流式转换 ====================
|
||||
|
||||
def convert_stream_chunk(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
state: Optional["StreamConversionState"] = None,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
将 Claude SSE 事件转换为 OpenAI 格式
|
||||
|
||||
Args:
|
||||
chunk: Claude SSE 事件
|
||||
state: 流式转换状态
|
||||
|
||||
Returns:
|
||||
OpenAI 格式的 SSE chunk 列表
|
||||
"""
|
||||
from src.core.api_format.conversion.state import StreamConversionState
|
||||
|
||||
if state is None:
|
||||
state = StreamConversionState()
|
||||
|
||||
result = self._convert_single_event(chunk, state.model, state.message_id)
|
||||
if result is None:
|
||||
return []
|
||||
return [result]
|
||||
|
||||
def _convert_single_event(
|
||||
self,
|
||||
event: Dict[str, Any],
|
||||
model: str = "",
|
||||
message_id: Optional[str] = None,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
转换单个 Claude SSE 事件为 OpenAI 格式
|
||||
|
||||
Args:
|
||||
event: Claude SSE 事件
|
||||
model: 模型名称
|
||||
message_id: 消息 ID
|
||||
|
||||
Returns:
|
||||
OpenAI 格式的 SSE chunk,如果无法转换返回 None
|
||||
"""
|
||||
event_type = event.get("type")
|
||||
chunk_id = f"chatcmpl-{(message_id or 'stream')[-8:]}"
|
||||
|
||||
if event_type == "message_start":
|
||||
message = event.get("message", {})
|
||||
return self._base_chunk(
|
||||
chunk_id,
|
||||
model or message.get("model", ""),
|
||||
{"role": "assistant"},
|
||||
)
|
||||
|
||||
if event_type == "content_block_start":
|
||||
content_block = event.get("content_block", {})
|
||||
if content_block.get("type") == self.CONTENT_TYPE_TOOL_USE:
|
||||
delta = {
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": event.get("index", 0),
|
||||
"id": content_block.get("id", ""),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": content_block.get("name", ""),
|
||||
"arguments": "",
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
return self._base_chunk(chunk_id, model, delta)
|
||||
return None
|
||||
|
||||
if event_type == "content_block_delta":
|
||||
delta_payload = event.get("delta") or {}
|
||||
delta_type = delta_payload.get("type")
|
||||
|
||||
if delta_type == "text_delta":
|
||||
delta = {"content": delta_payload.get("text", "")}
|
||||
return self._base_chunk(chunk_id, model, delta)
|
||||
|
||||
if delta_type == "input_json_delta":
|
||||
delta = {
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": event.get("index", 0),
|
||||
"function": {"arguments": delta_payload.get("partial_json", "")},
|
||||
}
|
||||
]
|
||||
}
|
||||
return self._base_chunk(chunk_id, model, delta)
|
||||
return None
|
||||
|
||||
if event_type == "message_delta":
|
||||
delta = event.get("delta") or {}
|
||||
stop_reason = delta.get("stop_reason")
|
||||
finish_reason = self.STOP_REASON_MAP.get(stop_reason, stop_reason)
|
||||
return self._base_chunk(chunk_id, model, {}, finish_reason=finish_reason)
|
||||
|
||||
if event_type == "message_stop":
|
||||
return self._base_chunk(chunk_id, model, {}, finish_reason="stop")
|
||||
|
||||
return None
|
||||
|
||||
def _base_chunk(
|
||||
self,
|
||||
chunk_id: str,
|
||||
model: str,
|
||||
delta: Dict[str, Any],
|
||||
finish_reason: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""构建基础 OpenAI chunk"""
|
||||
return {
|
||||
"id": chunk_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": int(time.time()),
|
||||
"model": model,
|
||||
"system_fingerprint": None,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": delta,
|
||||
"finish_reason": finish_reason,
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
# ==================== 工具方法 ====================
|
||||
|
||||
def _extract_text_content(
|
||||
self, content: Optional[Union[str, List[Dict[str, Any]]]]
|
||||
) -> Optional[str]:
|
||||
"""提取文本内容"""
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
parts = [
|
||||
block.get("text", "")
|
||||
for block in content
|
||||
if block.get("type") == self.CONTENT_TYPE_TEXT
|
||||
]
|
||||
return "\n\n".join(filter(None, parts)) or None
|
||||
return None
|
||||
|
||||
def _render_tool_content(self, tool_content: Any) -> str:
|
||||
"""渲染工具内容"""
|
||||
if isinstance(tool_content, list):
|
||||
return json.dumps(tool_content, ensure_ascii=False)
|
||||
return str(tool_content)
|
||||
|
||||
|
||||
__all__ = ["ClaudeToOpenAIConverter"]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,492 +0,0 @@
|
||||
"""
|
||||
OpenAI -> Claude 格式转换器
|
||||
|
||||
将 OpenAI Chat Completions API 格式转换为 Claude Messages API 格式。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.core.api_format.conversion.state import StreamConversionState
|
||||
|
||||
|
||||
class OpenAIToClaudeConverter:
|
||||
"""
|
||||
OpenAI -> Claude 格式转换器
|
||||
|
||||
支持:
|
||||
- 请求转换:OpenAI Chat Request -> Claude Request
|
||||
- 响应转换:OpenAI Chat Response -> Claude Response
|
||||
- 流式转换:OpenAI SSE -> Claude SSE
|
||||
"""
|
||||
|
||||
# 内容类型常量
|
||||
CONTENT_TYPE_TEXT = "text"
|
||||
CONTENT_TYPE_IMAGE = "image"
|
||||
CONTENT_TYPE_TOOL_USE = "tool_use"
|
||||
CONTENT_TYPE_TOOL_RESULT = "tool_result"
|
||||
|
||||
# 停止原因映射(OpenAI -> Claude)
|
||||
FINISH_REASON_MAP = {
|
||||
"stop": "end_turn",
|
||||
"length": "max_tokens",
|
||||
"tool_calls": "tool_use",
|
||||
"function_call": "tool_use",
|
||||
"content_filter": "end_turn",
|
||||
}
|
||||
|
||||
def __init__(self, model_mapping: Optional[Dict[str, str]] = None):
|
||||
"""
|
||||
Args:
|
||||
model_mapping: OpenAI 模型到 Claude 模型的映射
|
||||
"""
|
||||
self._model_mapping = model_mapping or {}
|
||||
|
||||
# ==================== 请求转换 ====================
|
||||
|
||||
def convert_request(self, request: Union[Dict[str, Any], Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
将 OpenAI 请求转换为 Claude 格式
|
||||
|
||||
Args:
|
||||
request: OpenAI 请求(Dict 或 Pydantic 模型)
|
||||
|
||||
Returns:
|
||||
Claude 格式的请求字典
|
||||
"""
|
||||
if hasattr(request, "model_dump"):
|
||||
data = request.model_dump(exclude_none=True)
|
||||
else:
|
||||
data = dict(request)
|
||||
|
||||
# 模型映射
|
||||
model = data.get("model", "")
|
||||
claude_model = self._model_mapping.get(model, model)
|
||||
|
||||
# 处理消息
|
||||
system_content: Optional[str] = None
|
||||
claude_messages: List[Dict[str, Any]] = []
|
||||
|
||||
for message in data.get("messages", []):
|
||||
role = message.get("role")
|
||||
|
||||
# 提取 system 消息
|
||||
if role == "system":
|
||||
system_content = self._collapse_content(message.get("content"))
|
||||
continue
|
||||
|
||||
# 转换其他消息
|
||||
converted = self._convert_message(message)
|
||||
if converted:
|
||||
claude_messages.append(converted)
|
||||
|
||||
# 构建 Claude 请求
|
||||
result: Dict[str, Any] = {
|
||||
"model": claude_model,
|
||||
"messages": claude_messages,
|
||||
"max_tokens": data.get("max_tokens") or 4096,
|
||||
}
|
||||
|
||||
# 可选参数
|
||||
if data.get("temperature") is not None:
|
||||
result["temperature"] = data["temperature"]
|
||||
if data.get("top_p") is not None:
|
||||
result["top_p"] = data["top_p"]
|
||||
if data.get("stream"):
|
||||
result["stream"] = data["stream"]
|
||||
if data.get("stop"):
|
||||
result["stop_sequences"] = self._convert_stop(data["stop"])
|
||||
if system_content:
|
||||
result["system"] = system_content
|
||||
|
||||
# 工具转换
|
||||
tools = self._convert_tools(data.get("tools"))
|
||||
if tools:
|
||||
result["tools"] = tools
|
||||
|
||||
tool_choice = self._convert_tool_choice(data.get("tool_choice"))
|
||||
if tool_choice:
|
||||
result["tool_choice"] = tool_choice
|
||||
|
||||
return result
|
||||
|
||||
def _convert_message(self, message: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||
"""转换单条消息"""
|
||||
role = message.get("role")
|
||||
|
||||
if role == "user":
|
||||
return self._convert_user_message(message)
|
||||
if role == "assistant":
|
||||
return self._convert_assistant_message(message)
|
||||
if role == "tool":
|
||||
return self._convert_tool_message(message)
|
||||
|
||||
return None
|
||||
|
||||
def _convert_user_message(self, message: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""转换用户消息"""
|
||||
content = message.get("content")
|
||||
|
||||
if isinstance(content, str) or content is None:
|
||||
return {"role": "user", "content": content or ""}
|
||||
|
||||
# 转换内容数组
|
||||
claude_content: List[Dict[str, Any]] = []
|
||||
for item in content:
|
||||
item_type = item.get("type")
|
||||
|
||||
if item_type == "text":
|
||||
claude_content.append(
|
||||
{"type": self.CONTENT_TYPE_TEXT, "text": item.get("text", "")}
|
||||
)
|
||||
elif item_type == "image_url":
|
||||
image_url = (item.get("image_url") or {}).get("url", "")
|
||||
claude_content.append(self._convert_image_url(image_url))
|
||||
|
||||
return {"role": "user", "content": claude_content}
|
||||
|
||||
def _convert_assistant_message(self, message: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""转换助手消息"""
|
||||
content_blocks: List[Dict[str, Any]] = []
|
||||
|
||||
# 处理文本内容
|
||||
content = message.get("content")
|
||||
if isinstance(content, str):
|
||||
content_blocks.append({"type": self.CONTENT_TYPE_TEXT, "text": content})
|
||||
elif isinstance(content, list):
|
||||
for part in content:
|
||||
if part.get("type") == "text":
|
||||
content_blocks.append(
|
||||
{"type": self.CONTENT_TYPE_TEXT, "text": part.get("text", "")}
|
||||
)
|
||||
|
||||
# 处理工具调用
|
||||
for tool_call in message.get("tool_calls") or []:
|
||||
if tool_call.get("type") == "function":
|
||||
function = tool_call.get("function", {})
|
||||
arguments = function.get("arguments", "{}")
|
||||
try:
|
||||
input_data = json.loads(arguments)
|
||||
except json.JSONDecodeError:
|
||||
input_data = {"raw": arguments}
|
||||
|
||||
content_blocks.append(
|
||||
{
|
||||
"type": self.CONTENT_TYPE_TOOL_USE,
|
||||
"id": tool_call.get("id", ""),
|
||||
"name": function.get("name", ""),
|
||||
"input": input_data,
|
||||
}
|
||||
)
|
||||
|
||||
# 简化单文本内容
|
||||
if not content_blocks:
|
||||
return {"role": "assistant", "content": ""}
|
||||
if len(content_blocks) == 1 and content_blocks[0]["type"] == self.CONTENT_TYPE_TEXT:
|
||||
return {"role": "assistant", "content": content_blocks[0]["text"]}
|
||||
|
||||
return {"role": "assistant", "content": content_blocks}
|
||||
|
||||
def _convert_tool_message(self, message: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""转换工具结果消息"""
|
||||
tool_content = message.get("content", "")
|
||||
|
||||
# 尝试解析 JSON
|
||||
parsed_content = tool_content
|
||||
if isinstance(tool_content, str):
|
||||
try:
|
||||
parsed_content = json.loads(tool_content)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
tool_block = {
|
||||
"type": self.CONTENT_TYPE_TOOL_RESULT,
|
||||
"tool_use_id": message.get("tool_call_id", ""),
|
||||
"content": parsed_content,
|
||||
}
|
||||
|
||||
return {"role": "user", "content": [tool_block]}
|
||||
|
||||
def _convert_tools(
|
||||
self, tools: Optional[List[Dict[str, Any]]]
|
||||
) -> Optional[List[Dict[str, Any]]]:
|
||||
"""转换工具定义"""
|
||||
if not tools:
|
||||
return None
|
||||
|
||||
result: List[Dict[str, Any]] = []
|
||||
for tool in tools:
|
||||
if tool.get("type") != "function":
|
||||
continue
|
||||
|
||||
function = tool.get("function", {})
|
||||
result.append(
|
||||
{
|
||||
"name": function.get("name", ""),
|
||||
"description": function.get("description"),
|
||||
"input_schema": function.get("parameters") or {},
|
||||
}
|
||||
)
|
||||
|
||||
return result if result else None
|
||||
|
||||
def _convert_tool_choice(
|
||||
self, tool_choice: Optional[Union[str, Dict[str, Any]]]
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""转换工具选择"""
|
||||
if tool_choice is None:
|
||||
return None
|
||||
if tool_choice == "none":
|
||||
return {"type": "none"}
|
||||
if tool_choice == "auto":
|
||||
return {"type": "auto"}
|
||||
if tool_choice == "required":
|
||||
return {"type": "any"}
|
||||
if isinstance(tool_choice, dict) and tool_choice.get("type") == "function":
|
||||
function = tool_choice.get("function", {})
|
||||
return {"type": "tool_use", "name": function.get("name", "")}
|
||||
|
||||
return {"type": "auto"}
|
||||
|
||||
def _convert_image_url(self, image_url: str) -> Dict[str, Any]:
|
||||
"""转换图片 URL"""
|
||||
if image_url.startswith("data:"):
|
||||
header, _, data = image_url.partition(",")
|
||||
media_type = "image/jpeg"
|
||||
if ";" in header:
|
||||
media_type = header.split(";")[0].split(":")[-1]
|
||||
|
||||
return {
|
||||
"type": self.CONTENT_TYPE_IMAGE,
|
||||
"source": {
|
||||
"type": "base64",
|
||||
"media_type": media_type,
|
||||
"data": data,
|
||||
},
|
||||
}
|
||||
|
||||
return {"type": self.CONTENT_TYPE_TEXT, "text": f"[Image: {image_url}]"}
|
||||
|
||||
def _convert_stop(self, stop: Optional[Union[str, List[str]]]) -> Optional[List[str]]:
|
||||
"""转换停止序列"""
|
||||
if stop is None:
|
||||
return None
|
||||
if isinstance(stop, str):
|
||||
return [stop]
|
||||
return stop
|
||||
|
||||
# ==================== 响应转换 ====================
|
||||
|
||||
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
将 OpenAI 响应转换为 Claude 格式
|
||||
|
||||
Args:
|
||||
response: OpenAI 响应字典
|
||||
|
||||
Returns:
|
||||
Claude 格式的响应字典
|
||||
"""
|
||||
choices = response.get("choices", [])
|
||||
if not choices:
|
||||
return self._empty_claude_response(response)
|
||||
|
||||
choice = choices[0]
|
||||
message = choice.get("message", {})
|
||||
|
||||
# 构建 content 数组
|
||||
content: List[Dict[str, Any]] = []
|
||||
|
||||
# 处理文本
|
||||
text_content = message.get("content")
|
||||
if text_content:
|
||||
content.append(
|
||||
{
|
||||
"type": self.CONTENT_TYPE_TEXT,
|
||||
"text": text_content,
|
||||
}
|
||||
)
|
||||
|
||||
# 处理工具调用
|
||||
for tool_call in message.get("tool_calls") or []:
|
||||
if tool_call.get("type") == "function":
|
||||
function = tool_call.get("function", {})
|
||||
arguments = function.get("arguments", "{}")
|
||||
try:
|
||||
input_data = json.loads(arguments)
|
||||
except json.JSONDecodeError:
|
||||
input_data = {"raw": arguments}
|
||||
|
||||
content.append(
|
||||
{
|
||||
"type": self.CONTENT_TYPE_TOOL_USE,
|
||||
"id": tool_call.get("id", ""),
|
||||
"name": function.get("name", ""),
|
||||
"input": input_data,
|
||||
}
|
||||
)
|
||||
|
||||
# 转换 finish_reason
|
||||
finish_reason = choice.get("finish_reason")
|
||||
stop_reason = self.FINISH_REASON_MAP.get(finish_reason, "end_turn")
|
||||
|
||||
# 转换 usage
|
||||
usage = response.get("usage", {})
|
||||
claude_usage = {
|
||||
"input_tokens": usage.get("prompt_tokens", 0),
|
||||
"output_tokens": usage.get("completion_tokens", 0),
|
||||
}
|
||||
|
||||
return {
|
||||
"id": f"msg_{response.get('id', uuid.uuid4().hex[:8])}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": response.get("model", ""),
|
||||
"content": content,
|
||||
"stop_reason": stop_reason,
|
||||
"stop_sequence": None,
|
||||
"usage": claude_usage,
|
||||
}
|
||||
|
||||
def _empty_claude_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""构建空的 Claude 响应"""
|
||||
return {
|
||||
"id": f"msg_{response.get('id', uuid.uuid4().hex[:8])}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": response.get("model", ""),
|
||||
"content": [],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 0, "output_tokens": 0},
|
||||
}
|
||||
|
||||
# ==================== 流式转换 ====================
|
||||
|
||||
def convert_stream_chunk(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
state: Optional["StreamConversionState"] = None,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
将 OpenAI SSE chunk 转换为 Claude SSE 事件
|
||||
|
||||
Args:
|
||||
chunk: OpenAI SSE chunk
|
||||
state: 流式转换状态
|
||||
|
||||
Returns:
|
||||
Claude SSE 事件列表
|
||||
"""
|
||||
from src.core.api_format.conversion.state import StreamConversionState
|
||||
|
||||
if state is None:
|
||||
state = StreamConversionState()
|
||||
|
||||
events: List[Dict[str, Any]] = []
|
||||
|
||||
choices = chunk.get("choices") or []
|
||||
if not choices:
|
||||
return events
|
||||
|
||||
choice = choices[0]
|
||||
delta = choice.get("delta") or {}
|
||||
finish_reason = choice.get("finish_reason")
|
||||
|
||||
# 处理角色(第一个 chunk)
|
||||
role = delta.get("role")
|
||||
if role and not state.message_started:
|
||||
msg_id = state.message_id or f"msg_{uuid.uuid4().hex[:8]}"
|
||||
events.append(
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": msg_id,
|
||||
"type": "message",
|
||||
"role": role,
|
||||
"model": state.model,
|
||||
"content": [],
|
||||
"stop_reason": None,
|
||||
"stop_sequence": None,
|
||||
},
|
||||
}
|
||||
)
|
||||
state.message_started = True
|
||||
|
||||
# 处理文本内容
|
||||
content_delta = delta.get("content")
|
||||
if isinstance(content_delta, str):
|
||||
events.append(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": content_delta},
|
||||
}
|
||||
)
|
||||
|
||||
# 处理工具调用
|
||||
tool_calls = delta.get("tool_calls") or []
|
||||
for tool_call in tool_calls:
|
||||
index = tool_call.get("index", 0)
|
||||
|
||||
# 工具调用开始
|
||||
if "id" in tool_call:
|
||||
function = tool_call.get("function", {})
|
||||
events.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": index,
|
||||
"content_block": {
|
||||
"type": self.CONTENT_TYPE_TOOL_USE,
|
||||
"id": tool_call["id"],
|
||||
"name": function.get("name", ""),
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
# 工具调用参数增量
|
||||
function = tool_call.get("function", {})
|
||||
if "arguments" in function:
|
||||
events.append(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": index,
|
||||
"delta": {
|
||||
"type": "input_json_delta",
|
||||
"partial_json": function.get("arguments", ""),
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
# 处理结束
|
||||
if finish_reason:
|
||||
stop_reason = self.FINISH_REASON_MAP.get(finish_reason, "end_turn")
|
||||
events.append(
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": stop_reason},
|
||||
}
|
||||
)
|
||||
|
||||
return events
|
||||
|
||||
# ==================== 工具方法 ====================
|
||||
|
||||
def _collapse_content(
|
||||
self, content: Optional[Union[str, List[Dict[str, Any]]]]
|
||||
) -> Optional[str]:
|
||||
"""折叠内容为字符串"""
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if not content:
|
||||
return None
|
||||
|
||||
text_parts = [part.get("text", "") for part in content if part.get("type") == "text"]
|
||||
return "\n\n".join(filter(None, text_parts)) or None
|
||||
|
||||
|
||||
__all__ = ["OpenAIToClaudeConverter"]
|
||||
133
src/core/api_format/conversion/field_mappings.py
Normal file
133
src/core/api_format/conversion/field_mappings.py
Normal file
@@ -0,0 +1,133 @@
|
||||
"""
|
||||
字段映射配置(集中定义)
|
||||
|
||||
该文件用于承载:
|
||||
- role/stop_reason/usage/error 的常见映射表
|
||||
|
||||
注意:
|
||||
- conversion 层只负责 body 结构转换,不维护 model_in_body/stream_in_body/auth_header 等元数据;
|
||||
这些应复用 `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]] = {
|
||||
"OPENAI": {
|
||||
"user": "user",
|
||||
"assistant": "assistant",
|
||||
"system": "system",
|
||||
"developer": "developer",
|
||||
"tool": "tool",
|
||||
},
|
||||
"CLAUDE": {"user": "user", "assistant": "assistant"},
|
||||
"GEMINI": {"user": "user", "assistant": "model"},
|
||||
}
|
||||
|
||||
|
||||
# 停止原因映射(internal -> provider),未知值使用 UNKNOWN 并写入 extra/raw
|
||||
STOP_REASON_MAPPINGS: Dict[str, Dict[str, str]] = {
|
||||
"CLAUDE": {
|
||||
"end_turn": "end_turn",
|
||||
"max_tokens": "max_tokens",
|
||||
"stop_sequence": "stop_sequence",
|
||||
"tool_use": "tool_use",
|
||||
# Claude 通常以错误/阻断体现,这里仅兜底
|
||||
"content_filtered": "end_turn",
|
||||
"unknown": "end_turn",
|
||||
},
|
||||
"OPENAI": {
|
||||
"end_turn": "stop",
|
||||
"max_tokens": "length",
|
||||
"stop_sequence": "stop",
|
||||
"tool_use": "tool_calls",
|
||||
"content_filtered": "content_filter",
|
||||
"unknown": "stop",
|
||||
},
|
||||
"GEMINI": {
|
||||
"end_turn": "STOP",
|
||||
"max_tokens": "MAX_TOKENS",
|
||||
"stop_sequence": "STOP",
|
||||
# Gemini finishReason 对工具调用并没有稳定等价枚举,这里保守兜底为 STOP
|
||||
"tool_use": "STOP",
|
||||
"content_filtered": "SAFETY",
|
||||
"unknown": "OTHER",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
# 使用量字段映射(provider usage field -> internal UsageInfo field)
|
||||
USAGE_FIELD_MAPPINGS: Dict[str, Dict[str, str]] = {
|
||||
"CLAUDE": {
|
||||
"input_tokens": "input_tokens",
|
||||
"output_tokens": "output_tokens",
|
||||
"cache_read_input_tokens": "cache_read_tokens",
|
||||
"cache_creation_input_tokens": "cache_write_tokens",
|
||||
},
|
||||
"OPENAI": {
|
||||
"prompt_tokens": "input_tokens",
|
||||
"completion_tokens": "output_tokens",
|
||||
"total_tokens": "total_tokens",
|
||||
},
|
||||
"GEMINI": {
|
||||
"promptTokenCount": "input_tokens",
|
||||
"candidatesTokenCount": "output_tokens",
|
||||
"totalTokenCount": "total_tokens",
|
||||
"cachedContentTokenCount": "cache_read_tokens",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
# 错误类型映射(provider -> internal ErrorType.value)
|
||||
ERROR_TYPE_MAPPINGS: Dict[str, Dict[str, str]] = {
|
||||
"CLAUDE": {
|
||||
"invalid_request_error": "invalid_request",
|
||||
"authentication_error": "authentication",
|
||||
"permission_error": "permission_denied",
|
||||
"not_found_error": "not_found",
|
||||
"rate_limit_error": "rate_limit",
|
||||
"timeout_error": "server_error",
|
||||
"overloaded_error": "overloaded",
|
||||
"billing_error": "permission_denied",
|
||||
"api_error": "server_error",
|
||||
},
|
||||
"OPENAI": {
|
||||
"invalid_request_error": "invalid_request",
|
||||
"invalid_api_key": "authentication",
|
||||
"insufficient_quota": "rate_limit",
|
||||
"rate_limit_exceeded": "rate_limit",
|
||||
"server_error": "server_error",
|
||||
"context_length_exceeded": "context_length_exceeded",
|
||||
"content_policy_violation": "content_filtered",
|
||||
},
|
||||
"GEMINI": {
|
||||
"INVALID_ARGUMENT": "invalid_request",
|
||||
"UNAUTHENTICATED": "authentication",
|
||||
"PERMISSION_DENIED": "permission_denied",
|
||||
"NOT_FOUND": "not_found",
|
||||
"RESOURCE_EXHAUSTED": "rate_limit",
|
||||
"INTERNAL": "server_error",
|
||||
"UNAVAILABLE": "overloaded",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
# 可重试的错误类型(internal ErrorType.value)
|
||||
RETRYABLE_ERROR_TYPES: Set[str] = {
|
||||
"rate_limit",
|
||||
"overloaded",
|
||||
"server_error",
|
||||
}
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ROLE_MAPPINGS",
|
||||
"STOP_REASON_MAPPINGS",
|
||||
"USAGE_FIELD_MAPPINGS",
|
||||
"ERROR_TYPE_MAPPINGS",
|
||||
"RETRYABLE_ERROR_TYPES",
|
||||
]
|
||||
|
||||
297
src/core/api_format/conversion/internal.py
Normal file
297
src/core/api_format/conversion/internal.py
Normal file
@@ -0,0 +1,297 @@
|
||||
"""
|
||||
格式转换内部表示(Internal / Canonical Format)
|
||||
|
||||
该模块定义 Hub-and-Spoke 架构的“中间表示法”,用于把不同 Provider 的请求/响应/流式事件
|
||||
统一映射到稳定的内部结构,再转换为目标格式。
|
||||
|
||||
设计原则:
|
||||
- 类型安全:尽量用 dataclass + Enum 表达语义,便于 IDE/静态检查
|
||||
- 可扩展:未知/不可逆字段写入 extra/raw,避免静默丢失
|
||||
- 兼容优先:UnknownBlock 在内部保留,但默认在输出阶段丢弃(可观测、可随时调整策略)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, FrozenSet, List, Optional, Union
|
||||
|
||||
|
||||
class Role(str, Enum):
|
||||
USER = "user"
|
||||
ASSISTANT = "assistant"
|
||||
SYSTEM = "system"
|
||||
DEVELOPER = "developer"
|
||||
TOOL = "tool"
|
||||
UNKNOWN = "unknown"
|
||||
|
||||
|
||||
class ContentType(str, Enum):
|
||||
TEXT = "text"
|
||||
IMAGE = "image"
|
||||
TOOL_USE = "tool_use"
|
||||
TOOL_RESULT = "tool_result"
|
||||
UNKNOWN = "unknown"
|
||||
|
||||
|
||||
class StopReason(str, Enum):
|
||||
END_TURN = "end_turn"
|
||||
MAX_TOKENS = "max_tokens"
|
||||
STOP_SEQUENCE = "stop_sequence"
|
||||
TOOL_USE = "tool_use"
|
||||
# Claude streaming 里会出现(官方文档枚举):pause_turn / refusal
|
||||
PAUSE_TURN = "pause_turn"
|
||||
REFUSAL = "refusal"
|
||||
CONTENT_FILTERED = "content_filtered"
|
||||
UNKNOWN = "unknown"
|
||||
|
||||
|
||||
class ErrorType(str, Enum):
|
||||
INVALID_REQUEST = "invalid_request"
|
||||
AUTHENTICATION = "authentication"
|
||||
PERMISSION_DENIED = "permission_denied"
|
||||
NOT_FOUND = "not_found"
|
||||
RATE_LIMIT = "rate_limit"
|
||||
OVERLOADED = "overloaded"
|
||||
SERVER_ERROR = "server_error"
|
||||
CONTENT_FILTERED = "content_filtered"
|
||||
CONTEXT_LENGTH_EXCEEDED = "context_length_exceeded"
|
||||
UNKNOWN = "unknown"
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextBlock:
|
||||
"""文本内容块"""
|
||||
|
||||
type: ContentType = field(default=ContentType.TEXT, init=False)
|
||||
text: str = ""
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ImageBlock:
|
||||
"""图片内容块"""
|
||||
|
||||
type: ContentType = field(default=ContentType.IMAGE, init=False)
|
||||
# base64 编码的图片数据(二选一)
|
||||
data: Optional[str] = None
|
||||
media_type: Optional[str] = None
|
||||
# 或者 URL 引用
|
||||
url: Optional[str] = None
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
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)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolResultBlock:
|
||||
"""工具结果内容块"""
|
||||
|
||||
type: ContentType = field(default=ContentType.TOOL_RESULT, init=False)
|
||||
tool_use_id: str = "" # 对应的 ToolUseBlock.tool_id
|
||||
# 工具输出可能是纯文本,也可能是结构化 JSON(Gemini functionResponse 等)
|
||||
output: Any = None
|
||||
content_text: Optional[str] = None
|
||||
is_error: bool = False
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
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)
|
||||
|
||||
|
||||
ContentBlock = Union[TextBlock, ImageBlock, ToolUseBlock, ToolResultBlock, UnknownBlock]
|
||||
|
||||
|
||||
@dataclass
|
||||
class InternalMessage:
|
||||
"""统一的消息表示"""
|
||||
|
||||
role: Role
|
||||
content: List[ContentBlock] # 统一使用列表,纯文本用单个 TextBlock
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolDefinition:
|
||||
"""统一的工具定义"""
|
||||
|
||||
name: str
|
||||
description: Optional[str] = None
|
||||
parameters: Optional[Dict[str, Any]] = None # JSON Schema
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
class ToolChoiceType(str, Enum):
|
||||
AUTO = "auto"
|
||||
NONE = "none"
|
||||
REQUIRED = "required"
|
||||
TOOL = "tool"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolChoice:
|
||||
"""统一的工具选择"""
|
||||
|
||||
type: ToolChoiceType
|
||||
tool_name: Optional[str] = None
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class InstructionSegment:
|
||||
"""系统/开发者指令段(用于保留 OpenAI system/developer 结构与顺序)"""
|
||||
|
||||
role: Role # 仅允许 Role.SYSTEM / Role.DEVELOPER
|
||||
text: str = ""
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class InternalRequest:
|
||||
"""统一的请求表示"""
|
||||
|
||||
model: str
|
||||
messages: List[InternalMessage]
|
||||
|
||||
# 指令层:保留 system/developer 结构与顺序
|
||||
instructions: List[InstructionSegment] = field(default_factory=list)
|
||||
|
||||
# 兼容字段:instructions 的 join 文本(无 role 标签),用于 Claude/Gemini 这类仅接受字符串 system 的格式
|
||||
system: Optional[str] = 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
|
||||
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) # 未识别字段透传
|
||||
|
||||
def to_debug_dict(self) -> Dict[str, Any]:
|
||||
"""用于日志和调试的简化表示"""
|
||||
return {
|
||||
"model": self.model,
|
||||
"instruction_count": len(self.instructions),
|
||||
"message_count": len(self.messages),
|
||||
"has_system": bool(self.instructions) or bool(self.system),
|
||||
"max_tokens": self.max_tokens,
|
||||
"stream": self.stream,
|
||||
"tool_count": len(self.tools) if self.tools else 0,
|
||||
"extra_keys": list(self.extra.keys()),
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class UsageInfo:
|
||||
"""统一的使用量信息"""
|
||||
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
total_tokens: int = 0
|
||||
cache_read_tokens: int = 0
|
||||
cache_write_tokens: int = 0
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
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)
|
||||
|
||||
def to_debug_dict(self) -> Dict[str, Any]:
|
||||
"""用于日志和调试的简化表示"""
|
||||
usage = None
|
||||
if self.usage:
|
||||
usage = {
|
||||
"input": self.usage.input_tokens,
|
||||
"output": self.usage.output_tokens,
|
||||
}
|
||||
return {
|
||||
"id": self.id,
|
||||
"model": self.model,
|
||||
"content_block_count": len(self.content),
|
||||
"stop_reason": self.stop_reason.value if self.stop_reason else None,
|
||||
"usage": usage,
|
||||
"extra_keys": list(self.extra.keys()),
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class InternalError:
|
||||
"""统一的错误表示"""
|
||||
|
||||
type: ErrorType
|
||||
message: str
|
||||
code: Optional[str] = None # 原始错误码
|
||||
param: Optional[str] = None # 导致错误的参数
|
||||
retryable: bool = False # 是否可重试
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def to_debug_dict(self) -> Dict[str, Any]:
|
||||
"""用于日志和调试"""
|
||||
return {
|
||||
"type": self.type.value,
|
||||
"message": self.message,
|
||||
"code": self.code,
|
||||
"param": self.param,
|
||||
"retryable": self.retryable,
|
||||
"extra": self.extra,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FormatCapabilities:
|
||||
supports_stream: bool = True
|
||||
supports_error_conversion: bool = True
|
||||
supports_tools: bool = True
|
||||
supports_images: bool = False
|
||||
supported_features: FrozenSet[str] = field(default_factory=frozenset)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"Role",
|
||||
"ContentType",
|
||||
"StopReason",
|
||||
"ErrorType",
|
||||
"ToolChoiceType",
|
||||
"TextBlock",
|
||||
"ImageBlock",
|
||||
"ToolUseBlock",
|
||||
"ToolResultBlock",
|
||||
"UnknownBlock",
|
||||
"ContentBlock",
|
||||
"InternalMessage",
|
||||
"InstructionSegment",
|
||||
"ToolDefinition",
|
||||
"ToolChoice",
|
||||
"InternalRequest",
|
||||
"UsageInfo",
|
||||
"InternalResponse",
|
||||
"InternalError",
|
||||
"FormatCapabilities",
|
||||
]
|
||||
|
||||
84
src/core/api_format/conversion/normalizer.py
Normal file
84
src/core/api_format/conversion/normalizer.py
Normal file
@@ -0,0 +1,84 @@
|
||||
"""
|
||||
格式标准化器接口(FormatNormalizer)
|
||||
|
||||
每个格式(OpenAI/Claude/Gemini)实现一个 Normalizer,将 provider 结构转换到 internal,
|
||||
再从 internal 输出到目标格式。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from .internal import FormatCapabilities, InternalError, InternalRequest, InternalResponse
|
||||
from .stream_events import InternalStreamEvent
|
||||
from .stream_state import StreamState
|
||||
|
||||
|
||||
class FormatNormalizer(ABC):
|
||||
"""格式标准化器基类"""
|
||||
|
||||
FORMAT_ID: str # 如 "CLAUDE", "OPENAI", "GEMINI"
|
||||
capabilities: FormatCapabilities
|
||||
|
||||
# ============ 请求转换 ============
|
||||
|
||||
@abstractmethod
|
||||
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
|
||||
"""将格式特定请求转换为内部表示"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
|
||||
"""将内部表示转换为格式特定请求"""
|
||||
raise NotImplementedError
|
||||
|
||||
# ============ 响应转换 ============
|
||||
|
||||
@abstractmethod
|
||||
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse:
|
||||
"""将格式特定响应转换为内部表示"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def response_from_internal(self, internal: InternalResponse) -> Dict[str, Any]:
|
||||
"""将内部表示转换为格式特定响应"""
|
||||
raise NotImplementedError
|
||||
|
||||
# ============ 流式转换(可选) ============
|
||||
|
||||
def stream_chunk_to_internal(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
state: StreamState,
|
||||
) -> List[InternalStreamEvent]:
|
||||
"""将格式特定流式块转换为内部事件"""
|
||||
raise NotImplementedError
|
||||
|
||||
def stream_event_from_internal(
|
||||
self,
|
||||
event: InternalStreamEvent,
|
||||
state: StreamState,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""将内部事件转换为格式特定流式块"""
|
||||
raise NotImplementedError
|
||||
|
||||
# ============ 错误转换(可选) ============
|
||||
|
||||
def is_error_response(self, response: Dict[str, Any]) -> bool:
|
||||
"""基于 body 的兜底判断(不可靠),子类可覆盖"""
|
||||
return False
|
||||
|
||||
def error_to_internal(self, error_response: Dict[str, Any]) -> InternalError:
|
||||
"""将格式特定错误转换为内部表示"""
|
||||
raise NotImplementedError
|
||||
|
||||
def error_from_internal(self, internal: InternalError) -> Dict[str, Any]:
|
||||
"""将内部错误表示转换为格式特定错误"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
__all__ = [
|
||||
"FormatNormalizer",
|
||||
]
|
||||
|
||||
12
src/core/api_format/conversion/normalizers/__init__.py
Normal file
12
src/core/api_format/conversion/normalizers/__init__.py
Normal file
@@ -0,0 +1,12 @@
|
||||
"""
|
||||
Normalizers
|
||||
|
||||
实现各格式 <-> internal 的标准化器。
|
||||
|
||||
本目录在 Phase 1 仅创建结构;具体实现将在 Phase 2+ 补齐。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
__all__: list[str] = []
|
||||
|
||||
993
src/core/api_format/conversion/normalizers/claude.py
Normal file
993
src/core/api_format/conversion/normalizers/claude.py
Normal file
@@ -0,0 +1,993 @@
|
||||
"""
|
||||
Claude Messages API Normalizer
|
||||
|
||||
负责:
|
||||
- Claude Messages request/response <-> Internal 表示转换
|
||||
- 可选:Claude streaming event <-> InternalStreamEvent
|
||||
- 可选:Claude error <-> InternalError
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from src.core.api_format.conversion.field_mappings import (
|
||||
ERROR_TYPE_MAPPINGS,
|
||||
RETRYABLE_ERROR_TYPES,
|
||||
STOP_REASON_MAPPINGS,
|
||||
USAGE_FIELD_MAPPINGS,
|
||||
)
|
||||
from src.core.api_format.conversion.internal import (
|
||||
ContentBlock,
|
||||
ContentType,
|
||||
ErrorType,
|
||||
FormatCapabilities,
|
||||
ImageBlock,
|
||||
InstructionSegment,
|
||||
InternalError,
|
||||
InternalMessage,
|
||||
InternalRequest,
|
||||
InternalResponse,
|
||||
Role,
|
||||
StopReason,
|
||||
TextBlock,
|
||||
ToolChoice,
|
||||
ToolChoiceType,
|
||||
ToolDefinition,
|
||||
ToolResultBlock,
|
||||
ToolUseBlock,
|
||||
UnknownBlock,
|
||||
UsageInfo,
|
||||
)
|
||||
from src.core.api_format.conversion.normalizer import FormatNormalizer
|
||||
from src.core.api_format.conversion.stream_events import (
|
||||
ContentBlockStartEvent,
|
||||
ContentBlockStopEvent,
|
||||
ContentDeltaEvent,
|
||||
ErrorEvent,
|
||||
InternalStreamEvent,
|
||||
MessageStartEvent,
|
||||
MessageStopEvent,
|
||||
ToolCallDeltaEvent,
|
||||
)
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
class ClaudeNormalizer(FormatNormalizer):
|
||||
FORMAT_ID = "CLAUDE"
|
||||
capabilities = FormatCapabilities(
|
||||
supports_stream=True,
|
||||
supports_error_conversion=True,
|
||||
supports_tools=True,
|
||||
supports_images=True,
|
||||
)
|
||||
|
||||
_CLAUDE_STOP_TO_INTERNAL: Dict[str, StopReason] = {
|
||||
"end_turn": StopReason.END_TURN,
|
||||
"max_tokens": StopReason.MAX_TOKENS,
|
||||
"stop_sequence": StopReason.STOP_SEQUENCE,
|
||||
"tool_use": StopReason.TOOL_USE,
|
||||
"pause_turn": StopReason.PAUSE_TURN,
|
||||
"refusal": StopReason.REFUSAL,
|
||||
"content_filtered": StopReason.CONTENT_FILTERED,
|
||||
}
|
||||
|
||||
_ERROR_TYPE_TO_CLAUDE: Dict[ErrorType, str] = {
|
||||
ErrorType.INVALID_REQUEST: "invalid_request_error",
|
||||
ErrorType.AUTHENTICATION: "authentication_error",
|
||||
ErrorType.PERMISSION_DENIED: "permission_error",
|
||||
ErrorType.NOT_FOUND: "not_found_error",
|
||||
ErrorType.RATE_LIMIT: "rate_limit_error",
|
||||
ErrorType.OVERLOADED: "overloaded_error",
|
||||
ErrorType.SERVER_ERROR: "api_error",
|
||||
ErrorType.CONTENT_FILTERED: "invalid_request_error",
|
||||
ErrorType.CONTEXT_LENGTH_EXCEEDED: "invalid_request_error",
|
||||
ErrorType.UNKNOWN: "api_error",
|
||||
}
|
||||
|
||||
# =========================
|
||||
# Requests
|
||||
# =========================
|
||||
|
||||
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
|
||||
model = str(request.get("model") or "")
|
||||
dropped: Dict[str, int] = {}
|
||||
|
||||
instructions: List[InstructionSegment] = []
|
||||
|
||||
# 顶层 system 先进入 instructions(保持确定性优先级)
|
||||
sys_value = request.get("system")
|
||||
sys_text, sys_dropped = self._collapse_claude_system(sys_value)
|
||||
self._merge_dropped(dropped, sys_dropped)
|
||||
if sys_text:
|
||||
instructions.append(InstructionSegment(role=Role.SYSTEM, text=sys_text))
|
||||
|
||||
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
|
||||
continue
|
||||
|
||||
role = str(msg.get("role") or "unknown")
|
||||
|
||||
# 兼容:少数客户端可能把 system/developer 混进 messages[]
|
||||
if role in ("system", "developer"):
|
||||
text, md = self._collapse_claude_system(msg.get("content"))
|
||||
self._merge_dropped(dropped, md)
|
||||
if text:
|
||||
instructions.append(
|
||||
InstructionSegment(
|
||||
role=Role.SYSTEM if role == "system" else Role.DEVELOPER,
|
||||
text=text,
|
||||
extra=self._extract_extra(msg, {"role", "content"}),
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
imsg, md = self._claude_message_to_internal(msg)
|
||||
self._merge_dropped(dropped, md)
|
||||
if imsg is not None:
|
||||
messages.append(imsg)
|
||||
|
||||
system_text = self._join_instructions(instructions)
|
||||
|
||||
tools = self._claude_tools_to_internal(request.get("tools"))
|
||||
tool_choice = self._claude_tool_choice_to_internal(request.get("tool_choice"))
|
||||
|
||||
internal = InternalRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
instructions=instructions,
|
||||
system=system_text,
|
||||
max_tokens=self._optional_int(request.get("max_tokens")),
|
||||
temperature=self._optional_float(request.get("temperature")),
|
||||
top_p=self._optional_float(request.get("top_p")),
|
||||
top_k=self._optional_int(request.get("top_k")),
|
||||
stop_sequences=self._coerce_str_list(request.get("stop_sequences")),
|
||||
stream=bool(request.get("stream") or False),
|
||||
tools=tools,
|
||||
tool_choice=tool_choice,
|
||||
extra={"claude": self._extract_extra(request, {"messages"})},
|
||||
)
|
||||
|
||||
if dropped:
|
||||
internal.extra.setdefault("raw", {})["dropped_blocks"] = dropped
|
||||
|
||||
return internal
|
||||
|
||||
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]
|
||||
|
||||
result: Dict[str, Any] = {
|
||||
"model": internal.model,
|
||||
"messages": out_messages,
|
||||
"max_tokens": internal.max_tokens if internal.max_tokens is not None else 4096,
|
||||
}
|
||||
|
||||
if system_text:
|
||||
result["system"] = system_text
|
||||
|
||||
if internal.temperature is not None:
|
||||
result["temperature"] = internal.temperature
|
||||
if internal.top_p is not None:
|
||||
result["top_p"] = internal.top_p
|
||||
if internal.top_k is not None:
|
||||
result["top_k"] = internal.top_k
|
||||
if internal.stop_sequences:
|
||||
result["stop_sequences"] = list(internal.stop_sequences)
|
||||
if internal.stream:
|
||||
result["stream"] = True
|
||||
|
||||
if internal.tools:
|
||||
result["tools"] = [
|
||||
{
|
||||
"name": t.name,
|
||||
"description": t.description,
|
||||
"input_schema": t.parameters or {},
|
||||
**(t.extra.get("claude") or {}),
|
||||
}
|
||||
for t in internal.tools
|
||||
]
|
||||
|
||||
if internal.tool_choice:
|
||||
result["tool_choice"] = self._tool_choice_to_claude(internal.tool_choice)
|
||||
|
||||
# 恢复 Claude 特有字段(如 metadata)
|
||||
claude_extra = internal.extra.get("claude") if isinstance(internal.extra, dict) else None
|
||||
if isinstance(claude_extra, dict):
|
||||
if "metadata" in claude_extra:
|
||||
result["metadata"] = claude_extra["metadata"]
|
||||
|
||||
return result
|
||||
|
||||
# =========================
|
||||
# Responses
|
||||
# =========================
|
||||
|
||||
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
|
||||
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] = {}
|
||||
if raw_stop is not None:
|
||||
extra.setdefault("raw", {})["stop_reason"] = raw_stop
|
||||
|
||||
internal = InternalResponse(
|
||||
id=rid,
|
||||
model=model,
|
||||
content=blocks,
|
||||
stop_reason=stop_reason,
|
||||
usage=usage_info,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
if dropped:
|
||||
internal.extra.setdefault("raw", {})["dropped_blocks"] = dropped
|
||||
|
||||
return internal
|
||||
|
||||
def response_from_internal(self, internal: InternalResponse) -> Dict[str, Any]:
|
||||
cid = internal.id or "unknown"
|
||||
if not cid.startswith("msg_"):
|
||||
cid = f"msg_{cid}"
|
||||
|
||||
content: List[Dict[str, Any]] = []
|
||||
for b in internal.content:
|
||||
if isinstance(b, TextBlock):
|
||||
if b.text:
|
||||
content.append({"type": "text", "text": b.text})
|
||||
continue
|
||||
if isinstance(b, ToolUseBlock):
|
||||
content.append(
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": b.tool_id,
|
||||
"name": b.tool_name,
|
||||
"input": b.tool_input or {},
|
||||
}
|
||||
)
|
||||
continue
|
||||
if isinstance(b, ImageBlock):
|
||||
if b.data and b.media_type:
|
||||
content.append(
|
||||
{
|
||||
"type": "image",
|
||||
"source": {
|
||||
"type": "base64",
|
||||
"media_type": b.media_type,
|
||||
"data": b.data,
|
||||
},
|
||||
}
|
||||
)
|
||||
elif b.url:
|
||||
content.append({"type": "text", "text": f"[Image: {b.url}]"})
|
||||
continue
|
||||
# Unknown/ToolResult 默认丢弃
|
||||
|
||||
stop_reason = None
|
||||
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}
|
||||
if internal.usage:
|
||||
usage = {
|
||||
"input_tokens": int(internal.usage.input_tokens),
|
||||
"output_tokens": int(internal.usage.output_tokens),
|
||||
}
|
||||
if internal.usage.cache_read_tokens:
|
||||
usage["cache_read_input_tokens"] = int(internal.usage.cache_read_tokens)
|
||||
if internal.usage.cache_write_tokens:
|
||||
usage["cache_creation_input_tokens"] = int(internal.usage.cache_write_tokens)
|
||||
|
||||
return {
|
||||
"id": cid,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": internal.model,
|
||||
"content": content,
|
||||
"stop_reason": stop_reason,
|
||||
"stop_sequence": None,
|
||||
"usage": usage,
|
||||
}
|
||||
|
||||
# =========================
|
||||
# Streaming (Claude SSE events)
|
||||
# =========================
|
||||
|
||||
def stream_chunk_to_internal(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
state: StreamState,
|
||||
) -> List[InternalStreamEvent]:
|
||||
ss = state.substate(self.FORMAT_ID)
|
||||
events: List[InternalStreamEvent] = []
|
||||
|
||||
event_type = chunk.get("type")
|
||||
if event_type is None:
|
||||
return events
|
||||
event_type = str(event_type)
|
||||
|
||||
if event_type == "ping":
|
||||
return events
|
||||
|
||||
if event_type == "message_start":
|
||||
message_raw = chunk.get("message")
|
||||
message: Dict[str, Any] = message_raw if isinstance(message_raw, dict) else {}
|
||||
msg_id = str(message.get("id") or "")
|
||||
model = str(message.get("model") or "")
|
||||
state.message_id = msg_id or state.message_id
|
||||
state.model = model or state.model
|
||||
ss["message_started"] = True
|
||||
ss.setdefault("block_index_to_tool_id", {})
|
||||
events.append(MessageStartEvent(message_id=msg_id, model=model))
|
||||
return events
|
||||
|
||||
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 {}
|
||||
btype = str(block.get("type") or "unknown")
|
||||
|
||||
if btype == "text":
|
||||
events.append(ContentBlockStartEvent(block_index=index, block_type=ContentType.TEXT))
|
||||
return events
|
||||
|
||||
if btype == "tool_use":
|
||||
tool_id = str(block.get("id") or "")
|
||||
tool_name = str(block.get("name") or "")
|
||||
mapping = ss.get("block_index_to_tool_id")
|
||||
if isinstance(mapping, dict):
|
||||
mapping[index] = tool_id
|
||||
events.append(
|
||||
ContentBlockStartEvent(
|
||||
block_index=index,
|
||||
block_type=ContentType.TOOL_USE,
|
||||
tool_id=tool_id or None,
|
||||
tool_name=tool_name or None,
|
||||
)
|
||||
)
|
||||
return events
|
||||
|
||||
events.append(
|
||||
ContentBlockStartEvent(
|
||||
block_index=index,
|
||||
block_type=ContentType.UNKNOWN,
|
||||
extra={"raw": {"claude_block_type": btype, "content_block": block}},
|
||||
)
|
||||
)
|
||||
return events
|
||||
|
||||
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 {}
|
||||
dtype = str(delta.get("type") or "unknown")
|
||||
|
||||
if dtype == "text_delta":
|
||||
text = delta.get("text")
|
||||
if text is None:
|
||||
return events
|
||||
events.append(ContentDeltaEvent(block_index=index, text_delta=str(text)))
|
||||
return events
|
||||
|
||||
if dtype == "input_json_delta":
|
||||
partial = delta.get("partial_json")
|
||||
if partial is None:
|
||||
return events
|
||||
mapping = ss.get("block_index_to_tool_id")
|
||||
tool_id = ""
|
||||
if isinstance(mapping, dict):
|
||||
tool_id = str(mapping.get(index) or "")
|
||||
events.append(ToolCallDeltaEvent(block_index=index, tool_id=tool_id, input_delta=str(partial)))
|
||||
return events
|
||||
|
||||
return events
|
||||
|
||||
if event_type == "content_block_stop":
|
||||
index = int(chunk.get("index") or 0)
|
||||
events.append(ContentBlockStopEvent(block_index=index))
|
||||
return events
|
||||
|
||||
if event_type == "message_delta":
|
||||
delta_raw2 = chunk.get("delta")
|
||||
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)
|
||||
usage = chunk.get("usage")
|
||||
if isinstance(usage, dict):
|
||||
ss["usage"] = usage
|
||||
return events
|
||||
|
||||
if event_type == "message_stop":
|
||||
raw_stop = ss.get("stop_reason")
|
||||
stop_reason: Optional[StopReason] = 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"))
|
||||
events.append(MessageStopEvent(stop_reason=stop_reason, usage=usage_info))
|
||||
return events
|
||||
|
||||
if event_type == "error":
|
||||
internal_error = self.error_to_internal(chunk)
|
||||
events.append(ErrorEvent(error=internal_error))
|
||||
return events
|
||||
|
||||
return events
|
||||
|
||||
def stream_event_from_internal(
|
||||
self,
|
||||
event: InternalStreamEvent,
|
||||
state: StreamState,
|
||||
) -> List[Dict[str, Any]]:
|
||||
ss = state.substate(self.FORMAT_ID)
|
||||
out: List[Dict[str, Any]] = []
|
||||
|
||||
if isinstance(event, MessageStartEvent):
|
||||
state.message_id = event.message_id or state.message_id
|
||||
state.model = event.model or state.model
|
||||
ss.setdefault("block_index_to_tool_id", {})
|
||||
message_obj: Dict[str, Any] = {
|
||||
"id": state.message_id or "msg_stream",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": state.model or "",
|
||||
"content": [],
|
||||
"stop_reason": None,
|
||||
"stop_sequence": None,
|
||||
}
|
||||
# Claude CLI 客户端对接口格式要求严格,必须包含 usage 字段
|
||||
# 即使没有 usage 信息也要提供默认值
|
||||
if event.usage:
|
||||
message_obj["usage"] = self._usage_to_claude(event.usage)
|
||||
else:
|
||||
message_obj["usage"] = {"input_tokens": 0, "output_tokens": 0}
|
||||
out.append({"type": "message_start", "message": message_obj})
|
||||
return out
|
||||
|
||||
if isinstance(event, ContentBlockStartEvent):
|
||||
if event.block_type == ContentType.TEXT:
|
||||
out.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": int(event.block_index),
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
}
|
||||
)
|
||||
return out
|
||||
|
||||
if event.block_type == ContentType.TOOL_USE:
|
||||
tool_id = event.tool_id or ""
|
||||
tool_name = event.tool_name or ""
|
||||
mapping = ss.get("block_index_to_tool_id")
|
||||
if isinstance(mapping, dict):
|
||||
mapping[int(event.block_index)] = tool_id
|
||||
out.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": int(event.block_index),
|
||||
"content_block": {"type": "tool_use", "id": tool_id, "name": tool_name},
|
||||
}
|
||||
)
|
||||
return out
|
||||
|
||||
return out
|
||||
|
||||
if isinstance(event, ContentDeltaEvent):
|
||||
if event.text_delta:
|
||||
out.append(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": int(event.block_index),
|
||||
"delta": {"type": "text_delta", "text": event.text_delta},
|
||||
}
|
||||
)
|
||||
return out
|
||||
|
||||
if isinstance(event, ToolCallDeltaEvent):
|
||||
if event.input_delta:
|
||||
out.append(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": int(event.block_index),
|
||||
"delta": {"type": "input_json_delta", "partial_json": event.input_delta},
|
||||
}
|
||||
)
|
||||
return out
|
||||
|
||||
if isinstance(event, ContentBlockStopEvent):
|
||||
out.append({"type": "content_block_stop", "index": int(event.block_index)})
|
||||
return out
|
||||
|
||||
if isinstance(event, MessageStopEvent):
|
||||
stop_reason = None
|
||||
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] = {
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": stop_reason},
|
||||
}
|
||||
|
||||
if event.usage:
|
||||
msg_delta["usage"] = self._usage_to_claude(event.usage)
|
||||
|
||||
out.append(msg_delta)
|
||||
out.append({"type": "message_stop"})
|
||||
return out
|
||||
|
||||
if isinstance(event, ErrorEvent):
|
||||
out.append(self.error_from_internal(event.error))
|
||||
return out
|
||||
|
||||
return out
|
||||
|
||||
# =========================
|
||||
# Error conversion
|
||||
# =========================
|
||||
|
||||
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] = {}
|
||||
if isinstance(error_response, dict):
|
||||
err_raw = error_response.get("error")
|
||||
err = err_raw if isinstance(err_raw, dict) else {}
|
||||
|
||||
raw_type = err.get("type")
|
||||
mapped = ERROR_TYPE_MAPPINGS.get("CLAUDE", {}).get(str(raw_type), ErrorType.UNKNOWN.value)
|
||||
internal_type = self._error_type_from_value(mapped)
|
||||
retryable = internal_type.value in RETRYABLE_ERROR_TYPES
|
||||
|
||||
return InternalError(
|
||||
type=internal_type,
|
||||
message=str(err.get("message") or ""),
|
||||
code=err.get("code") if err.get("code") is None else str(err.get("code")),
|
||||
param=err.get("param") if err.get("param") is None else str(err.get("param")),
|
||||
retryable=retryable,
|
||||
extra={"claude": {"error": err}, "raw": {"type": raw_type}},
|
||||
)
|
||||
|
||||
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}
|
||||
if internal.param is not None:
|
||||
payload["param"] = internal.param
|
||||
if internal.code is not None:
|
||||
payload["code"] = internal.code
|
||||
return {"type": "error", "error": payload}
|
||||
|
||||
# =========================
|
||||
# Helpers
|
||||
# =========================
|
||||
|
||||
def _claude_message_to_internal(self, msg: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]:
|
||||
dropped: Dict[str, int] = {}
|
||||
role_raw = str(msg.get("role") or "unknown")
|
||||
|
||||
if role_raw == "user":
|
||||
role = Role.USER
|
||||
elif role_raw == "assistant":
|
||||
role = Role.ASSISTANT
|
||||
else:
|
||||
role = Role.UNKNOWN
|
||||
|
||||
blocks, bd = self._claude_content_to_blocks(msg.get("content"))
|
||||
self._merge_dropped(dropped, bd)
|
||||
|
||||
return (
|
||||
InternalMessage(
|
||||
role=role,
|
||||
content=blocks,
|
||||
extra=self._extract_extra(msg, {"role", "content"}),
|
||||
),
|
||||
dropped,
|
||||
)
|
||||
|
||||
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):
|
||||
return ([TextBlock(text=content)] if content else []), dropped
|
||||
if not isinstance(content, list):
|
||||
dropped["claude_content_non_list"] = dropped.get("claude_content_non_list", 0) + 1
|
||||
return [], dropped
|
||||
|
||||
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
|
||||
continue
|
||||
|
||||
btype = str(block.get("type") or "unknown")
|
||||
if btype == "text":
|
||||
text = str(block.get("text") or "")
|
||||
if text:
|
||||
blocks.append(TextBlock(text=text, extra=self._extract_extra(block, {"type", "text"})))
|
||||
continue
|
||||
|
||||
if btype == "image":
|
||||
src_raw = block.get("source")
|
||||
src: Dict[str, Any] = src_raw if isinstance(src_raw, dict) else {}
|
||||
stype = src.get("type")
|
||||
if stype == "base64":
|
||||
data = src.get("data")
|
||||
media_type = src.get("media_type")
|
||||
if isinstance(data, str) and data and isinstance(media_type, str) and media_type:
|
||||
blocks.append(ImageBlock(data=data, media_type=media_type))
|
||||
continue
|
||||
dropped["claude_image_unsupported"] = dropped.get("claude_image_unsupported", 0) + 1
|
||||
blocks.append(UnknownBlock(raw_type="image", payload=block))
|
||||
continue
|
||||
|
||||
if btype == "tool_use":
|
||||
tool_id = str(block.get("id") or "")
|
||||
tool_name = str(block.get("name") or "")
|
||||
tool_input = block.get("input")
|
||||
if not isinstance(tool_input, dict):
|
||||
tool_input = {"raw": tool_input}
|
||||
blocks.append(
|
||||
ToolUseBlock(
|
||||
tool_id=tool_id,
|
||||
tool_name=tool_name,
|
||||
tool_input=tool_input,
|
||||
extra={"claude": self._extract_extra(block, {"type", "id", "name", "input"})},
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
if btype == "tool_result":
|
||||
tool_use_id = str(block.get("tool_use_id") or "")
|
||||
is_error = bool(block.get("is_error") or False)
|
||||
raw_content = block.get("content")
|
||||
blocks.append(self._tool_result_from_claude(tool_use_id, raw_content, is_error, block))
|
||||
continue
|
||||
|
||||
dropped_key = f"claude_block:{btype}"
|
||||
dropped[dropped_key] = dropped.get(dropped_key, 0) + 1
|
||||
blocks.append(UnknownBlock(raw_type=btype, payload=block))
|
||||
|
||||
return blocks, dropped
|
||||
|
||||
def _tool_result_from_claude(
|
||||
self,
|
||||
tool_use_id: str,
|
||||
raw_content: Any,
|
||||
is_error: bool,
|
||||
raw_block: Dict[str, Any],
|
||||
) -> ToolResultBlock:
|
||||
if raw_content is None:
|
||||
return ToolResultBlock(
|
||||
tool_use_id=tool_use_id,
|
||||
output=None,
|
||||
content_text=None,
|
||||
is_error=is_error,
|
||||
extra={"claude": raw_block},
|
||||
)
|
||||
|
||||
if isinstance(raw_content, str):
|
||||
parsed: Any = None
|
||||
try:
|
||||
parsed = json.loads(raw_content)
|
||||
except json.JSONDecodeError:
|
||||
parsed = None
|
||||
|
||||
if parsed is not None:
|
||||
return ToolResultBlock(
|
||||
tool_use_id=tool_use_id,
|
||||
output=parsed,
|
||||
content_text=None,
|
||||
is_error=is_error,
|
||||
extra={"raw": {"content": raw_content}, "claude": raw_block},
|
||||
)
|
||||
|
||||
return ToolResultBlock(
|
||||
tool_use_id=tool_use_id,
|
||||
output=None,
|
||||
content_text=raw_content,
|
||||
is_error=is_error,
|
||||
extra={"claude": raw_block},
|
||||
)
|
||||
|
||||
if isinstance(raw_content, list):
|
||||
text_parts: List[str] = []
|
||||
for part in raw_content:
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
text = part.get("text")
|
||||
if text:
|
||||
text_parts.append(str(text))
|
||||
collapsed = "\n\n".join(text_parts) if text_parts else None
|
||||
|
||||
return ToolResultBlock(
|
||||
tool_use_id=tool_use_id,
|
||||
output=None,
|
||||
content_text=collapsed,
|
||||
is_error=is_error,
|
||||
extra={"raw": {"content": raw_content}, "claude": raw_block},
|
||||
)
|
||||
|
||||
return ToolResultBlock(
|
||||
tool_use_id=tool_use_id,
|
||||
output=raw_content,
|
||||
content_text=None,
|
||||
is_error=is_error,
|
||||
extra={"claude": raw_block},
|
||||
)
|
||||
|
||||
def _collapse_claude_system(self, system_value: Any) -> Tuple[Optional[str], 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] = []
|
||||
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
|
||||
continue
|
||||
if item.get("type") == "text":
|
||||
text = item.get("text")
|
||||
if text:
|
||||
texts.append(str(text))
|
||||
else:
|
||||
dropped_key = f"claude_system_item:{item.get('type')}"
|
||||
dropped[dropped_key] = dropped.get(dropped_key, 0) + 1
|
||||
joined = "\n\n".join(texts)
|
||||
return (joined or None), dropped
|
||||
|
||||
dropped["claude_system_unsupported"] = dropped.get("claude_system_unsupported", 0) + 1
|
||||
return None, dropped
|
||||
|
||||
def _join_instructions(self, instructions: List[InstructionSegment]) -> Optional[str]:
|
||||
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]]:
|
||||
if not tools or not isinstance(tools, list):
|
||||
return None
|
||||
|
||||
out: List[ToolDefinition] = []
|
||||
for tool in tools:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
name = str(tool.get("name") or "")
|
||||
if not name:
|
||||
continue
|
||||
out.append(
|
||||
ToolDefinition(
|
||||
name=name,
|
||||
description=tool.get("description"),
|
||||
parameters=tool.get("input_schema") if isinstance(tool.get("input_schema"), dict) else None,
|
||||
extra={"claude": self._extract_extra(tool, {"name", "description", "input_schema"})},
|
||||
)
|
||||
)
|
||||
return out or None
|
||||
|
||||
def _claude_tool_choice_to_internal(self, tool_choice: Any) -> Optional[ToolChoice]:
|
||||
if tool_choice is None:
|
||||
return None
|
||||
if not isinstance(tool_choice, dict):
|
||||
return ToolChoice(type=ToolChoiceType.AUTO, extra={"raw": tool_choice})
|
||||
|
||||
ctype = str(tool_choice.get("type") or "auto")
|
||||
if ctype == "none":
|
||||
return ToolChoice(type=ToolChoiceType.NONE, extra={"claude": tool_choice})
|
||||
if ctype == "auto":
|
||||
return ToolChoice(type=ToolChoiceType.AUTO, extra={"claude": tool_choice})
|
||||
if ctype in ("any", "required"):
|
||||
return ToolChoice(type=ToolChoiceType.REQUIRED, extra={"claude": tool_choice})
|
||||
if ctype in ("tool_use", "tool"):
|
||||
name = str(tool_choice.get("name") or "")
|
||||
return ToolChoice(type=ToolChoiceType.TOOL, tool_name=name, 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]:
|
||||
if tool_choice.type == ToolChoiceType.NONE:
|
||||
return {"type": "none"}
|
||||
if tool_choice.type == ToolChoiceType.AUTO:
|
||||
return {"type": "auto"}
|
||||
if tool_choice.type == ToolChoiceType.REQUIRED:
|
||||
return {"type": "any"}
|
||||
if tool_choice.type == ToolChoiceType.TOOL:
|
||||
return {"type": "tool_use", "name": tool_choice.tool_name or ""}
|
||||
return {"type": "auto"}
|
||||
|
||||
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] = []
|
||||
|
||||
for b in msg.content:
|
||||
if isinstance(b, UnknownBlock):
|
||||
continue
|
||||
|
||||
if isinstance(b, TextBlock):
|
||||
if b.text:
|
||||
text_parts.append(b.text)
|
||||
continue
|
||||
|
||||
if isinstance(b, ImageBlock):
|
||||
if role != "user":
|
||||
if b.url:
|
||||
text_parts.append(f"[Image: {b.url}]")
|
||||
elif b.media_type and b.data:
|
||||
text_parts.append("[Image]")
|
||||
continue
|
||||
|
||||
if b.data and b.media_type:
|
||||
blocks.append(
|
||||
{
|
||||
"type": "image",
|
||||
"source": {"type": "base64", "media_type": b.media_type, "data": b.data},
|
||||
}
|
||||
)
|
||||
elif b.url:
|
||||
text_parts.append(f"[Image: {b.url}]")
|
||||
continue
|
||||
|
||||
if isinstance(b, ToolUseBlock) and role == "assistant":
|
||||
blocks.append(
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": b.tool_id,
|
||||
"name": b.tool_name,
|
||||
"input": b.tool_input or {},
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
if isinstance(b, ToolResultBlock) and role == "user":
|
||||
if b.content_text is not None:
|
||||
content: Any = b.content_text
|
||||
elif b.output is None:
|
||||
content = ""
|
||||
elif isinstance(b.output, str):
|
||||
content = b.output
|
||||
else:
|
||||
content = b.output
|
||||
|
||||
blocks.append(
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": b.tool_use_id,
|
||||
"content": content,
|
||||
"is_error": bool(b.is_error),
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
if text_parts:
|
||||
if blocks:
|
||||
blocks = [{"type": "text", "text": "\n".join(text_parts)}] + blocks
|
||||
else:
|
||||
return {"role": role, "content": "\n".join(text_parts)}
|
||||
|
||||
return {"role": role, "content": blocks}
|
||||
|
||||
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):
|
||||
role = Role.USER
|
||||
normalized.append(InternalMessage(role=role, content=m.content, extra=m.extra))
|
||||
|
||||
if not normalized:
|
||||
return []
|
||||
|
||||
if normalized[0].role != Role.USER:
|
||||
normalized = [InternalMessage(role=Role.USER, content=[])] + normalized
|
||||
|
||||
merged: List[InternalMessage] = []
|
||||
for m in normalized:
|
||||
if merged and merged[-1].role == m.role:
|
||||
merged[-1].content.extend(m.content)
|
||||
continue
|
||||
merged.append(m)
|
||||
|
||||
return merged
|
||||
|
||||
def _claude_usage_to_internal(self, usage: Any) -> Optional[UsageInfo]:
|
||||
if not isinstance(usage, dict):
|
||||
return None
|
||||
|
||||
mapping = USAGE_FIELD_MAPPINGS.get("CLAUDE", {})
|
||||
fields: Dict[str, int] = {}
|
||||
extra = self._extract_extra(usage, set(mapping.keys()))
|
||||
|
||||
for provider_key, internal_key in mapping.items():
|
||||
if provider_key in usage and usage.get(provider_key) is not None:
|
||||
try:
|
||||
fields[internal_key] = int(usage.get(provider_key) or 0)
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
|
||||
if "total_tokens" not in fields:
|
||||
fields["total_tokens"] = int(fields.get("input_tokens", 0) + fields.get("output_tokens", 0))
|
||||
|
||||
return UsageInfo(
|
||||
input_tokens=int(fields.get("input_tokens", 0)),
|
||||
output_tokens=int(fields.get("output_tokens", 0)),
|
||||
total_tokens=int(fields.get("total_tokens", 0)),
|
||||
cache_read_tokens=int(fields.get("cache_read_tokens", 0)),
|
||||
cache_write_tokens=int(fields.get("cache_write_tokens", 0)),
|
||||
extra={"claude": extra} if extra else {},
|
||||
)
|
||||
|
||||
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),
|
||||
}
|
||||
if usage.cache_read_tokens:
|
||||
result["cache_read_input_tokens"] = int(usage.cache_read_tokens)
|
||||
if usage.cache_write_tokens:
|
||||
result["cache_creation_input_tokens"] = int(usage.cache_write_tokens)
|
||||
return result
|
||||
|
||||
def _error_type_from_value(self, value: str) -> ErrorType:
|
||||
try:
|
||||
return ErrorType(value)
|
||||
except ValueError:
|
||||
return ErrorType.UNKNOWN
|
||||
|
||||
def _optional_int(self, value: Any) -> Optional[int]:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
def _optional_float(self, value: Any) -> Optional[float]:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
def _coerce_str_list(self, value: Any) -> Optional[List[str]]:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
return [value]
|
||||
if isinstance(value, list):
|
||||
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]:
|
||||
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:
|
||||
for k, v in source.items():
|
||||
target[k] = target.get(k, 0) + int(v)
|
||||
|
||||
|
||||
__all__ = ["ClaudeNormalizer"]
|
||||
18
src/core/api_format/conversion/normalizers/claude_cli.py
Normal file
18
src/core/api_format/conversion/normalizers/claude_cli.py
Normal file
@@ -0,0 +1,18 @@
|
||||
"""
|
||||
Claude CLI Normalizer
|
||||
|
||||
CLAUDE_CLI 的请求/响应 body 与 CLAUDE 一致(Anthropic Messages API),差异主要在认证头。
|
||||
因此这里复用 ClaudeNormalizer 的转换逻辑,仅更换 FORMAT_ID。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer
|
||||
|
||||
|
||||
class ClaudeCliNormalizer(ClaudeNormalizer):
|
||||
FORMAT_ID = "CLAUDE_CLI"
|
||||
|
||||
|
||||
__all__ = ["ClaudeCliNormalizer"]
|
||||
|
||||
944
src/core/api_format/conversion/normalizers/gemini.py
Normal file
944
src/core/api_format/conversion/normalizers/gemini.py
Normal file
@@ -0,0 +1,944 @@
|
||||
"""
|
||||
Gemini (GenerateContent / streamGenerateContent) Normalizer
|
||||
|
||||
负责:
|
||||
- Gemini request/response <-> Internal 表示转换
|
||||
- 可选:Gemini streaming chunk <-> InternalStreamEvent
|
||||
- 可选:Gemini error <-> InternalError
|
||||
|
||||
说明:
|
||||
- 请求体字段在本项目中同时兼容 snake_case(历史转换器产物)与 camelCase(官方/客户端输入)。
|
||||
- 响应/流式通常为 camelCase(candidates/finishReason/usageMetadata/modelVersion)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from src.core.api_format.conversion.field_mappings import (
|
||||
ERROR_TYPE_MAPPINGS,
|
||||
RETRYABLE_ERROR_TYPES,
|
||||
STOP_REASON_MAPPINGS,
|
||||
USAGE_FIELD_MAPPINGS,
|
||||
)
|
||||
from src.core.api_format.conversion.internal import (
|
||||
ContentBlock,
|
||||
ContentType,
|
||||
ErrorType,
|
||||
FormatCapabilities,
|
||||
ImageBlock,
|
||||
InstructionSegment,
|
||||
InternalError,
|
||||
InternalMessage,
|
||||
InternalRequest,
|
||||
InternalResponse,
|
||||
Role,
|
||||
StopReason,
|
||||
TextBlock,
|
||||
ToolChoice,
|
||||
ToolChoiceType,
|
||||
ToolDefinition,
|
||||
ToolResultBlock,
|
||||
ToolUseBlock,
|
||||
UnknownBlock,
|
||||
UsageInfo,
|
||||
)
|
||||
from src.core.api_format.conversion.normalizer import FormatNormalizer
|
||||
from src.core.api_format.conversion.stream_events import (
|
||||
ContentBlockStartEvent,
|
||||
ContentBlockStopEvent,
|
||||
ContentDeltaEvent,
|
||||
ErrorEvent,
|
||||
InternalStreamEvent,
|
||||
MessageStartEvent,
|
||||
MessageStopEvent,
|
||||
ToolCallDeltaEvent,
|
||||
)
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
class GeminiNormalizer(FormatNormalizer):
|
||||
FORMAT_ID = "GEMINI"
|
||||
capabilities = FormatCapabilities(
|
||||
supports_stream=True,
|
||||
supports_error_conversion=True,
|
||||
supports_tools=True,
|
||||
supports_images=True,
|
||||
)
|
||||
|
||||
_FINISH_REASON_TO_STOP: Dict[str, StopReason] = {
|
||||
"STOP": StopReason.END_TURN,
|
||||
"MAX_TOKENS": StopReason.MAX_TOKENS,
|
||||
"SAFETY": StopReason.CONTENT_FILTERED,
|
||||
"RECITATION": StopReason.CONTENT_FILTERED,
|
||||
"MALFORMED_FUNCTION_CALL": StopReason.TOOL_USE,
|
||||
"OTHER": StopReason.UNKNOWN,
|
||||
}
|
||||
|
||||
_ERROR_TYPE_TO_GEMINI_STATUS: Dict[ErrorType, str] = {
|
||||
ErrorType.INVALID_REQUEST: "INVALID_ARGUMENT",
|
||||
ErrorType.AUTHENTICATION: "UNAUTHENTICATED",
|
||||
ErrorType.PERMISSION_DENIED: "PERMISSION_DENIED",
|
||||
ErrorType.NOT_FOUND: "NOT_FOUND",
|
||||
ErrorType.RATE_LIMIT: "RESOURCE_EXHAUSTED",
|
||||
ErrorType.OVERLOADED: "UNAVAILABLE",
|
||||
ErrorType.SERVER_ERROR: "INTERNAL",
|
||||
ErrorType.CONTENT_FILTERED: "FAILED_PRECONDITION",
|
||||
ErrorType.CONTEXT_LENGTH_EXCEEDED: "INVALID_ARGUMENT",
|
||||
ErrorType.UNKNOWN: "INTERNAL",
|
||||
}
|
||||
|
||||
# =========================
|
||||
# Requests
|
||||
# =========================
|
||||
|
||||
def request_to_internal(self, request: Dict[str, Any]) -> InternalRequest:
|
||||
model = str(request.get("model") or "")
|
||||
dropped: Dict[str, int] = {}
|
||||
|
||||
instructions: List[InstructionSegment] = []
|
||||
system_text, sys_dropped = self._collapse_system_instruction(
|
||||
request.get("system_instruction")
|
||||
if "system_instruction" in request
|
||||
else request.get("systemInstruction")
|
||||
)
|
||||
self._merge_dropped(dropped, sys_dropped)
|
||||
if system_text:
|
||||
instructions.append(InstructionSegment(role=Role.SYSTEM, text=system_text))
|
||||
|
||||
messages: List[InternalMessage] = []
|
||||
contents = request.get("contents") or []
|
||||
if isinstance(contents, list):
|
||||
for content in contents:
|
||||
if not isinstance(content, dict):
|
||||
dropped["gemini_content_non_dict"] = dropped.get("gemini_content_non_dict", 0) + 1
|
||||
continue
|
||||
imsg, md = self._content_to_internal_message(content)
|
||||
self._merge_dropped(dropped, md)
|
||||
if imsg is not None:
|
||||
messages.append(imsg)
|
||||
else:
|
||||
dropped["gemini_contents_non_list"] = dropped.get("gemini_contents_non_list", 0) + 1
|
||||
|
||||
generation_config = self._get_generation_config(request)
|
||||
|
||||
max_tokens = self._optional_int(
|
||||
generation_config.get("max_output_tokens")
|
||||
if isinstance(generation_config, dict)
|
||||
else None
|
||||
)
|
||||
temperature = self._optional_float(
|
||||
generation_config.get("temperature")
|
||||
if isinstance(generation_config, dict)
|
||||
else None
|
||||
)
|
||||
top_p = self._optional_float(
|
||||
generation_config.get("top_p") if isinstance(generation_config, dict) else None
|
||||
)
|
||||
top_k = self._optional_int(
|
||||
generation_config.get("top_k") if isinstance(generation_config, dict) else None
|
||||
)
|
||||
stop_sequences = None
|
||||
if isinstance(generation_config, dict):
|
||||
stop_sequences = self._coerce_str_list(generation_config.get("stop_sequences"))
|
||||
|
||||
tools = self._gemini_tools_to_internal(request.get("tools"))
|
||||
tool_choice = self._gemini_tool_config_to_tool_choice(
|
||||
request.get("tool_config")
|
||||
if "tool_config" in request
|
||||
else request.get("toolConfig")
|
||||
)
|
||||
|
||||
internal = InternalRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
instructions=instructions,
|
||||
system=self._join_instructions(instructions),
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
top_k=top_k,
|
||||
stop_sequences=stop_sequences,
|
||||
stream=bool(request.get("stream") or False),
|
||||
tools=tools,
|
||||
tool_choice=tool_choice,
|
||||
extra={"gemini": self._extract_extra(request, {"contents"})},
|
||||
)
|
||||
|
||||
if dropped:
|
||||
internal.extra.setdefault("raw", {})["dropped_blocks"] = dropped
|
||||
|
||||
return internal
|
||||
|
||||
def request_from_internal(self, internal: InternalRequest) -> Dict[str, Any]:
|
||||
system_text = internal.system or self._join_instructions(internal.instructions)
|
||||
|
||||
# tools/tool_choice
|
||||
tools = None
|
||||
if internal.tools:
|
||||
tools = [
|
||||
{
|
||||
"function_declarations": [
|
||||
{
|
||||
"name": t.name,
|
||||
"description": t.description,
|
||||
"parameters": t.parameters or {},
|
||||
**(t.extra.get("gemini_function_declaration") or {}),
|
||||
}
|
||||
for t in internal.tools
|
||||
]
|
||||
}
|
||||
]
|
||||
|
||||
tool_config = None
|
||||
if internal.tool_choice:
|
||||
tool_config = self._tool_choice_to_gemini_tool_config(internal.tool_choice)
|
||||
|
||||
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:
|
||||
generation_config["temperature"] = internal.temperature
|
||||
if internal.top_p is not None:
|
||||
generation_config["top_p"] = internal.top_p
|
||||
if internal.top_k is not None:
|
||||
generation_config["top_k"] = internal.top_k
|
||||
if internal.stop_sequences:
|
||||
generation_config["stop_sequences"] = list(internal.stop_sequences)
|
||||
|
||||
contents: List[Dict[str, Any]] = []
|
||||
for msg in internal.messages:
|
||||
contents.append(self._internal_message_to_content(msg))
|
||||
|
||||
result: Dict[str, Any] = {
|
||||
"contents": contents,
|
||||
}
|
||||
|
||||
# Gemini Chat 模式 model 可能在 URL 路径中;这里仅在 internal.model 存在时回写
|
||||
if internal.model:
|
||||
result["model"] = internal.model
|
||||
|
||||
if system_text:
|
||||
result["system_instruction"] = {"parts": [{"text": system_text}]}
|
||||
|
||||
if generation_config:
|
||||
result["generation_config"] = generation_config
|
||||
|
||||
if tools:
|
||||
result["tools"] = tools
|
||||
|
||||
if tool_config:
|
||||
result["tool_config"] = tool_config
|
||||
|
||||
return result
|
||||
|
||||
# =========================
|
||||
# Responses
|
||||
# =========================
|
||||
|
||||
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 "")
|
||||
|
||||
candidates = response.get("candidates") or []
|
||||
candidate0 = candidates[0] if isinstance(candidates, list) and candidates else {}
|
||||
candidate0 = candidate0 if isinstance(candidate0, dict) else {}
|
||||
|
||||
content = candidate0.get("content") if isinstance(candidate0, dict) else None
|
||||
content = content if isinstance(content, dict) else {}
|
||||
|
||||
blocks, dropped = self._parts_to_blocks(content.get("parts"))
|
||||
|
||||
finish_reason = candidate0.get("finishReason")
|
||||
stop_reason = None
|
||||
if finish_reason is not None:
|
||||
stop_reason = self._FINISH_REASON_TO_STOP.get(str(finish_reason), StopReason.UNKNOWN)
|
||||
|
||||
usage_info = self._usage_metadata_to_internal(response.get("usageMetadata"))
|
||||
|
||||
extra: Dict[str, Any] = {}
|
||||
if finish_reason is not None:
|
||||
extra.setdefault("raw", {})["finishReason"] = finish_reason
|
||||
|
||||
internal = InternalResponse(
|
||||
id=rid,
|
||||
model=model,
|
||||
content=blocks,
|
||||
stop_reason=stop_reason,
|
||||
usage=usage_info,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
if dropped:
|
||||
internal.extra.setdefault("raw", {})["dropped_blocks"] = dropped
|
||||
|
||||
return internal
|
||||
|
||||
def response_from_internal(self, internal: InternalResponse) -> Dict[str, Any]:
|
||||
parts: List[Dict[str, Any]] = []
|
||||
for b in internal.content:
|
||||
if isinstance(b, TextBlock):
|
||||
if b.text:
|
||||
parts.append({"text": b.text})
|
||||
continue
|
||||
if isinstance(b, ToolUseBlock):
|
||||
parts.append(
|
||||
{
|
||||
"functionCall": {
|
||||
"name": b.tool_name,
|
||||
"args": b.tool_input or {},
|
||||
}
|
||||
}
|
||||
)
|
||||
continue
|
||||
if isinstance(b, ImageBlock):
|
||||
if b.data and b.media_type:
|
||||
parts.append(
|
||||
{
|
||||
"inlineData": {
|
||||
"mimeType": b.media_type,
|
||||
"data": b.data,
|
||||
}
|
||||
}
|
||||
)
|
||||
elif b.url:
|
||||
parts.append({"text": f"[Image: {b.url}]"})
|
||||
continue
|
||||
# Unknown/ToolResult 默认丢弃
|
||||
|
||||
finish_reason = None
|
||||
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] = {}
|
||||
if internal.usage:
|
||||
usage_metadata = {
|
||||
"promptTokenCount": int(internal.usage.input_tokens),
|
||||
"candidatesTokenCount": int(internal.usage.output_tokens),
|
||||
"totalTokenCount": int(internal.usage.total_tokens or (internal.usage.input_tokens + internal.usage.output_tokens)),
|
||||
}
|
||||
if internal.usage.cache_read_tokens:
|
||||
usage_metadata["cachedContentTokenCount"] = int(internal.usage.cache_read_tokens)
|
||||
|
||||
candidate: Dict[str, Any] = {
|
||||
"content": {"parts": parts, "role": "model"},
|
||||
"index": 0,
|
||||
}
|
||||
if finish_reason is not None:
|
||||
candidate["finishReason"] = finish_reason
|
||||
|
||||
out: Dict[str, Any] = {
|
||||
"candidates": [candidate],
|
||||
"modelVersion": internal.model or "gemini",
|
||||
}
|
||||
|
||||
if usage_metadata:
|
||||
out["usageMetadata"] = usage_metadata
|
||||
|
||||
# id 不是 Gemini 标准字段,但内部可能携带;保守保留
|
||||
if internal.id:
|
||||
out["id"] = internal.id
|
||||
|
||||
return out
|
||||
|
||||
# =========================
|
||||
# Streaming
|
||||
# =========================
|
||||
|
||||
def stream_chunk_to_internal(self, chunk: Dict[str, Any], state: StreamState) -> List[InternalStreamEvent]:
|
||||
ss = state.substate(self.FORMAT_ID)
|
||||
events: List[InternalStreamEvent] = []
|
||||
|
||||
if not ss.get("message_started"):
|
||||
# 保留初始化时设置的 model(客户端请求的模型),仅在空时用上游值
|
||||
model = state.model or str(chunk.get("modelVersion") or "")
|
||||
if not state.model:
|
||||
state.model = model
|
||||
state.message_id = state.message_id or "gemini"
|
||||
ss["message_started"] = True
|
||||
ss.setdefault("text_block_started", False)
|
||||
ss.setdefault("accumulated_text", "")
|
||||
ss.setdefault("next_block_index", 1) # 0 预留给文本
|
||||
events.append(MessageStartEvent(message_id=state.message_id, model=model))
|
||||
|
||||
candidates = chunk.get("candidates") or []
|
||||
if not isinstance(candidates, list) or not candidates:
|
||||
return events
|
||||
|
||||
candidate0 = candidates[0] if candidates else {}
|
||||
candidate0 = candidate0 if isinstance(candidate0, dict) else {}
|
||||
|
||||
content = candidate0.get("content") if isinstance(candidate0, dict) else None
|
||||
content = content if isinstance(content, dict) else {}
|
||||
parts = content.get("parts") or []
|
||||
|
||||
# parts -> events
|
||||
if isinstance(parts, list):
|
||||
for part in parts:
|
||||
if not isinstance(part, dict):
|
||||
continue
|
||||
|
||||
# text(兼容:delta 或累积)
|
||||
text = part.get("text")
|
||||
if isinstance(text, str) and text:
|
||||
prev = str(ss.get("accumulated_text") or "")
|
||||
if text.startswith(prev):
|
||||
delta = text[len(prev) :]
|
||||
ss["accumulated_text"] = text
|
||||
else:
|
||||
delta = text
|
||||
ss["accumulated_text"] = prev + delta
|
||||
|
||||
if delta:
|
||||
if not ss.get("text_block_started"):
|
||||
ss["text_block_started"] = True
|
||||
events.append(ContentBlockStartEvent(block_index=0, block_type=ContentType.TEXT))
|
||||
events.append(ContentDeltaEvent(block_index=0, text_delta=delta))
|
||||
continue
|
||||
|
||||
# functionCall(stream response 常见 camelCase)
|
||||
func_call = part.get("functionCall")
|
||||
if func_call is None:
|
||||
func_call = part.get("function_call")
|
||||
|
||||
if isinstance(func_call, dict):
|
||||
name = str(func_call.get("name") or "")
|
||||
args = func_call.get("args")
|
||||
if not isinstance(args, dict):
|
||||
args = {}
|
||||
|
||||
block_index = int(ss.get("next_block_index") or 1)
|
||||
ss["next_block_index"] = block_index + 1
|
||||
|
||||
events.append(
|
||||
ContentBlockStartEvent(
|
||||
block_index=block_index,
|
||||
block_type=ContentType.TOOL_USE,
|
||||
tool_id=None,
|
||||
tool_name=name or None,
|
||||
)
|
||||
)
|
||||
if args:
|
||||
events.append(
|
||||
ToolCallDeltaEvent(
|
||||
block_index=block_index,
|
||||
tool_id="",
|
||||
input_delta=json.dumps(args, ensure_ascii=False),
|
||||
)
|
||||
)
|
||||
events.append(ContentBlockStopEvent(block_index=block_index))
|
||||
continue
|
||||
|
||||
finish_reason = candidate0.get("finishReason")
|
||||
if finish_reason is not None:
|
||||
stop_reason = self._FINISH_REASON_TO_STOP.get(str(finish_reason), StopReason.UNKNOWN)
|
||||
usage_info = self._usage_metadata_to_internal(chunk.get("usageMetadata"))
|
||||
# 先补齐 content_block_stop(仅 text block),再发送 MessageStop
|
||||
if ss.get("text_block_started") and not ss.get("text_block_stopped"):
|
||||
ss["text_block_stopped"] = True
|
||||
events.append(ContentBlockStopEvent(block_index=0))
|
||||
events.append(MessageStopEvent(stop_reason=stop_reason, usage=usage_info))
|
||||
|
||||
if "error" in chunk:
|
||||
events.append(ErrorEvent(error=self.error_to_internal(chunk)))
|
||||
|
||||
return events
|
||||
|
||||
def stream_event_from_internal(
|
||||
self,
|
||||
event: InternalStreamEvent,
|
||||
state: StreamState,
|
||||
) -> List[Dict[str, Any]]:
|
||||
ss = state.substate(self.FORMAT_ID)
|
||||
out: List[Dict[str, Any]] = []
|
||||
|
||||
def base_chunk(parts: List[Dict[str, Any]]) -> Dict[str, Any]:
|
||||
return {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": parts, "role": "model"},
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
"modelVersion": state.model or "",
|
||||
}
|
||||
|
||||
if isinstance(event, MessageStartEvent):
|
||||
state.message_id = event.message_id or state.message_id
|
||||
state.model = event.model or state.model
|
||||
ss.setdefault("tool_blocks", {})
|
||||
return out
|
||||
|
||||
if isinstance(event, ContentDeltaEvent):
|
||||
if event.text_delta:
|
||||
out.append(base_chunk([{"text": event.text_delta}]))
|
||||
return out
|
||||
|
||||
if isinstance(event, ContentBlockStartEvent) and event.block_type == ContentType.TOOL_USE:
|
||||
tool_blocks = ss.get("tool_blocks")
|
||||
if not isinstance(tool_blocks, dict):
|
||||
tool_blocks = {}
|
||||
ss["tool_blocks"] = tool_blocks
|
||||
|
||||
tool_blocks[int(event.block_index)] = {
|
||||
"name": event.tool_name or "",
|
||||
"json": "",
|
||||
}
|
||||
return out
|
||||
|
||||
if isinstance(event, ToolCallDeltaEvent):
|
||||
tool_blocks = ss.get("tool_blocks")
|
||||
if isinstance(tool_blocks, dict):
|
||||
entry = tool_blocks.get(int(event.block_index))
|
||||
if isinstance(entry, dict):
|
||||
entry["json"] = str(entry.get("json") or "") + (event.input_delta or "")
|
||||
return out
|
||||
|
||||
if isinstance(event, ContentBlockStopEvent):
|
||||
tool_blocks = ss.get("tool_blocks")
|
||||
if not isinstance(tool_blocks, dict):
|
||||
return out
|
||||
|
||||
entry = tool_blocks.get(int(event.block_index))
|
||||
if not isinstance(entry, dict):
|
||||
return out
|
||||
|
||||
name = str(entry.get("name") or "")
|
||||
raw_json = str(entry.get("json") or "")
|
||||
args: Dict[str, Any] = {}
|
||||
if raw_json:
|
||||
try:
|
||||
parsed = json.loads(raw_json)
|
||||
if isinstance(parsed, dict):
|
||||
args = parsed
|
||||
except json.JSONDecodeError:
|
||||
args = {}
|
||||
|
||||
out.append(base_chunk([{"functionCall": {"name": name, "args": args}}]))
|
||||
return out
|
||||
|
||||
if isinstance(event, MessageStopEvent):
|
||||
finish_reason = None
|
||||
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([])
|
||||
if finish_reason is not None:
|
||||
chunk["candidates"][0]["finishReason"] = finish_reason
|
||||
|
||||
if event.usage:
|
||||
chunk["usageMetadata"] = {
|
||||
"promptTokenCount": int(event.usage.input_tokens),
|
||||
"candidatesTokenCount": int(event.usage.output_tokens),
|
||||
"totalTokenCount": int(event.usage.total_tokens or (event.usage.input_tokens + event.usage.output_tokens)),
|
||||
}
|
||||
if event.usage.cache_read_tokens:
|
||||
chunk["usageMetadata"]["cachedContentTokenCount"] = int(event.usage.cache_read_tokens)
|
||||
|
||||
out.append(chunk)
|
||||
return out
|
||||
|
||||
if isinstance(event, ErrorEvent):
|
||||
out.append(self.error_from_internal(event.error))
|
||||
return out
|
||||
|
||||
return out
|
||||
|
||||
# =========================
|
||||
# Error conversion
|
||||
# =========================
|
||||
|
||||
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:
|
||||
err = error_response.get("error") if isinstance(error_response, dict) else None
|
||||
err = err if isinstance(err, dict) else {}
|
||||
|
||||
raw_status = err.get("status")
|
||||
mapped = ERROR_TYPE_MAPPINGS.get("GEMINI", {}).get(str(raw_status), ErrorType.UNKNOWN.value)
|
||||
internal_type = self._error_type_from_value(mapped)
|
||||
retryable = internal_type.value in RETRYABLE_ERROR_TYPES
|
||||
|
||||
code_value = err.get("code")
|
||||
code_str = None
|
||||
if code_value is not None:
|
||||
code_str = str(code_value)
|
||||
|
||||
return InternalError(
|
||||
type=internal_type,
|
||||
message=str(err.get("message") or ""),
|
||||
code=code_str,
|
||||
param=None,
|
||||
retryable=retryable,
|
||||
extra={"gemini": {"error": err}, "raw": {"status": raw_status}},
|
||||
)
|
||||
|
||||
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] = {
|
||||
"code": 400 if internal.type == ErrorType.INVALID_REQUEST else 500,
|
||||
"message": internal.message,
|
||||
"status": status,
|
||||
}
|
||||
return {"error": payload}
|
||||
|
||||
# =========================
|
||||
# Helpers
|
||||
# =========================
|
||||
|
||||
def _content_to_internal_message(self, content: Dict[str, Any]) -> Tuple[Optional[InternalMessage], Dict[str, int]]:
|
||||
dropped: Dict[str, int] = {}
|
||||
|
||||
role_raw = str(content.get("role") or "user")
|
||||
if role_raw == "model":
|
||||
role = Role.ASSISTANT
|
||||
elif role_raw == "user":
|
||||
role = Role.USER
|
||||
else:
|
||||
role = Role.UNKNOWN
|
||||
|
||||
blocks, bd = self._parts_to_blocks(content.get("parts"))
|
||||
self._merge_dropped(dropped, bd)
|
||||
|
||||
return (
|
||||
InternalMessage(
|
||||
role=role,
|
||||
content=blocks,
|
||||
extra=self._extract_extra(content, {"role", "parts"}),
|
||||
),
|
||||
dropped,
|
||||
)
|
||||
|
||||
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] = []
|
||||
for part in parts:
|
||||
if not isinstance(part, dict):
|
||||
dropped["gemini_part_non_dict"] = dropped.get("gemini_part_non_dict", 0) + 1
|
||||
continue
|
||||
|
||||
if "text" in part:
|
||||
text = part.get("text")
|
||||
if isinstance(text, str) and text:
|
||||
blocks.append(TextBlock(text=text, extra=self._extract_extra(part, {"text"})))
|
||||
continue
|
||||
|
||||
inline = part.get("inline_data")
|
||||
if inline is None:
|
||||
inline = part.get("inlineData")
|
||||
if isinstance(inline, dict):
|
||||
mime_type = inline.get("mime_type") if "mime_type" in inline else inline.get("mimeType")
|
||||
data = inline.get("data")
|
||||
if isinstance(mime_type, str) and mime_type and isinstance(data, str) and data:
|
||||
blocks.append(ImageBlock(data=data, media_type=mime_type))
|
||||
else:
|
||||
dropped["gemini_inline_data_invalid"] = dropped.get("gemini_inline_data_invalid", 0) + 1
|
||||
blocks.append(UnknownBlock(raw_type="inline_data", payload=part))
|
||||
continue
|
||||
|
||||
func_call = part.get("function_call")
|
||||
if func_call is None:
|
||||
func_call = part.get("functionCall")
|
||||
if isinstance(func_call, dict):
|
||||
name = str(func_call.get("name") or "")
|
||||
args = func_call.get("args")
|
||||
if not isinstance(args, dict):
|
||||
args = {}
|
||||
blocks.append(
|
||||
ToolUseBlock(
|
||||
tool_id=f"toolu_{name}" if name else "toolu_0",
|
||||
tool_name=name,
|
||||
tool_input=args,
|
||||
extra={"gemini": part},
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
func_resp = part.get("function_response")
|
||||
if func_resp is None:
|
||||
func_resp = part.get("functionResponse")
|
||||
if isinstance(func_resp, dict):
|
||||
name = str(func_resp.get("name") or "")
|
||||
response = func_resp.get("response")
|
||||
output: Any = None
|
||||
content_text: Optional[str] = None
|
||||
|
||||
# 兼容历史:response 常见结构为 {"result": ...}
|
||||
if isinstance(response, dict) and "result" in response:
|
||||
output = response.get("result")
|
||||
if isinstance(output, str):
|
||||
content_text = output
|
||||
output = None
|
||||
else:
|
||||
output = response
|
||||
|
||||
blocks.append(
|
||||
ToolResultBlock(
|
||||
tool_use_id=name,
|
||||
output=output,
|
||||
content_text=content_text,
|
||||
is_error=False,
|
||||
extra={"gemini": part},
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# 其它:Unknown
|
||||
raw_type = next(iter(part.keys()), "unknown")
|
||||
dropped_key = f"gemini_part:{raw_type}"
|
||||
dropped[dropped_key] = dropped.get(dropped_key, 0) + 1
|
||||
blocks.append(UnknownBlock(raw_type=str(raw_type), payload=part))
|
||||
|
||||
return blocks, dropped
|
||||
|
||||
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]] = []
|
||||
for b in msg.content:
|
||||
if isinstance(b, UnknownBlock):
|
||||
continue
|
||||
|
||||
if isinstance(b, TextBlock):
|
||||
if b.text:
|
||||
parts.append({"text": b.text})
|
||||
continue
|
||||
|
||||
if isinstance(b, ImageBlock):
|
||||
if b.data and b.media_type:
|
||||
parts.append({"inline_data": {"mime_type": b.media_type, "data": b.data}})
|
||||
elif b.url:
|
||||
parts.append({"text": f"[Image: {b.url}]"})
|
||||
continue
|
||||
|
||||
if isinstance(b, ToolUseBlock) and role == "model":
|
||||
parts.append({"function_call": {"name": b.tool_name, "args": b.tool_input or {}}})
|
||||
continue
|
||||
|
||||
if isinstance(b, ToolResultBlock) and role == "user":
|
||||
# 兼容旧转换器:name 直接使用 tool_use_id,response 固定包一层 result
|
||||
value: Any
|
||||
if b.content_text is not None:
|
||||
value = b.content_text
|
||||
elif b.output is None:
|
||||
value = ""
|
||||
else:
|
||||
value = b.output
|
||||
|
||||
parts.append(
|
||||
{
|
||||
"function_response": {
|
||||
"name": b.tool_use_id,
|
||||
"response": {"result": value},
|
||||
}
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
return {"role": role, "parts": parts}
|
||||
|
||||
def _collapse_system_instruction(self, system_instruction: Any) -> Tuple[Optional[str], Dict[str, int]]:
|
||||
dropped: Dict[str, int] = {}
|
||||
if system_instruction is None:
|
||||
return None, dropped
|
||||
|
||||
# 支持 {"parts": [{"text": ...}, ...]}
|
||||
if isinstance(system_instruction, dict):
|
||||
parts = system_instruction.get("parts")
|
||||
if isinstance(parts, list):
|
||||
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")))
|
||||
joined = "".join(texts)
|
||||
return (joined or None), dropped
|
||||
|
||||
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]:
|
||||
# 兼容 snake_case 与 camelCase
|
||||
gc = request.get("generation_config") if "generation_config" in request else request.get("generationConfig")
|
||||
if not isinstance(gc, dict):
|
||||
return {}
|
||||
|
||||
# 统一内部使用 snake_case key
|
||||
def pick(*keys: str) -> Any:
|
||||
for k in keys:
|
||||
if k in gc:
|
||||
return gc.get(k)
|
||||
return None
|
||||
|
||||
normalized: Dict[str, Any] = {}
|
||||
normalized["max_output_tokens"] = pick("max_output_tokens", "maxOutputTokens")
|
||||
normalized["temperature"] = pick("temperature")
|
||||
normalized["top_p"] = pick("top_p", "topP")
|
||||
normalized["top_k"] = pick("top_k", "topK")
|
||||
normalized["stop_sequences"] = pick("stop_sequences", "stopSequences")
|
||||
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]]:
|
||||
if not tools or not isinstance(tools, list):
|
||||
return None
|
||||
|
||||
out: List[ToolDefinition] = []
|
||||
for tool in tools:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
|
||||
decls = tool.get("function_declarations")
|
||||
if decls is None:
|
||||
decls = tool.get("functionDeclarations")
|
||||
|
||||
if not isinstance(decls, list):
|
||||
continue
|
||||
|
||||
for decl in decls:
|
||||
if not isinstance(decl, dict):
|
||||
continue
|
||||
name = str(decl.get("name") or "")
|
||||
if not name:
|
||||
continue
|
||||
out.append(
|
||||
ToolDefinition(
|
||||
name=name,
|
||||
description=decl.get("description"),
|
||||
parameters=decl.get("parameters") if isinstance(decl.get("parameters"), dict) else None,
|
||||
extra={"gemini_function_declaration": self._extract_extra(decl, {"name", "description", "parameters"})},
|
||||
)
|
||||
)
|
||||
|
||||
return out or None
|
||||
|
||||
def _gemini_tool_config_to_tool_choice(self, tool_config: Any) -> Optional[ToolChoice]:
|
||||
if tool_config is None:
|
||||
return None
|
||||
if not isinstance(tool_config, dict):
|
||||
return ToolChoice(type=ToolChoiceType.AUTO, extra={"raw": tool_config})
|
||||
|
||||
cfg = tool_config.get("function_calling_config")
|
||||
if cfg is None:
|
||||
cfg = tool_config.get("functionCallingConfig")
|
||||
|
||||
if not isinstance(cfg, dict):
|
||||
return ToolChoice(type=ToolChoiceType.AUTO, extra={"gemini": tool_config})
|
||||
|
||||
mode = str(cfg.get("mode") or "AUTO").upper()
|
||||
allowed = cfg.get("allowed_function_names")
|
||||
if allowed is None:
|
||||
allowed = cfg.get("allowedFunctionNames")
|
||||
|
||||
if mode == "NONE":
|
||||
return ToolChoice(type=ToolChoiceType.NONE, extra={"gemini": tool_config})
|
||||
if mode in ("ANY", "REQUIRED"):
|
||||
return ToolChoice(type=ToolChoiceType.REQUIRED, extra={"gemini": tool_config})
|
||||
if isinstance(allowed, list) and len(allowed) == 1:
|
||||
return ToolChoice(type=ToolChoiceType.TOOL, tool_name=str(allowed[0] or ""), 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]:
|
||||
mode = "AUTO"
|
||||
cfg: Dict[str, Any] = {}
|
||||
|
||||
if tool_choice.type == ToolChoiceType.NONE:
|
||||
mode = "NONE"
|
||||
elif tool_choice.type == ToolChoiceType.REQUIRED:
|
||||
mode = "ANY"
|
||||
elif tool_choice.type == ToolChoiceType.TOOL:
|
||||
mode = "ANY"
|
||||
cfg["allowed_function_names"] = [tool_choice.tool_name or ""]
|
||||
|
||||
cfg["mode"] = mode
|
||||
return {"function_calling_config": cfg}
|
||||
|
||||
def _usage_metadata_to_internal(self, usage_metadata: Any) -> Optional[UsageInfo]:
|
||||
if not isinstance(usage_metadata, dict):
|
||||
return None
|
||||
|
||||
mapping = USAGE_FIELD_MAPPINGS.get("GEMINI", {})
|
||||
fields: Dict[str, int] = {}
|
||||
extra = self._extract_extra(usage_metadata, set(mapping.keys()))
|
||||
|
||||
# promptTokenCount/candidatesTokenCount/totalTokenCount/cachedContentTokenCount
|
||||
for provider_key, internal_key in mapping.items():
|
||||
if provider_key in usage_metadata and usage_metadata.get(provider_key) is not None:
|
||||
try:
|
||||
fields[internal_key] = int(usage_metadata.get(provider_key) or 0)
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
|
||||
# thoughtsTokenCount(如果存在,按 handler 的口径并入 output_tokens)
|
||||
thoughts = usage_metadata.get("thoughtsTokenCount")
|
||||
if thoughts is not None:
|
||||
try:
|
||||
fields["output_tokens"] = int(fields.get("output_tokens", 0) + int(thoughts or 0))
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
|
||||
if "total_tokens" not in fields:
|
||||
fields["total_tokens"] = int(fields.get("input_tokens", 0) + fields.get("output_tokens", 0))
|
||||
|
||||
return UsageInfo(
|
||||
input_tokens=int(fields.get("input_tokens", 0)),
|
||||
output_tokens=int(fields.get("output_tokens", 0)),
|
||||
total_tokens=int(fields.get("total_tokens", 0)),
|
||||
cache_read_tokens=int(fields.get("cache_read_tokens", 0)),
|
||||
cache_write_tokens=0,
|
||||
extra={"gemini": extra} if extra else {},
|
||||
)
|
||||
|
||||
def _join_instructions(self, instructions: List[InstructionSegment]) -> Optional[str]:
|
||||
parts = [seg.text for seg in instructions if seg.text]
|
||||
joined = "\n\n".join(parts)
|
||||
return joined or None
|
||||
|
||||
def _error_type_from_value(self, value: str) -> ErrorType:
|
||||
try:
|
||||
return ErrorType(value)
|
||||
except ValueError:
|
||||
return ErrorType.UNKNOWN
|
||||
|
||||
def _optional_int(self, value: Any) -> Optional[int]:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
def _optional_float(self, value: Any) -> Optional[float]:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
def _coerce_str_list(self, value: Any) -> Optional[List[str]]:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
return [value]
|
||||
if isinstance(value, list):
|
||||
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]:
|
||||
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:
|
||||
for k, v in source.items():
|
||||
target[k] = target.get(k, 0) + int(v)
|
||||
|
||||
|
||||
__all__ = ["GeminiNormalizer"]
|
||||
18
src/core/api_format/conversion/normalizers/gemini_cli.py
Normal file
18
src/core/api_format/conversion/normalizers/gemini_cli.py
Normal file
@@ -0,0 +1,18 @@
|
||||
"""
|
||||
Gemini CLI Normalizer
|
||||
|
||||
GEMINI_CLI 的请求/响应 body 与 GEMINI 一致(Google Gemini API),差异主要在鉴权/UA 等请求层。
|
||||
因此这里复用 GeminiNormalizer 的转换逻辑,仅更换 FORMAT_ID。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||
|
||||
|
||||
class GeminiCliNormalizer(GeminiNormalizer):
|
||||
FORMAT_ID = "GEMINI_CLI"
|
||||
|
||||
|
||||
__all__ = ["GeminiCliNormalizer"]
|
||||
|
||||
1065
src/core/api_format/conversion/normalizers/openai.py
Normal file
1065
src/core/api_format/conversion/normalizers/openai.py
Normal file
File diff suppressed because it is too large
Load Diff
880
src/core/api_format/conversion/normalizers/openai_cli.py
Normal file
880
src/core/api_format/conversion/normalizers/openai_cli.py
Normal file
@@ -0,0 +1,880 @@
|
||||
"""
|
||||
OpenAI CLI / Responses Normalizer (OPENAI_CLI)
|
||||
|
||||
目标:
|
||||
- 将 OpenAI Responses API(/v1/responses)映射到 InternalRequest / InternalResponse
|
||||
- 支持流式事件:response.output_text.delta / response.completed 等
|
||||
|
||||
说明:
|
||||
- 这里实现的是“最佳努力”的最小可用映射,重点覆盖文本与 usage。
|
||||
- 未识别的字段会进入 extra/raw,未知内容块保留在 internal,但默认输出阶段会丢弃。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
from src.core.api_format.conversion.field_mappings import (
|
||||
ERROR_TYPE_MAPPINGS,
|
||||
RETRYABLE_ERROR_TYPES,
|
||||
)
|
||||
from src.core.api_format.conversion.internal import (
|
||||
ContentBlock,
|
||||
ContentType,
|
||||
ErrorType,
|
||||
FormatCapabilities,
|
||||
InstructionSegment,
|
||||
InternalError,
|
||||
InternalMessage,
|
||||
InternalRequest,
|
||||
InternalResponse,
|
||||
Role,
|
||||
StopReason,
|
||||
TextBlock,
|
||||
ToolChoice,
|
||||
ToolChoiceType,
|
||||
ToolDefinition,
|
||||
ToolResultBlock,
|
||||
ToolUseBlock,
|
||||
UnknownBlock,
|
||||
UsageInfo,
|
||||
)
|
||||
from src.core.api_format.conversion.normalizer import FormatNormalizer
|
||||
from src.core.api_format.conversion.stream_events import (
|
||||
ContentBlockStartEvent,
|
||||
ContentBlockStopEvent,
|
||||
ContentDeltaEvent,
|
||||
ErrorEvent,
|
||||
InternalStreamEvent,
|
||||
MessageStartEvent,
|
||||
MessageStopEvent,
|
||||
ToolCallDeltaEvent,
|
||||
UnknownStreamEvent,
|
||||
)
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
class OpenAICliNormalizer(FormatNormalizer):
|
||||
FORMAT_ID = "OPENAI_CLI"
|
||||
capabilities = FormatCapabilities(
|
||||
supports_stream=True,
|
||||
supports_error_conversion=True,
|
||||
supports_tools=True,
|
||||
supports_images=True,
|
||||
)
|
||||
|
||||
_ERROR_TYPE_TO_OPENAI: Dict[ErrorType, str] = {
|
||||
ErrorType.INVALID_REQUEST: "invalid_request_error",
|
||||
ErrorType.AUTHENTICATION: "invalid_api_key",
|
||||
ErrorType.PERMISSION_DENIED: "invalid_request_error",
|
||||
ErrorType.NOT_FOUND: "not_found",
|
||||
ErrorType.RATE_LIMIT: "rate_limit_exceeded",
|
||||
ErrorType.OVERLOADED: "server_error",
|
||||
ErrorType.SERVER_ERROR: "server_error",
|
||||
ErrorType.CONTENT_FILTERED: "content_policy_violation",
|
||||
ErrorType.CONTEXT_LENGTH_EXCEEDED: "context_length_exceeded",
|
||||
ErrorType.UNKNOWN: "server_error",
|
||||
}
|
||||
|
||||
# =========================
|
||||
# Requests
|
||||
# =========================
|
||||
|
||||
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
|
||||
if isinstance(instructions_text, str) and instructions_text.strip():
|
||||
system_text = instructions_text
|
||||
instructions.append(InstructionSegment(role=Role.SYSTEM, text=instructions_text))
|
||||
|
||||
messages = self._input_to_internal_messages(request.get("input"))
|
||||
|
||||
tools = self._tools_to_internal(request.get("tools"))
|
||||
tool_choice = self._tool_choice_to_internal(request.get("tool_choice"))
|
||||
|
||||
max_tokens = self._optional_int(
|
||||
request.get("max_output_tokens", request.get("max_tokens"))
|
||||
)
|
||||
|
||||
internal = InternalRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
instructions=instructions,
|
||||
system=system_text,
|
||||
max_tokens=max_tokens,
|
||||
temperature=self._optional_float(request.get("temperature")),
|
||||
top_p=self._optional_float(request.get("top_p")),
|
||||
stop_sequences=self._coerce_str_list(request.get("stop")),
|
||||
stream=bool(request.get("stream") or False),
|
||||
tools=tools,
|
||||
tool_choice=tool_choice,
|
||||
extra={"openai_cli": self._extract_extra(request, {"input"})},
|
||||
)
|
||||
|
||||
return internal
|
||||
|
||||
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),
|
||||
}
|
||||
|
||||
instructions_text = self._join_instructions(internal)
|
||||
if instructions_text:
|
||||
result["instructions"] = instructions_text
|
||||
|
||||
if internal.max_tokens is not None:
|
||||
# Responses API 使用 max_output_tokens;兼容层仍可能接受 max_tokens
|
||||
result["max_output_tokens"] = internal.max_tokens
|
||||
if internal.temperature is not None:
|
||||
result["temperature"] = internal.temperature
|
||||
if internal.top_p is not None:
|
||||
result["top_p"] = internal.top_p
|
||||
if internal.stop_sequences:
|
||||
result["stop"] = list(internal.stop_sequences)
|
||||
if internal.stream:
|
||||
result["stream"] = True
|
||||
|
||||
if internal.tools:
|
||||
result["tools"] = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": t.name,
|
||||
"description": t.description,
|
||||
"parameters": t.parameters or {},
|
||||
**(t.extra.get("openai_function") or {}),
|
||||
},
|
||||
**(t.extra.get("openai_tool") or {}),
|
||||
}
|
||||
for t in internal.tools
|
||||
]
|
||||
|
||||
if internal.tool_choice:
|
||||
result["tool_choice"] = self._tool_choice_to_openai(internal.tool_choice)
|
||||
|
||||
return result
|
||||
|
||||
# =========================
|
||||
# Responses
|
||||
# =========================
|
||||
|
||||
def response_to_internal(self, response: Dict[str, Any]) -> InternalResponse:
|
||||
payload = self._unwrap_response_object(response)
|
||||
|
||||
rid = str(payload.get("id") or "")
|
||||
model = str(payload.get("model") or "")
|
||||
|
||||
blocks, extra = self._extract_output_text_blocks(payload)
|
||||
usage = self._usage_to_internal(payload.get("usage"))
|
||||
|
||||
stop_reason = StopReason.UNKNOWN
|
||||
status = payload.get("status")
|
||||
if isinstance(status, str) and status == "completed":
|
||||
stop_reason = StopReason.END_TURN
|
||||
|
||||
return InternalResponse(
|
||||
id=rid,
|
||||
model=model,
|
||||
content=blocks,
|
||||
stop_reason=stop_reason,
|
||||
usage=usage,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def response_from_internal(self, internal: InternalResponse) -> Dict[str, Any]:
|
||||
text = self._collapse_internal_text(internal.content)
|
||||
|
||||
output_message = {
|
||||
"type": "message",
|
||||
"id": f"msg_{internal.id or 'stream'}",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": text}],
|
||||
}
|
||||
|
||||
usage = internal.usage or UsageInfo()
|
||||
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),
|
||||
}
|
||||
|
||||
return {
|
||||
"id": internal.id or "resp",
|
||||
"object": "response",
|
||||
"created": int(time.time()),
|
||||
"model": internal.model or "",
|
||||
"status": "completed",
|
||||
"output": [output_message],
|
||||
"usage": usage_obj,
|
||||
}
|
||||
|
||||
# =========================
|
||||
# Stream conversion
|
||||
# =========================
|
||||
|
||||
def stream_chunk_to_internal(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
state: StreamState,
|
||||
) -> List[InternalStreamEvent]:
|
||||
ss = state.substate(self.FORMAT_ID)
|
||||
events: List[InternalStreamEvent] = []
|
||||
|
||||
# 统一错误结构(最佳努力)
|
||||
if isinstance(chunk, dict) and "error" in chunk:
|
||||
try:
|
||||
events.append(ErrorEvent(error=self.error_to_internal(chunk)))
|
||||
except Exception:
|
||||
pass
|
||||
return events
|
||||
|
||||
etype = str(chunk.get("type") or "")
|
||||
|
||||
# 尽量在首次事件补齐 message_start
|
||||
if not ss.get("message_started"):
|
||||
resp_obj = chunk.get("response")
|
||||
resp_obj = resp_obj if isinstance(resp_obj, dict) else {}
|
||||
msg_id = str(resp_obj.get("id") or chunk.get("id") or state.message_id or "")
|
||||
model = str(resp_obj.get("model") or chunk.get("model") or state.model or "")
|
||||
if msg_id or model or etype:
|
||||
state.message_id = msg_id
|
||||
state.model = model
|
||||
ss["message_started"] = True
|
||||
ss.setdefault("text_block_started", False)
|
||||
ss.setdefault("text_block_stopped", False)
|
||||
events.append(MessageStartEvent(message_id=msg_id, model=model))
|
||||
|
||||
# response.created:响应创建事件,message_start 已在上面处理
|
||||
if etype == "response.created":
|
||||
return events
|
||||
|
||||
# 文本增量:response.output_text.delta
|
||||
if etype in ("response.output_text.delta", "response.outtext.delta"):
|
||||
delta = chunk.get("delta")
|
||||
delta_text = ""
|
||||
if isinstance(delta, str):
|
||||
delta_text = delta
|
||||
elif isinstance(delta, dict) and isinstance(delta.get("text"), str):
|
||||
delta_text = str(delta.get("text") or "")
|
||||
|
||||
if delta_text:
|
||||
if not ss.get("text_block_started"):
|
||||
ss["text_block_started"] = True
|
||||
events.append(ContentBlockStartEvent(block_index=0, block_type=ContentType.TEXT))
|
||||
events.append(ContentDeltaEvent(block_index=0, text_delta=delta_text))
|
||||
return events
|
||||
|
||||
# 文本完成:response.output_text.done(可选)
|
||||
if etype == "response.output_text.done":
|
||||
if ss.get("text_block_started") and not ss.get("text_block_stopped"):
|
||||
ss["text_block_stopped"] = True
|
||||
events.append(ContentBlockStopEvent(block_index=0))
|
||||
return events
|
||||
|
||||
# 完成:response.completed(包含 usage)
|
||||
if etype == "response.completed":
|
||||
resp_obj = chunk.get("response")
|
||||
resp_obj = resp_obj if isinstance(resp_obj, dict) else {}
|
||||
usage = self._usage_to_internal(resp_obj.get("usage") or chunk.get("usage"))
|
||||
|
||||
if ss.get("text_block_started") and not ss.get("text_block_stopped"):
|
||||
ss["text_block_stopped"] = True
|
||||
events.append(ContentBlockStopEvent(block_index=0))
|
||||
|
||||
events.append(MessageStopEvent(stop_reason=StopReason.END_TURN, usage=usage))
|
||||
return events
|
||||
|
||||
# 失败:response.failed(最佳努力)
|
||||
if etype == "response.failed":
|
||||
try:
|
||||
events.append(ErrorEvent(error=self.error_to_internal(chunk)))
|
||||
except Exception:
|
||||
pass
|
||||
return events
|
||||
|
||||
# response.in_progress:状态更新,不产生内容事件
|
||||
if etype == "response.in_progress":
|
||||
# 更新 state 中的元数据(如果有)
|
||||
resp_obj = chunk.get("response")
|
||||
if isinstance(resp_obj, dict):
|
||||
if resp_obj.get("id"):
|
||||
state.message_id = str(resp_obj.get("id"))
|
||||
if resp_obj.get("model"):
|
||||
state.model = str(resp_obj.get("model"))
|
||||
return []
|
||||
|
||||
# response.output_item.added:新输出项添加(如 message、function_call 等)
|
||||
if etype == "response.output_item.added":
|
||||
item = chunk.get("item")
|
||||
if isinstance(item, dict):
|
||||
item_type = item.get("type")
|
||||
# function_call 输出项
|
||||
if item_type == "function_call":
|
||||
if not ss.get("tool_block_started"):
|
||||
ss["tool_block_started"] = True
|
||||
ss["current_tool_id"] = item.get("call_id") or item.get("id") or ""
|
||||
ss["current_tool_name"] = item.get("name") or ""
|
||||
events.append(ContentBlockStartEvent(
|
||||
block_index=ss.get("block_index", 0),
|
||||
block_type=ContentType.TOOL_USE,
|
||||
extra={"tool_id": ss["current_tool_id"], "tool_name": ss["current_tool_name"]},
|
||||
))
|
||||
ss["block_index"] = ss.get("block_index", 0) + 1
|
||||
# message 输出项
|
||||
elif item_type == "message":
|
||||
# 通常在 response.created 时已处理,这里可以忽略或更新状态
|
||||
pass
|
||||
return events
|
||||
|
||||
# response.output_item.done:输出项完成
|
||||
if etype == "response.output_item.done":
|
||||
item = chunk.get("item")
|
||||
if isinstance(item, dict):
|
||||
item_type = item.get("type")
|
||||
if item_type == "function_call" and ss.get("tool_block_started"):
|
||||
ss["tool_block_started"] = False
|
||||
events.append(ContentBlockStopEvent(block_index=ss.get("block_index", 1) - 1))
|
||||
return events
|
||||
|
||||
# response.function_call_arguments.delta:工具调用参数增量
|
||||
if etype == "response.function_call_arguments.delta":
|
||||
delta = chunk.get("delta") or ""
|
||||
if delta:
|
||||
events.append(ToolCallDeltaEvent(
|
||||
block_index=ss.get("block_index", 1) - 1,
|
||||
tool_id=ss.get("current_tool_id", ""),
|
||||
input_delta=delta,
|
||||
))
|
||||
return events
|
||||
|
||||
# response.function_call_arguments.done:工具调用参数完成
|
||||
if etype == "response.function_call_arguments.done":
|
||||
# 参数已完整,不产生额外事件
|
||||
return []
|
||||
|
||||
# response.content_part.added / response.content_part.done:内容部分事件
|
||||
if etype in ("response.content_part.added", "response.content_part.done"):
|
||||
# 通常伴随 output_text 事件,这里可以忽略
|
||||
return []
|
||||
|
||||
# response.reasoning_summary_text.delta:推理摘要增量
|
||||
if etype == "response.reasoning_summary_text.delta":
|
||||
# 保留为 UnknownStreamEvent,让下游决定是否使用
|
||||
return [UnknownStreamEvent(raw_type=etype, payload=chunk)]
|
||||
|
||||
# response.reasoning_summary_text.done:推理摘要完成
|
||||
if etype == "response.reasoning_summary_text.done":
|
||||
return [UnknownStreamEvent(raw_type=etype, payload=chunk)]
|
||||
|
||||
if etype:
|
||||
return [UnknownStreamEvent(raw_type=etype, payload=chunk)]
|
||||
return [UnknownStreamEvent(raw_type="unknown", payload=chunk)]
|
||||
|
||||
def stream_event_from_internal(
|
||||
self,
|
||||
event: InternalStreamEvent,
|
||||
state: StreamState,
|
||||
) -> List[Dict[str, Any]]:
|
||||
ss = state.substate(self.FORMAT_ID)
|
||||
out: List[Dict[str, Any]] = []
|
||||
|
||||
def event_block(payload: Dict[str, Any]) -> Dict[str, Any]:
|
||||
# OpenAI Responses SSE 的 payload 通常自带 type 字段;这里强制保证
|
||||
return payload
|
||||
|
||||
if isinstance(event, MessageStartEvent):
|
||||
state.message_id = event.message_id or state.message_id or "resp_stream"
|
||||
state.model = event.model or state.model or ""
|
||||
ss.setdefault("collected_text", "")
|
||||
out.append(
|
||||
event_block(
|
||||
{
|
||||
"type": "response.created",
|
||||
"response": {
|
||||
"id": state.message_id,
|
||||
"object": "response",
|
||||
"created": int(time.time()),
|
||||
"model": state.model,
|
||||
"status": "in_progress",
|
||||
"output": [],
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
return out
|
||||
|
||||
if isinstance(event, ContentDeltaEvent):
|
||||
if event.text_delta:
|
||||
ss["collected_text"] = str(ss.get("collected_text") or "") + event.text_delta
|
||||
out.append(
|
||||
event_block(
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"delta": event.text_delta,
|
||||
}
|
||||
)
|
||||
)
|
||||
return out
|
||||
|
||||
if isinstance(event, MessageStopEvent):
|
||||
final_text = str(ss.get("collected_text") or "")
|
||||
response_obj = self.response_from_internal(
|
||||
InternalResponse(
|
||||
id=state.message_id or "resp",
|
||||
model=state.model or "",
|
||||
content=[TextBlock(text=final_text)] if final_text else [],
|
||||
stop_reason=event.stop_reason or StopReason.END_TURN,
|
||||
usage=event.usage or UsageInfo(),
|
||||
)
|
||||
)
|
||||
out.append(event_block({"type": "response.completed", "response": response_obj}))
|
||||
return out
|
||||
|
||||
if isinstance(event, ErrorEvent):
|
||||
err_payload = self.error_from_internal(event.error)
|
||||
err_payload["type"] = "response.failed"
|
||||
out.append(event_block(err_payload))
|
||||
return out
|
||||
|
||||
# 其他事件:Responses SSE 无直接对应,跳过
|
||||
return out
|
||||
|
||||
# =========================
|
||||
# Error conversion
|
||||
# =========================
|
||||
|
||||
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:
|
||||
err = error_response.get("error") if isinstance(error_response, dict) else None
|
||||
err = err if isinstance(err, dict) else {}
|
||||
|
||||
raw_type = err.get("type")
|
||||
mapped = ERROR_TYPE_MAPPINGS.get("OPENAI", {}).get(str(raw_type), ErrorType.UNKNOWN.value)
|
||||
internal_type = self._error_type_from_value(mapped)
|
||||
retryable = internal_type.value in RETRYABLE_ERROR_TYPES
|
||||
|
||||
return InternalError(
|
||||
type=internal_type,
|
||||
message=str(err.get("message") or ""),
|
||||
code=err.get("code") if err.get("code") is None else str(err.get("code")),
|
||||
param=err.get("param") if err.get("param") is None else str(err.get("param")),
|
||||
retryable=retryable,
|
||||
extra={"openai_cli": {"error": err}, "raw": {"type": raw_type}},
|
||||
)
|
||||
|
||||
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}
|
||||
if internal.code is not None:
|
||||
payload["code"] = internal.code
|
||||
if internal.param is not None:
|
||||
payload["param"] = internal.param
|
||||
return {"error": payload}
|
||||
|
||||
# =========================
|
||||
# Helpers
|
||||
# =========================
|
||||
|
||||
def _unwrap_response_object(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||
if not isinstance(response, dict):
|
||||
return {}
|
||||
resp_inner = response.get("response")
|
||||
if isinstance(resp_inner, dict) and isinstance(response.get("type"), str):
|
||||
# 例如:{"type": "response.completed", "response": {...}}
|
||||
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] = []
|
||||
|
||||
output = payload.get("output")
|
||||
if isinstance(output, list):
|
||||
for item in output:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if item.get("type") == "message":
|
||||
content = item.get("content")
|
||||
if isinstance(content, list):
|
||||
for part in content:
|
||||
if not isinstance(part, dict):
|
||||
continue
|
||||
ptype = str(part.get("type") or "")
|
||||
if ptype in ("output_text", "text") and isinstance(part.get("text"), str):
|
||||
text_parts.append(part.get("text") or "")
|
||||
continue
|
||||
|
||||
if item.get("type") in ("output_text", "text") and isinstance(item.get("text"), str):
|
||||
text_parts.append(item.get("text") or "")
|
||||
|
||||
# 兼容:部分实现可能直接给 output_text
|
||||
if not text_parts and isinstance(payload.get("output_text"), str):
|
||||
text_parts.append(payload.get("output_text") or "")
|
||||
|
||||
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 {}
|
||||
return blocks, extra
|
||||
|
||||
def _usage_to_internal(self, usage: Any) -> UsageInfo:
|
||||
if not isinstance(usage, dict):
|
||||
return UsageInfo()
|
||||
input_tokens = int(usage.get("input_tokens") or 0)
|
||||
output_tokens = int(usage.get("output_tokens") or 0)
|
||||
total_tokens = int(usage.get("total_tokens") or (input_tokens + output_tokens))
|
||||
return UsageInfo(
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=total_tokens,
|
||||
extra={"openai_cli": {"usage": usage}},
|
||||
)
|
||||
|
||||
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]:
|
||||
if input_data is None:
|
||||
return []
|
||||
|
||||
# input: "text"
|
||||
if isinstance(input_data, str):
|
||||
return [InternalMessage(role=Role.USER, content=[TextBlock(text=input_data)])]
|
||||
|
||||
# input: {"messages": [...]}
|
||||
if isinstance(input_data, dict) and isinstance(input_data.get("messages"), list):
|
||||
input_data = input_data.get("messages")
|
||||
|
||||
if not isinstance(input_data, list):
|
||||
return [InternalMessage(role=Role.USER, content=[UnknownBlock(raw_type="input", payload={"input": input_data})])]
|
||||
|
||||
messages: List[InternalMessage] = []
|
||||
for item in input_data:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
|
||||
item_type = str(item.get("type") or "")
|
||||
|
||||
# 标准 message(有 role 字段)
|
||||
if item_type == "message" or item.get("role"):
|
||||
role = self._role_from_value(item.get("role"))
|
||||
blocks = self._responses_content_to_blocks(item.get("content"))
|
||||
messages.append(InternalMessage(role=role, content=blocks, extra=self._extract_extra(item, {"type", "role", "content"})))
|
||||
continue
|
||||
|
||||
# function_call -> assistant 消息 + ToolUseBlock
|
||||
if item_type == "function_call":
|
||||
tool_id = str(item.get("call_id") or item.get("id") or "")
|
||||
tool_name = str(item.get("name") or "")
|
||||
args_raw = item.get("arguments") or "{}"
|
||||
try:
|
||||
tool_input = json.loads(args_raw) if isinstance(args_raw, str) else (args_raw if isinstance(args_raw, dict) else {})
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
tool_input = {"_raw": args_raw}
|
||||
tool_block = ToolUseBlock(
|
||||
tool_id=tool_id,
|
||||
tool_name=tool_name,
|
||||
tool_input=tool_input,
|
||||
extra={"openai_cli": self._extract_extra(item, {"type", "call_id", "id", "name", "arguments"})},
|
||||
)
|
||||
messages.append(InternalMessage(role=Role.ASSISTANT, content=[tool_block]))
|
||||
continue
|
||||
|
||||
# function_call_output -> tool 消息 + ToolResultBlock
|
||||
if item_type == "function_call_output":
|
||||
tool_use_id = str(item.get("call_id") or item.get("id") or "")
|
||||
output = item.get("output")
|
||||
# output 可能是字符串或结构化数据
|
||||
content_text = output if isinstance(output, str) else None
|
||||
result_block = ToolResultBlock(
|
||||
tool_use_id=tool_use_id,
|
||||
output=output,
|
||||
content_text=content_text,
|
||||
extra={"openai_cli": self._extract_extra(item, {"type", "call_id", "id", "output"})},
|
||||
)
|
||||
messages.append(InternalMessage(role=Role.TOOL, content=[result_block]))
|
||||
continue
|
||||
|
||||
# reasoning -> assistant 消息,提取 summary 作为文本
|
||||
if item_type == "reasoning":
|
||||
summary_parts: List[str] = []
|
||||
summary = item.get("summary")
|
||||
if isinstance(summary, list):
|
||||
for s in summary:
|
||||
if isinstance(s, dict) and s.get("type") == "summary_text":
|
||||
text = s.get("text")
|
||||
if isinstance(text, str) and text:
|
||||
summary_parts.append(text)
|
||||
elif isinstance(s, str) and s:
|
||||
summary_parts.append(s)
|
||||
elif isinstance(summary, str) and summary:
|
||||
summary_parts.append(summary)
|
||||
|
||||
# 如果有 summary 文本,创建一个 UnknownBlock 保留原始结构
|
||||
reasoning_blocks: List[ContentBlock] = []
|
||||
if summary_parts:
|
||||
# 保留 reasoning 的 summary 作为 UnknownBlock,便于输出时决策
|
||||
reasoning_blocks.append(UnknownBlock(
|
||||
raw_type="reasoning",
|
||||
payload={"summary_text": "\n".join(summary_parts), "original": item},
|
||||
))
|
||||
else:
|
||||
reasoning_blocks.append(UnknownBlock(raw_type="reasoning", payload=item))
|
||||
messages.append(InternalMessage(
|
||||
role=Role.ASSISTANT,
|
||||
content=reasoning_blocks,
|
||||
extra={"openai_cli": {"type": "reasoning"}},
|
||||
))
|
||||
continue
|
||||
|
||||
# 其他未知类型 -> 保留为 UnknownBlock
|
||||
messages.append(InternalMessage(
|
||||
role=Role.UNKNOWN,
|
||||
content=[UnknownBlock(raw_type=item_type or "unknown", payload=item)],
|
||||
))
|
||||
|
||||
return messages
|
||||
|
||||
def _responses_content_to_blocks(self, content: Any) -> List[ContentBlock]:
|
||||
if content is None:
|
||||
return []
|
||||
if isinstance(content, str):
|
||||
return [TextBlock(text=content)]
|
||||
if isinstance(content, dict) and isinstance(content.get("text"), str):
|
||||
return [TextBlock(text=str(content.get("text") or ""))]
|
||||
|
||||
if not isinstance(content, list):
|
||||
return [UnknownBlock(raw_type="content", payload={"content": content})]
|
||||
|
||||
blocks: List[ContentBlock] = []
|
||||
for part in content:
|
||||
if isinstance(part, str):
|
||||
if part:
|
||||
blocks.append(TextBlock(text=part))
|
||||
continue
|
||||
if not isinstance(part, dict):
|
||||
continue
|
||||
ptype = str(part.get("type") or "")
|
||||
if ptype in ("input_text", "output_text", "text") and isinstance(part.get("text"), str):
|
||||
text = part.get("text") or ""
|
||||
if text:
|
||||
blocks.append(TextBlock(text=text))
|
||||
continue
|
||||
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]] = []
|
||||
for msg in messages:
|
||||
# ToolUseBlock -> function_call
|
||||
for block in msg.content:
|
||||
if isinstance(block, ToolUseBlock):
|
||||
out.append({
|
||||
"type": "function_call",
|
||||
"call_id": block.tool_id,
|
||||
"name": block.tool_name,
|
||||
"arguments": json.dumps(block.tool_input, ensure_ascii=False) if block.tool_input else "{}",
|
||||
})
|
||||
continue
|
||||
|
||||
if isinstance(block, ToolResultBlock):
|
||||
out.append({
|
||||
"type": "function_call_output",
|
||||
"call_id": block.tool_use_id,
|
||||
"output": block.content_text if block.content_text is not None else block.output,
|
||||
})
|
||||
continue
|
||||
|
||||
# reasoning(UnknownBlock with raw_type="reasoning")
|
||||
if isinstance(block, UnknownBlock) and block.raw_type == "reasoning":
|
||||
payload = block.payload or {}
|
||||
original = payload.get("original")
|
||||
if isinstance(original, dict):
|
||||
# 尽量还原原始结构
|
||||
out.append(original)
|
||||
else:
|
||||
summary_text = payload.get("summary_text", "")
|
||||
out.append({
|
||||
"type": "reasoning",
|
||||
"summary": [{"type": "summary_text", "text": summary_text}] if summary_text else [],
|
||||
})
|
||||
continue
|
||||
|
||||
# 普通 message(TextBlock)
|
||||
role = self._role_to_openai(msg.role)
|
||||
content_items: List[Dict[str, Any]] = []
|
||||
has_text = False
|
||||
|
||||
for block in msg.content:
|
||||
if isinstance(block, (ToolUseBlock, ToolResultBlock)):
|
||||
continue # 已在上面处理
|
||||
if isinstance(block, UnknownBlock) and block.raw_type == "reasoning":
|
||||
continue # 已在上面处理
|
||||
if isinstance(block, UnknownBlock):
|
||||
continue # 跳过其他未知块
|
||||
if isinstance(block, TextBlock) and block.text:
|
||||
content_items.append({"type": "input_text", "text": block.text})
|
||||
has_text = True
|
||||
|
||||
if has_text:
|
||||
out.append({"type": "message", "role": role, "content": content_items})
|
||||
|
||||
return out
|
||||
|
||||
def _tools_to_internal(self, tools: Any) -> Optional[List[ToolDefinition]]:
|
||||
if not isinstance(tools, list):
|
||||
return None
|
||||
out: List[ToolDefinition] = []
|
||||
for tool in tools:
|
||||
if not isinstance(tool, dict):
|
||||
continue
|
||||
if tool.get("type") == "function" and isinstance(tool.get("function"), dict):
|
||||
fn = tool["function"]
|
||||
name = str(fn.get("name") or "")
|
||||
if not name:
|
||||
continue
|
||||
out.append(
|
||||
ToolDefinition(
|
||||
name=name,
|
||||
description=fn.get("description"),
|
||||
parameters=fn.get("parameters") if isinstance(fn.get("parameters"), dict) else None,
|
||||
extra={"openai_tool": self._extract_extra(tool, {"type", "function"}), "openai_function": self._extract_extra(fn, {"name", "description", "parameters"})},
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# 兼容:部分实现可能直接给 {name, description, parameters}
|
||||
name = str(tool.get("name") or "")
|
||||
if name:
|
||||
out.append(
|
||||
ToolDefinition(
|
||||
name=name,
|
||||
description=tool.get("description"),
|
||||
parameters=tool.get("parameters") if isinstance(tool.get("parameters"), dict) else None,
|
||||
extra={"openai_cli": self._extract_extra(tool, {"name", "description", "parameters"})},
|
||||
)
|
||||
)
|
||||
return out or None
|
||||
|
||||
def _tool_choice_to_internal(self, tool_choice: Any) -> Optional[ToolChoice]:
|
||||
if tool_choice is None:
|
||||
return None
|
||||
if isinstance(tool_choice, str):
|
||||
if tool_choice == "none":
|
||||
return ToolChoice(type=ToolChoiceType.NONE, extra={"openai_cli": {"tool_choice": tool_choice}})
|
||||
if tool_choice == "auto":
|
||||
return ToolChoice(type=ToolChoiceType.AUTO, extra={"openai_cli": {"tool_choice": tool_choice}})
|
||||
return ToolChoice(type=ToolChoiceType.AUTO, extra={"raw": tool_choice})
|
||||
|
||||
if isinstance(tool_choice, dict):
|
||||
# OpenAI 兼容结构:{"type":"function","function":{"name":"..."}}
|
||||
if tool_choice.get("type") == "function" and isinstance(tool_choice.get("function"), dict):
|
||||
name = str(tool_choice["function"].get("name") or "")
|
||||
return ToolChoice(type=ToolChoiceType.TOOL, tool_name=name, extra={"openai_cli": tool_choice})
|
||||
return ToolChoice(type=ToolChoiceType.AUTO, extra={"openai_cli": tool_choice})
|
||||
|
||||
return ToolChoice(type=ToolChoiceType.AUTO, extra={"raw": tool_choice})
|
||||
|
||||
def _tool_choice_to_openai(self, tool_choice: ToolChoice) -> Union[str, Dict[str, Any]]:
|
||||
if tool_choice.type == ToolChoiceType.NONE:
|
||||
return "none"
|
||||
if tool_choice.type == ToolChoiceType.AUTO:
|
||||
return "auto"
|
||||
if tool_choice.type == ToolChoiceType.REQUIRED:
|
||||
return "required"
|
||||
if tool_choice.type == ToolChoiceType.TOOL:
|
||||
return {"type": "function", "function": {"name": tool_choice.tool_name or ""}}
|
||||
return "auto"
|
||||
|
||||
def _role_from_value(self, role: Any) -> Role:
|
||||
value = str(role or "").lower()
|
||||
if value == "user":
|
||||
return Role.USER
|
||||
if value == "assistant":
|
||||
return Role.ASSISTANT
|
||||
if value == "system":
|
||||
return Role.SYSTEM
|
||||
if value == "developer":
|
||||
return Role.DEVELOPER
|
||||
if value == "tool":
|
||||
return Role.TOOL
|
||||
return Role.UNKNOWN
|
||||
|
||||
def _role_to_openai(self, role: Role) -> str:
|
||||
if role == Role.USER:
|
||||
return "user"
|
||||
if role == Role.ASSISTANT:
|
||||
return "assistant"
|
||||
if role == Role.SYSTEM:
|
||||
return "system"
|
||||
if role == Role.DEVELOPER:
|
||||
return "developer"
|
||||
if role == Role.TOOL:
|
||||
return "tool"
|
||||
return "user"
|
||||
|
||||
def _optional_int(self, value: Any) -> Optional[int]:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
def _optional_float(self, value: Any) -> Optional[float]:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
def _coerce_str_list(self, value: Any) -> Optional[List[str]]:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
return [value]
|
||||
if isinstance(value, list):
|
||||
out: List[str] = []
|
||||
for item in value:
|
||||
if item is None:
|
||||
continue
|
||||
out.append(str(item))
|
||||
return out
|
||||
return [str(value)]
|
||||
|
||||
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] = []
|
||||
for seg in internal.instructions:
|
||||
if seg.text:
|
||||
parts.append(seg.text)
|
||||
return "\n\n".join(parts)
|
||||
return internal.system or ""
|
||||
|
||||
def _error_type_from_value(self, value: str) -> ErrorType:
|
||||
for t in ErrorType:
|
||||
if t.value == value:
|
||||
return t
|
||||
return ErrorType.UNKNOWN
|
||||
|
||||
|
||||
__all__ = ["OpenAICliNormalizer"]
|
||||
@@ -1,54 +0,0 @@
|
||||
"""
|
||||
转换器协议定义
|
||||
|
||||
定义转换器必须实现的方法签名,用于类型检查和文档说明。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List, Protocol, runtime_checkable
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class RequestConverter(Protocol):
|
||||
"""请求转换器协议"""
|
||||
|
||||
def convert_request(self, request: Dict[str, Any]) -> Dict[str, Any]: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ResponseConverter(Protocol):
|
||||
"""响应转换器协议"""
|
||||
|
||||
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class StreamChunkConverter(Protocol):
|
||||
"""
|
||||
流式响应块转换器协议
|
||||
|
||||
统一签名:(chunk, state) -> List[Dict]
|
||||
|
||||
说明:
|
||||
- chunk: 单个流式事件/块
|
||||
- state: 跨 chunk 的转换状态(StreamConversionState 或 GeminiStreamConversionState)
|
||||
- 返回: 转换后的事件列表(可能 0-N 个)
|
||||
|
||||
注意:
|
||||
- 所有流式转换器必须实现此签名
|
||||
- state 参数用于跨 chunk 维护状态(如累积文本、消息 ID 等)
|
||||
"""
|
||||
|
||||
def convert_stream_chunk(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
state: Any,
|
||||
) -> List[Dict[str, Any]]: ...
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RequestConverter",
|
||||
"ResponseConverter",
|
||||
"StreamChunkConverter",
|
||||
]
|
||||
@@ -1,120 +1,68 @@
|
||||
"""
|
||||
格式转换器注册表(核心层)
|
||||
格式转换注册表(Canonical / Hub-and-Spoke)
|
||||
|
||||
自动管理不同 API 格式之间的转换器,支持:
|
||||
- 请求转换:客户端格式 → Provider 格式
|
||||
- 响应转换:Provider 格式 → 客户端格式
|
||||
实现路径:
|
||||
source -> internal -> target
|
||||
|
||||
说明:
|
||||
- 该注册表位于 core 层,避免 services 依赖 api/handlers。
|
||||
- 具体转换器的注册(例如 Claude/OpenAI/Gemini)应由应用启动层完成,
|
||||
或由 api 层的 bootstrap 逻辑完成,以保持依赖方向:api -> core。
|
||||
- 旧 N×N converters 已移除;这里是唯一的格式转换实现。
|
||||
- 转换失败将抛出 `FormatConversionError`(不再静默回退)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from contextlib import contextmanager
|
||||
from typing import TYPE_CHECKING, Any, Dict, Generator, Optional, Tuple, Union
|
||||
from typing import Any, Dict, Generator, List, Optional
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.core.metrics import format_conversion_duration_seconds, format_conversion_total
|
||||
|
||||
from .exceptions import FormatConversionError
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .state import (
|
||||
ClaudeStreamConversionState,
|
||||
GeminiStreamConversionState,
|
||||
OpenAIStreamConversionState,
|
||||
StreamConversionState,
|
||||
)
|
||||
from src.core.api_format.conversion.exceptions import FormatConversionError
|
||||
from src.core.api_format.conversion.normalizer import FormatNormalizer
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _track_conversion_metrics(
|
||||
direction: str, source: str, target: str
|
||||
direction: str,
|
||||
source: str,
|
||||
target: str,
|
||||
) -> Generator[None, None, None]:
|
||||
"""
|
||||
跟踪转换指标的上下文管理器
|
||||
|
||||
Args:
|
||||
direction: 转换方向(request/response/stream)
|
||||
source: 源格式(大写)
|
||||
target: 目标格式(大写)
|
||||
|
||||
Yields:
|
||||
None - 执行转换逻辑
|
||||
"""
|
||||
start = time.perf_counter()
|
||||
status = "success"
|
||||
try:
|
||||
yield
|
||||
format_conversion_total.labels(direction, source, target, "success").inc()
|
||||
except Exception:
|
||||
status = "error"
|
||||
format_conversion_total.labels(direction, source, target, "error").inc()
|
||||
raise
|
||||
finally:
|
||||
format_conversion_total.labels(direction, source, target, status).inc()
|
||||
format_conversion_duration_seconds.labels(direction, source, target).observe(
|
||||
time.perf_counter() - start
|
||||
)
|
||||
|
||||
|
||||
class FormatConverterRegistry:
|
||||
"""
|
||||
格式转换器注册表
|
||||
|
||||
管理不同 API 格式之间的双向转换器
|
||||
"""
|
||||
class FormatConversionRegistry:
|
||||
"""基于 Normalizer 的格式转换注册表"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
# key: (source_format, target_format), value: converter instance
|
||||
self._converters: Dict[Tuple[str, str], Any] = {}
|
||||
self._normalizers: Dict[str, FormatNormalizer] = {}
|
||||
|
||||
def register(
|
||||
self,
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
converter: Any,
|
||||
) -> None:
|
||||
"""
|
||||
注册格式转换器
|
||||
def register(self, normalizer: FormatNormalizer) -> None:
|
||||
self._normalizers[str(normalizer.FORMAT_ID).upper()] = normalizer
|
||||
logger.info(f"[FormatConversionRegistry] 注册 normalizer: {normalizer.FORMAT_ID}")
|
||||
|
||||
Args:
|
||||
source_format: 源格式(如 "CLAUDE", "OPENAI", "GEMINI")
|
||||
target_format: 目标格式
|
||||
converter: 转换器实例(需要有 convert_request/convert_response 方法)
|
||||
"""
|
||||
key = (source_format.upper(), target_format.upper())
|
||||
self._converters[key] = converter
|
||||
logger.info(f"[ConverterRegistry] 注册转换器: {source_format} -> {target_format}")
|
||||
def get_normalizer(self, format_id: str) -> Optional[FormatNormalizer]:
|
||||
return self._normalizers.get(str(format_id).upper())
|
||||
|
||||
def get_converter(
|
||||
self,
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
) -> Optional[Any]:
|
||||
"""
|
||||
获取转换器
|
||||
def _require_normalizer(self, format_id: str) -> FormatNormalizer:
|
||||
normalizer = self.get_normalizer(format_id)
|
||||
if normalizer is None:
|
||||
raise FormatConversionError(format_id, format_id, f"未注册 Normalizer: {format_id}")
|
||||
return normalizer
|
||||
|
||||
Args:
|
||||
source_format: 源格式
|
||||
target_format: 目标格式
|
||||
|
||||
Returns:
|
||||
转换器实例,如果不存在返回 None
|
||||
"""
|
||||
key = (source_format.upper(), target_format.upper())
|
||||
return self._converters.get(key)
|
||||
|
||||
def has_converter(
|
||||
self,
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
) -> bool:
|
||||
"""检查是否存在转换器"""
|
||||
key = (source_format.upper(), target_format.upper())
|
||||
return key in self._converters
|
||||
# ==================== 请求/响应转换(严格) ====================
|
||||
|
||||
def convert_request(
|
||||
self,
|
||||
@@ -122,43 +70,18 @@ class FormatConverterRegistry:
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
转换请求
|
||||
|
||||
Args:
|
||||
request: 原始请求字典
|
||||
source_format: 源格式(客户端格式)
|
||||
target_format: 目标格式(Provider 格式)
|
||||
|
||||
Returns:
|
||||
转换后的请求字典,如果无需转换或没有转换器则返回原始请求
|
||||
"""
|
||||
# 同格式无需转换
|
||||
if source_format.upper() == target_format.upper():
|
||||
if str(source_format).upper() == str(target_format).upper():
|
||||
return request
|
||||
|
||||
converter = self.get_converter(source_format, target_format)
|
||||
if converter is None:
|
||||
logger.warning(
|
||||
f"[ConverterRegistry] 未找到请求转换器: {source_format} -> {target_format},返回原始请求"
|
||||
)
|
||||
return request
|
||||
src = self._require_normalizer(source_format)
|
||||
tgt = self._require_normalizer(target_format)
|
||||
|
||||
if not hasattr(converter, "convert_request"):
|
||||
logger.warning(
|
||||
f"[ConverterRegistry] 转换器缺少 convert_request 方法: {source_format} -> {target_format}"
|
||||
)
|
||||
return request
|
||||
|
||||
try:
|
||||
converted: Dict[str, Any] = converter.convert_request(request)
|
||||
logger.debug(f"[ConverterRegistry] 请求转换成功: {source_format} -> {target_format}")
|
||||
return converted
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"[ConverterRegistry] 请求转换失败: {source_format} -> {target_format}: {e}"
|
||||
)
|
||||
return request
|
||||
with _track_conversion_metrics("request", str(source_format).upper(), str(target_format).upper()):
|
||||
try:
|
||||
internal = src.request_to_internal(request)
|
||||
return tgt.request_from_internal(internal)
|
||||
except Exception as e:
|
||||
raise FormatConversionError(source_format, target_format, str(e)) from e
|
||||
|
||||
def convert_response(
|
||||
self,
|
||||
@@ -166,292 +89,168 @@ class FormatConverterRegistry:
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
转换响应
|
||||
|
||||
Args:
|
||||
response: 原始响应字典
|
||||
source_format: 源格式(Provider 格式)
|
||||
target_format: 目标格式(客户端格式)
|
||||
|
||||
Returns:
|
||||
转换后的响应字典,如果无需转换或没有转换器则返回原始响应
|
||||
"""
|
||||
# 同格式无需转换
|
||||
if source_format.upper() == target_format.upper():
|
||||
if str(source_format).upper() == str(target_format).upper():
|
||||
return response
|
||||
|
||||
converter = self.get_converter(source_format, target_format)
|
||||
if converter is None:
|
||||
logger.warning(
|
||||
f"[ConverterRegistry] 未找到响应转换器: {source_format} -> {target_format},返回原始响应"
|
||||
src = self._require_normalizer(source_format)
|
||||
tgt = self._require_normalizer(target_format)
|
||||
|
||||
with _track_conversion_metrics("response", str(source_format).upper(), str(target_format).upper()):
|
||||
try:
|
||||
internal = src.response_to_internal(response)
|
||||
return tgt.response_from_internal(internal)
|
||||
except Exception as e:
|
||||
raise FormatConversionError(source_format, target_format, str(e)) from e
|
||||
|
||||
def convert_error_response(
|
||||
self,
|
||||
error_response: Dict[str, Any],
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
) -> Dict[str, Any]:
|
||||
if str(source_format).upper() == str(target_format).upper():
|
||||
return error_response
|
||||
|
||||
src = self._require_normalizer(source_format)
|
||||
tgt = self._require_normalizer(target_format)
|
||||
|
||||
if not (src.capabilities.supports_error_conversion and tgt.capabilities.supports_error_conversion):
|
||||
raise FormatConversionError(
|
||||
source_format,
|
||||
target_format,
|
||||
"source/target normalizer 不支持错误转换",
|
||||
)
|
||||
return response
|
||||
|
||||
if not hasattr(converter, "convert_response"):
|
||||
logger.warning(
|
||||
f"[ConverterRegistry] 转换器缺少 convert_response 方法: {source_format} -> {target_format}"
|
||||
)
|
||||
return response
|
||||
with _track_conversion_metrics("error", str(source_format).upper(), str(target_format).upper()):
|
||||
try:
|
||||
internal = src.error_to_internal(error_response)
|
||||
return tgt.error_from_internal(internal)
|
||||
except Exception as e:
|
||||
raise FormatConversionError(source_format, target_format, str(e)) from e
|
||||
|
||||
try:
|
||||
converted: Dict[str, Any] = converter.convert_response(response)
|
||||
logger.debug(f"[ConverterRegistry] 响应转换成功: {source_format} -> {target_format}")
|
||||
return converted
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"[ConverterRegistry] 响应转换失败: {source_format} -> {target_format}: {e}"
|
||||
)
|
||||
return response
|
||||
# ==================== 流式转换(严格) ====================
|
||||
|
||||
def convert_stream_chunk(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
state: Optional[Union["StreamConversionState", "GeminiStreamConversionState"]] = None,
|
||||
) -> list[Dict[str, Any]]:
|
||||
"""
|
||||
转换流式响应块
|
||||
|
||||
Args:
|
||||
chunk: 原始流式响应块
|
||||
source_format: 源格式(Provider 格式)
|
||||
target_format: 目标格式(客户端格式)
|
||||
state: 流式转换状态(StreamConversionState 或 GeminiStreamConversionState)
|
||||
|
||||
Returns:
|
||||
转换后的事件列表(可能 0-N 个),失败时返回原始 chunk 的单元素列表
|
||||
"""
|
||||
# 同格式无需转换
|
||||
if source_format.upper() == target_format.upper():
|
||||
state: Optional[StreamState] = None,
|
||||
) -> List[Dict[str, Any]]:
|
||||
if str(source_format).upper() == str(target_format).upper():
|
||||
return [chunk]
|
||||
|
||||
converter = self.get_converter(source_format, target_format)
|
||||
if converter is None:
|
||||
return [chunk]
|
||||
src = self._require_normalizer(source_format)
|
||||
tgt = self._require_normalizer(target_format)
|
||||
|
||||
# 使用流式转换方法
|
||||
if hasattr(converter, "convert_stream_chunk"):
|
||||
if not (src.capabilities.supports_stream and tgt.capabilities.supports_stream):
|
||||
raise FormatConversionError(
|
||||
source_format,
|
||||
target_format,
|
||||
"source/target normalizer 不支持流式转换",
|
||||
)
|
||||
|
||||
if state is None:
|
||||
state = StreamState()
|
||||
|
||||
with _track_conversion_metrics("stream", str(source_format).upper(), str(target_format).upper()):
|
||||
try:
|
||||
result: list[Dict[str, Any]] = converter.convert_stream_chunk(chunk, state)
|
||||
return result
|
||||
events = src.stream_chunk_to_internal(chunk, state)
|
||||
out: List[Dict[str, Any]] = []
|
||||
for event in events:
|
||||
out.extend(tgt.stream_event_from_internal(event, state))
|
||||
return out
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"[ConverterRegistry] 流式块转换失败: {source_format} -> {target_format}: {e}"
|
||||
)
|
||||
return [chunk]
|
||||
raise FormatConversionError(source_format, target_format, str(e)) from e
|
||||
|
||||
# 降级到普通响应转换(作为单个事件返回)
|
||||
if hasattr(converter, "convert_response"):
|
||||
try:
|
||||
converted: Dict[str, Any] = converter.convert_response(chunk)
|
||||
return [converted]
|
||||
except Exception:
|
||||
return [chunk]
|
||||
# ==================== 能力查询 ====================
|
||||
|
||||
return [chunk]
|
||||
def can_convert_request(self, source_format: str, target_format: str) -> bool:
|
||||
if str(source_format).upper() == str(target_format).upper():
|
||||
return True
|
||||
return self.get_normalizer(source_format) is not None and self.get_normalizer(target_format) is not None
|
||||
|
||||
def list_converters(self) -> list[Tuple[str, str]]:
|
||||
"""列出所有已注册的转换器"""
|
||||
return list(self._converters.keys())
|
||||
def can_convert_response(self, source_format: str, target_format: str) -> bool:
|
||||
return self.can_convert_request(source_format, target_format)
|
||||
|
||||
# ========== 能力查询方法 ==========
|
||||
|
||||
def can_convert_request(self, source: str, target: str) -> bool:
|
||||
"""检查是否支持请求转换"""
|
||||
converter = self.get_converter(source, target)
|
||||
return converter is not None and hasattr(converter, "convert_request")
|
||||
|
||||
def can_convert_response(self, source: str, target: str) -> bool:
|
||||
"""检查是否支持响应转换"""
|
||||
converter = self.get_converter(source, target)
|
||||
return converter is not None and hasattr(converter, "convert_response")
|
||||
|
||||
def can_convert_stream(self, source: str, target: str) -> bool:
|
||||
"""检查是否支持流式转换"""
|
||||
converter = self.get_converter(source, target)
|
||||
if converter is None:
|
||||
def can_convert_stream(self, source_format: str, target_format: str) -> bool:
|
||||
if str(source_format).upper() == str(target_format).upper():
|
||||
return True
|
||||
src = self.get_normalizer(source_format)
|
||||
tgt = self.get_normalizer(target_format)
|
||||
if src is None or tgt is None:
|
||||
return False
|
||||
return hasattr(converter, "convert_stream_chunk")
|
||||
return bool(src.capabilities.supports_stream and tgt.capabilities.supports_stream)
|
||||
|
||||
def can_convert_full(
|
||||
self,
|
||||
source: str,
|
||||
target: str,
|
||||
require_stream: bool = False,
|
||||
) -> bool:
|
||||
"""
|
||||
检查是否支持完整的双向转换
|
||||
def can_convert_error(self, source_format: str, target_format: str) -> bool:
|
||||
if str(source_format).upper() == str(target_format).upper():
|
||||
return True
|
||||
src = self.get_normalizer(source_format)
|
||||
tgt = self.get_normalizer(target_format)
|
||||
if src is None or tgt is None:
|
||||
return False
|
||||
return bool(src.capabilities.supports_error_conversion and tgt.capabilities.supports_error_conversion)
|
||||
|
||||
对于跨格式请求,需要:
|
||||
1. 请求转换:source -> target
|
||||
2. 响应转换:target -> source(注意方向相反)
|
||||
3. 流式转换(如果 require_stream=True):target -> source
|
||||
|
||||
Args:
|
||||
source: 客户端格式
|
||||
target: Provider 格式
|
||||
require_stream: 是否要求支持流式转换
|
||||
"""
|
||||
# 请求:client -> provider
|
||||
if not self.can_convert_request(source, target):
|
||||
def can_convert_full(self, format_a: str, format_b: str, *, require_stream: bool = False) -> bool:
|
||||
if not self.can_convert_request(format_a, format_b):
|
||||
return False
|
||||
# 响应:provider -> client(方向相反)
|
||||
if not self.can_convert_response(target, source):
|
||||
return False
|
||||
# 流式:provider -> client(方向相反)
|
||||
if require_stream and not self.can_convert_stream(target, source):
|
||||
if not self.can_convert_request(format_b, format_a):
|
||||
return False
|
||||
if require_stream:
|
||||
return self.can_convert_stream(format_a, format_b) and self.can_convert_stream(format_b, format_a)
|
||||
return True
|
||||
|
||||
def get_supported_targets(self, source: str) -> list[str]:
|
||||
"""获取指定源格式支持转换到的目标格式列表"""
|
||||
source_upper = source.upper()
|
||||
return [target for (src, target) in self._converters.keys() if src == source_upper]
|
||||
def list_normalizers(self) -> List[str]:
|
||||
return sorted(self._normalizers.keys())
|
||||
|
||||
# ========== 严格模式方法 ==========
|
||||
|
||||
def convert_request_strict(
|
||||
self,
|
||||
request: Dict[str, Any],
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
严格模式请求转换 - 失败时抛出异常
|
||||
|
||||
用于需要故障转移的场景:转换失败会抛出 FormatConversionError,
|
||||
让 Orchestrator 可以尝试下一个候选。
|
||||
|
||||
Raises:
|
||||
FormatConversionError: 转换失败时抛出
|
||||
"""
|
||||
source_upper = source_format.upper()
|
||||
target_upper = target_format.upper()
|
||||
|
||||
# 同格式无需转换
|
||||
if source_upper == target_upper:
|
||||
return request
|
||||
|
||||
converter = self.get_converter(source_format, target_format)
|
||||
if converter is None:
|
||||
raise FormatConversionError(source_format, target_format, "未找到转换器")
|
||||
|
||||
if not hasattr(converter, "convert_request"):
|
||||
raise FormatConversionError(
|
||||
source_format, target_format, "转换器缺少 convert_request 方法"
|
||||
)
|
||||
|
||||
with _track_conversion_metrics("request", source_upper, target_upper):
|
||||
try:
|
||||
converted: Dict[str, Any] = converter.convert_request(request)
|
||||
logger.debug(f"[ConverterRegistry] 请求转换成功: {source_format} -> {target_format}")
|
||||
return converted
|
||||
except FormatConversionError:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise FormatConversionError(source_format, target_format, str(e)) from e
|
||||
|
||||
def convert_response_strict(
|
||||
self,
|
||||
response: Dict[str, Any],
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
严格模式响应转换 - 失败时抛出异常
|
||||
|
||||
Raises:
|
||||
FormatConversionError: 转换失败时抛出
|
||||
"""
|
||||
source_upper = source_format.upper()
|
||||
target_upper = target_format.upper()
|
||||
|
||||
if source_upper == target_upper:
|
||||
return response
|
||||
|
||||
converter = self.get_converter(source_format, target_format)
|
||||
if converter is None:
|
||||
raise FormatConversionError(source_format, target_format, "未找到转换器")
|
||||
|
||||
if not hasattr(converter, "convert_response"):
|
||||
raise FormatConversionError(
|
||||
source_format, target_format, "转换器缺少 convert_response 方法"
|
||||
)
|
||||
|
||||
with _track_conversion_metrics("response", source_upper, target_upper):
|
||||
try:
|
||||
converted: Dict[str, Any] = converter.convert_response(response)
|
||||
logger.debug(f"[ConverterRegistry] 响应转换成功: {source_format} -> {target_format}")
|
||||
return converted
|
||||
except FormatConversionError:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise FormatConversionError(source_format, target_format, str(e)) from e
|
||||
|
||||
def convert_stream_chunk_strict(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
source_format: str,
|
||||
target_format: str,
|
||||
state: Optional[
|
||||
Union[
|
||||
"StreamConversionState",
|
||||
"GeminiStreamConversionState",
|
||||
"ClaudeStreamConversionState",
|
||||
"OpenAIStreamConversionState",
|
||||
]
|
||||
] = None,
|
||||
) -> list[Dict[str, Any]]:
|
||||
"""
|
||||
严格模式流式块转换 - 失败时抛出异常
|
||||
|
||||
Args:
|
||||
chunk: 流式响应块
|
||||
source_format: 源格式
|
||||
target_format: 目标格式
|
||||
state: 流式转换状态(StreamConversionState 或 GeminiStreamConversionState)
|
||||
|
||||
Returns:
|
||||
转换后的事件列表(可能 0-N 个)
|
||||
|
||||
Raises:
|
||||
FormatConversionError: 转换失败时抛出
|
||||
"""
|
||||
source_upper = source_format.upper()
|
||||
target_upper = target_format.upper()
|
||||
|
||||
if source_upper == target_upper:
|
||||
return [chunk]
|
||||
|
||||
converter = self.get_converter(source_format, target_format)
|
||||
if converter is None:
|
||||
raise FormatConversionError(source_format, target_format, "未找到转换器")
|
||||
|
||||
if not hasattr(converter, "convert_stream_chunk"):
|
||||
raise FormatConversionError(
|
||||
source_format, target_format, "转换器缺少 convert_stream_chunk 方法"
|
||||
)
|
||||
|
||||
with _track_conversion_metrics("stream", source_upper, target_upper):
|
||||
try:
|
||||
result: list[Dict[str, Any]] = converter.convert_stream_chunk(chunk, state)
|
||||
return result
|
||||
except FormatConversionError:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise FormatConversionError(
|
||||
source_format, target_format, f"流式块转换失败: {e}"
|
||||
) from e
|
||||
def get_supported_targets(self, source_format: str) -> List[str]:
|
||||
src = str(source_format).upper()
|
||||
if src not in self._normalizers:
|
||||
return []
|
||||
return [k for k in self._normalizers.keys() if k != src]
|
||||
|
||||
|
||||
# 全局单例
|
||||
converter_registry = FormatConverterRegistry()
|
||||
# 全局注册表(唯一实现)
|
||||
format_conversion_registry = FormatConversionRegistry()
|
||||
_DEFAULT_NORMALIZERS_REGISTERED = False
|
||||
_REGISTRATION_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def register_default_normalizers() -> None:
|
||||
"""注册默认 Normalizers(OPENAI/CLAUDE/GEMINI + *_CLI)"""
|
||||
global _DEFAULT_NORMALIZERS_REGISTERED # noqa: PLW0603 - module-level 缓存
|
||||
|
||||
# 快速路径:已注册则直接返回(无锁)
|
||||
if _DEFAULT_NORMALIZERS_REGISTERED:
|
||||
return
|
||||
|
||||
# 慢路径:加锁后双重检查
|
||||
with _REGISTRATION_LOCK:
|
||||
if _DEFAULT_NORMALIZERS_REGISTERED:
|
||||
return
|
||||
|
||||
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer
|
||||
from src.core.api_format.conversion.normalizers.claude_cli import ClaudeCliNormalizer
|
||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||
from src.core.api_format.conversion.normalizers.gemini_cli import GeminiCliNormalizer
|
||||
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||
from src.core.api_format.conversion.normalizers.openai_cli import OpenAICliNormalizer
|
||||
|
||||
format_conversion_registry.register(OpenAINormalizer())
|
||||
format_conversion_registry.register(OpenAICliNormalizer())
|
||||
format_conversion_registry.register(ClaudeNormalizer())
|
||||
format_conversion_registry.register(ClaudeCliNormalizer())
|
||||
format_conversion_registry.register(GeminiNormalizer())
|
||||
format_conversion_registry.register(GeminiCliNormalizer())
|
||||
|
||||
_DEFAULT_NORMALIZERS_REGISTERED = True
|
||||
logger.info(
|
||||
f"[FormatConversionRegistry] 已注册 {len(format_conversion_registry.list_normalizers())} 个 normalizer"
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"FormatConverterRegistry",
|
||||
"converter_registry",
|
||||
"FormatConversionError",
|
||||
"FormatConversionRegistry",
|
||||
"format_conversion_registry",
|
||||
"register_default_normalizers",
|
||||
]
|
||||
|
||||
@@ -1,114 +0,0 @@
|
||||
"""
|
||||
流式转换状态类
|
||||
|
||||
用于在多个 chunk 之间维护转换上下文,例如:
|
||||
- 是否已发送 message_start 事件
|
||||
- 累积的文本(用于计算增量)
|
||||
- 当前内容块索引
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamConversionState:
|
||||
"""
|
||||
SSE 流式转换状态(Claude <-> OpenAI)
|
||||
|
||||
用于跨 chunk 维护转换状态,确保生成正确的事件序列。
|
||||
"""
|
||||
|
||||
message_id: str = ""
|
||||
model: str = ""
|
||||
message_started: bool = False
|
||||
content_block_started: bool = False
|
||||
current_tool_index: int = 0
|
||||
|
||||
def reset(self) -> None:
|
||||
"""重置状态(重试时调用)"""
|
||||
self.message_started = False
|
||||
self.content_block_started = False
|
||||
self.current_tool_index = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class GeminiStreamConversionState:
|
||||
"""
|
||||
Gemini 流式转换状态(JSON 数组格式)
|
||||
|
||||
与 Claude/OpenAI 的 SSE 不同,Gemini 需要追踪额外状态:
|
||||
1. 累积文本(用于计算真正的增量)
|
||||
2. 内容块索引(用于工具调用)
|
||||
3. 是否已发送 message_start
|
||||
"""
|
||||
|
||||
message_id: str = ""
|
||||
model: str = ""
|
||||
accumulated_text: str = "" # 累积的文本(用于计算增量)
|
||||
message_started: bool = False
|
||||
content_block_started: bool = False
|
||||
current_block_index: int = 0
|
||||
tool_call_index: int = 0 # 工具调用计数
|
||||
has_sent_usage: bool = False # 是否已发送 usage
|
||||
|
||||
def reset(self) -> None:
|
||||
"""重置状态(重试时调用)"""
|
||||
self.accumulated_text = ""
|
||||
self.message_started = False
|
||||
self.content_block_started = False
|
||||
self.current_block_index = 0
|
||||
self.tool_call_index = 0
|
||||
self.has_sent_usage = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class ClaudeStreamConversionState:
|
||||
"""
|
||||
Claude -> Gemini 流式转换状态
|
||||
|
||||
用于将 Claude SSE 事件流转换为 Gemini JSON 流式响应
|
||||
"""
|
||||
|
||||
message_id: str = ""
|
||||
model: str = ""
|
||||
current_block_type: str = "" # 当前内容块类型(text/tool_use)
|
||||
current_block_index: int = 0
|
||||
current_tool_name: str = ""
|
||||
current_tool_id: str = ""
|
||||
accumulated_tool_input: str = "" # 累积的工具输入 JSON
|
||||
|
||||
def reset(self) -> None:
|
||||
"""重置状态(重试时调用)"""
|
||||
self.current_block_type = ""
|
||||
self.current_block_index = 0
|
||||
self.current_tool_name = ""
|
||||
self.current_tool_id = ""
|
||||
self.accumulated_tool_input = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class OpenAIStreamConversionState:
|
||||
"""
|
||||
OpenAI -> Gemini 流式转换状态
|
||||
|
||||
用于将 OpenAI SSE 事件流转换为 Gemini JSON 流式响应
|
||||
"""
|
||||
|
||||
model: str = ""
|
||||
current_tool_name: str = ""
|
||||
accumulated_tool_args: str = "" # 累积的工具参数 JSON
|
||||
|
||||
def reset(self) -> None:
|
||||
"""重置状态(重试时调用)"""
|
||||
self.current_tool_name = ""
|
||||
self.accumulated_tool_args = ""
|
||||
|
||||
|
||||
__all__ = [
|
||||
"StreamConversionState",
|
||||
"GeminiStreamConversionState",
|
||||
"ClaudeStreamConversionState",
|
||||
"OpenAIStreamConversionState",
|
||||
]
|
||||
147
src/core/api_format/conversion/stream_events.py
Normal file
147
src/core/api_format/conversion/stream_events.py
Normal file
@@ -0,0 +1,147 @@
|
||||
"""
|
||||
类型安全的流式事件定义(InternalStreamEvent)
|
||||
|
||||
用于把 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 .internal import ContentType, InternalError, StopReason, UsageInfo
|
||||
|
||||
|
||||
class StreamEventType(str, Enum):
|
||||
"""流式事件类型"""
|
||||
|
||||
MESSAGE_START = "message_start"
|
||||
CONTENT_BLOCK_START = "content_block_start"
|
||||
CONTENT_DELTA = "content_delta"
|
||||
TOOL_CALL_DELTA = "tool_call_delta"
|
||||
CONTENT_BLOCK_STOP = "content_block_stop"
|
||||
MESSAGE_STOP = "message_stop"
|
||||
USAGE = "usage"
|
||||
ERROR = "error"
|
||||
UNKNOWN = "unknown"
|
||||
|
||||
|
||||
@dataclass
|
||||
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)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ContentBlockStartEvent:
|
||||
"""内容块开始事件"""
|
||||
|
||||
type: StreamEventType = field(default=StreamEventType.CONTENT_BLOCK_START, init=False)
|
||||
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)
|
||||
|
||||
|
||||
@dataclass
|
||||
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)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolCallDeltaEvent:
|
||||
"""工具调用增量事件(工具输入 JSON 的字符串片段)"""
|
||||
|
||||
type: StreamEventType = field(default=StreamEventType.TOOL_CALL_DELTA, init=False)
|
||||
block_index: int = 0
|
||||
tool_id: str = ""
|
||||
input_delta: str = "" # JSON 字符串片段
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ContentBlockStopEvent:
|
||||
"""内容块结束事件"""
|
||||
|
||||
type: StreamEventType = field(default=StreamEventType.CONTENT_BLOCK_STOP, init=False)
|
||||
block_index: int = 0
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
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)
|
||||
|
||||
|
||||
@dataclass
|
||||
class UsageEvent:
|
||||
"""使用量事件"""
|
||||
|
||||
type: StreamEventType = field(default=StreamEventType.USAGE, init=False)
|
||||
usage: UsageInfo = field(default_factory=UsageInfo)
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ErrorEvent:
|
||||
"""错误事件"""
|
||||
|
||||
type: StreamEventType = field(default=StreamEventType.ERROR, init=False)
|
||||
error: InternalError
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
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)
|
||||
|
||||
|
||||
InternalStreamEvent = Union[
|
||||
MessageStartEvent,
|
||||
ContentBlockStartEvent,
|
||||
ContentDeltaEvent,
|
||||
ToolCallDeltaEvent,
|
||||
ContentBlockStopEvent,
|
||||
MessageStopEvent,
|
||||
UsageEvent,
|
||||
ErrorEvent,
|
||||
UnknownStreamEvent,
|
||||
]
|
||||
|
||||
|
||||
__all__ = [
|
||||
"StreamEventType",
|
||||
"MessageStartEvent",
|
||||
"ContentBlockStartEvent",
|
||||
"ContentDeltaEvent",
|
||||
"ToolCallDeltaEvent",
|
||||
"ContentBlockStopEvent",
|
||||
"MessageStopEvent",
|
||||
"UsageEvent",
|
||||
"ErrorEvent",
|
||||
"UnknownStreamEvent",
|
||||
"InternalStreamEvent",
|
||||
]
|
||||
50
src/core/api_format/conversion/stream_state.py
Normal file
50
src/core/api_format/conversion/stream_state.py
Normal file
@@ -0,0 +1,50 @@
|
||||
"""
|
||||
统一流式状态容器(StreamState)
|
||||
|
||||
目标:在多个 chunk 之间维护转换上下文,但避免把“某个格式特定的状态字段”固化在核心层。
|
||||
每个 Normalizer 通过 `substate(format_id)` 获取自己的隔离状态字典。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamState:
|
||||
"""
|
||||
统一的流式状态容器
|
||||
|
||||
关键点:
|
||||
- 不把具体格式字段固化为属性,避免 source/target 互相污染
|
||||
- 每个 Normalizer 只读写自己的隔离子状态:`state.substate(self.FORMAT_ID)`
|
||||
"""
|
||||
|
||||
# 可选:便于调试与链路追踪(不强依赖)
|
||||
model: str = ""
|
||||
message_id: str = ""
|
||||
|
||||
# Registry/调用层的通用扩展信息(与具体格式无关)
|
||||
extra: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
# 各 Normalizer 的隔离状态(key: FORMAT_ID)
|
||||
by_format: Dict[str, Dict[str, Any]] = field(default_factory=dict)
|
||||
|
||||
def substate(self, format_id: str) -> Dict[str, Any]:
|
||||
"""获取指定格式的隔离子状态"""
|
||||
key = str(format_id).upper()
|
||||
return self.by_format.setdefault(key, {})
|
||||
|
||||
def reset(self) -> None:
|
||||
"""重置状态(重试时调用)"""
|
||||
self.model = ""
|
||||
self.message_id = ""
|
||||
self.extra.clear()
|
||||
self.by_format.clear()
|
||||
|
||||
|
||||
__all__ = [
|
||||
"StreamState",
|
||||
]
|
||||
|
||||
@@ -16,7 +16,8 @@ def is_cli_format(format_id: Union[str, "APIFormat", None]) -> bool:
|
||||
"""
|
||||
判断是否为 CLI 透传格式
|
||||
|
||||
CLI 格式以 _CLI 结尾,不参与格式转换,请求直接透传。
|
||||
CLI 格式以 _CLI 结尾,表示该入口更偏向“CLI 兼容层”(鉴权/UA/路径差异等)。
|
||||
是否参与格式转换由转换层决定;当前项目已支持 CLI 格式参与转换。
|
||||
|
||||
Args:
|
||||
format_id: 格式标识符(字符串或 APIFormat 枚举)
|
||||
@@ -96,11 +97,13 @@ def is_same_format(
|
||||
|
||||
def is_convertible_format(format_id: Union[str, "APIFormat", None]) -> bool:
|
||||
"""
|
||||
判断是否为可转换格式(非 CLI)
|
||||
判断是否为可转换格式
|
||||
|
||||
可转换格式可以与其他格式进行双向转换。
|
||||
CLI 格式为透传模式,不参与转换。
|
||||
.. deprecated::
|
||||
此函数语义已退化(对非 None 输入总返回 True)。
|
||||
真正的可转换性应通过 format_conversion_registry.can_convert_*() 查询。
|
||||
保留此函数仅为向后兼容,不建议新代码使用。
|
||||
"""
|
||||
if format_id is None:
|
||||
return False
|
||||
return not is_cli_format(format_id)
|
||||
return True
|
||||
|
||||
@@ -156,9 +156,9 @@ async def lifespan(app: FastAPI):
|
||||
|
||||
# 注册格式转换器
|
||||
logger.info("注册格式转换器...")
|
||||
from src.core.api_format import register_all_converters
|
||||
from src.core.api_format.conversion.registry import register_default_normalizers
|
||||
|
||||
register_all_converters()
|
||||
register_default_normalizers()
|
||||
|
||||
# 初始化功能模块系统
|
||||
logger.info("初始化功能模块系统...")
|
||||
|
||||
@@ -30,7 +30,8 @@ from redis import Redis
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.config.settings import config
|
||||
from src.core.api_format import APIFormat, FormatConversionError
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.api_format.conversion.exceptions import FormatConversionError
|
||||
from src.core.error_utils import extract_error_message
|
||||
from src.core.exceptions import (
|
||||
ConcurrencyLimitError,
|
||||
|
||||
@@ -15,7 +15,7 @@ from unittest.mock import MagicMock, patch
|
||||
import pytest
|
||||
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.core.api_format import StreamConversionState
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
# Mock CliMessageHandlerBase 用于测试
|
||||
@@ -29,7 +29,12 @@ class MockCliHandler:
|
||||
events: list,
|
||||
) -> List[str]:
|
||||
"""复制自 CliMessageHandlerBase._convert_sse_line"""
|
||||
from src.core.api_format import converter_registry
|
||||
from src.core.api_format.conversion import (
|
||||
format_conversion_registry,
|
||||
register_default_normalizers,
|
||||
)
|
||||
|
||||
register_default_normalizers()
|
||||
|
||||
# 如果是空行或特殊控制行,直接返回
|
||||
if not line or line.strip() == "" or line == "data: [DONE]":
|
||||
@@ -50,7 +55,7 @@ class MockCliHandler:
|
||||
|
||||
# 初始化流式转换状态
|
||||
if ctx.stream_conversion_state is None:
|
||||
ctx.stream_conversion_state = StreamConversionState(
|
||||
ctx.stream_conversion_state = StreamState(
|
||||
model=ctx.mapped_model or ctx.model,
|
||||
message_id=ctx.response_id or ctx.request_id,
|
||||
)
|
||||
@@ -59,19 +64,13 @@ class MockCliHandler:
|
||||
client_format = ctx.client_api_format or ""
|
||||
|
||||
try:
|
||||
converted_events = converter_registry.convert_stream_chunk_strict(
|
||||
converted_events = format_conversion_registry.convert_stream_chunk(
|
||||
data_obj,
|
||||
provider_format,
|
||||
client_format,
|
||||
state=ctx.stream_conversion_state,
|
||||
)
|
||||
|
||||
if converted_events and not ctx.stream_conversion_state.message_started:
|
||||
for evt in converted_events:
|
||||
if evt.get("type") == "message_start" or evt.get("choices"):
|
||||
ctx.stream_conversion_state.message_started = True
|
||||
break
|
||||
|
||||
result = []
|
||||
for evt in converted_events:
|
||||
result.append(f"data: {json.dumps(evt, ensure_ascii=False)}")
|
||||
@@ -168,25 +167,13 @@ class TestConvertSseLineOneInManyOut:
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def setup_converters(self):
|
||||
"""注册测试用转换器"""
|
||||
from src.core.api_format import (
|
||||
ClaudeToOpenAIConverter,
|
||||
OpenAIToClaudeConverter,
|
||||
converter_registry,
|
||||
)
|
||||
"""确保 Canonical normalizers 已注册"""
|
||||
from src.core.api_format.conversion import register_default_normalizers
|
||||
|
||||
# 保存原始状态
|
||||
original_converters = converter_registry._converters.copy()
|
||||
|
||||
# 注册转换器
|
||||
converter_registry.register("OPENAI", "CLAUDE", OpenAIToClaudeConverter())
|
||||
converter_registry.register("CLAUDE", "OPENAI", ClaudeToOpenAIConverter())
|
||||
register_default_normalizers()
|
||||
|
||||
yield
|
||||
|
||||
# 恢复原始状态
|
||||
converter_registry._converters = original_converters
|
||||
|
||||
def test_openai_to_claude_conversion(self) -> None:
|
||||
"""测试 OpenAI -> Claude 流式转换"""
|
||||
handler = MockCliHandler()
|
||||
@@ -211,9 +198,6 @@ class TestConvertSseLineOneInManyOut:
|
||||
first_event = json.loads(result1[0][6:])
|
||||
assert first_event.get("type") == "message_start"
|
||||
|
||||
# 状态应更新
|
||||
assert ctx.stream_conversion_state.message_started is True
|
||||
|
||||
def test_claude_to_openai_conversion(self) -> None:
|
||||
"""测试 Claude -> OpenAI 流式转换"""
|
||||
handler = MockCliHandler()
|
||||
@@ -247,7 +231,6 @@ class TestConvertSseLineOneInManyOut:
|
||||
handler._convert_sse_line(ctx, f"data: {json.dumps(chunk1)}", [])
|
||||
|
||||
state_after_first = ctx.stream_conversion_state
|
||||
message_started_after_first = state_after_first.message_started
|
||||
|
||||
# 第二个 chunk
|
||||
chunk2 = {"choices": [{"delta": {"content": "hello"}}]}
|
||||
@@ -255,8 +238,6 @@ class TestConvertSseLineOneInManyOut:
|
||||
|
||||
# 状态应该是同一个对象
|
||||
assert ctx.stream_conversion_state is state_after_first
|
||||
# message_started 应该保持 True
|
||||
assert ctx.stream_conversion_state.message_started is True
|
||||
|
||||
|
||||
class TestStreamContextIntegration:
|
||||
@@ -265,9 +246,7 @@ class TestStreamContextIntegration:
|
||||
def test_stream_conversion_state_reset_on_retry(self) -> None:
|
||||
"""测试重试时重置流式转换状态"""
|
||||
ctx = StreamContext(model="test", api_format="OPENAI")
|
||||
ctx.stream_conversion_state = StreamConversionState(
|
||||
model="test", message_id="123", message_started=True
|
||||
)
|
||||
ctx.stream_conversion_state = StreamState(model="test", message_id="123")
|
||||
|
||||
ctx.reset_for_retry()
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ import pytest
|
||||
from src.api.handlers.base.response_parser import ParsedResponse, ResponseParser, StreamStats
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.api.handlers.base.stream_processor import StreamProcessor
|
||||
from src.core.api_format import register_all_converters
|
||||
from src.core.api_format.conversion import register_default_normalizers
|
||||
|
||||
|
||||
class DummyParser(ResponseParser):
|
||||
@@ -31,7 +31,7 @@ async def _empty_async_iter():
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_response_stream_converts_claude_to_openai() -> None:
|
||||
register_all_converters()
|
||||
register_default_normalizers()
|
||||
|
||||
ctx = StreamContext(model="test-model", api_format="OPENAI")
|
||||
ctx.client_api_format = "OPENAI"
|
||||
|
||||
@@ -6,7 +6,7 @@ import pytest
|
||||
from src.api.handlers.base.parsers import get_parser_for_format
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.api.handlers.base.stream_processor import StreamProcessor
|
||||
from src.core.api_format import GeminiToOpenAIConverter, converter_registry
|
||||
from src.core.api_format.conversion import register_default_normalizers
|
||||
|
||||
|
||||
class _DummyResponseCtx:
|
||||
@@ -26,77 +26,71 @@ async def _iter_bytes(chunks: list[bytes]) -> AsyncIterator[bytes]:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_processor_converts_gemini_json_lines_without_data_prefix() -> None:
|
||||
# Register only what we need, and restore after test.
|
||||
original_converters = converter_registry._converters.copy()
|
||||
try:
|
||||
converter_registry.register("GEMINI", "OPENAI", GeminiToOpenAIConverter())
|
||||
register_default_normalizers()
|
||||
|
||||
ctx = StreamContext(model="gemini-test", api_format="OPENAI")
|
||||
ctx.provider_api_format = "GEMINI"
|
||||
ctx.client_api_format = "OPENAI"
|
||||
ctx.needs_conversion = True
|
||||
ctx.request_id = "req_test"
|
||||
ctx.mapped_model = "gemini-test"
|
||||
ctx = StreamContext(model="gemini-test", api_format="OPENAI")
|
||||
ctx.provider_api_format = "GEMINI"
|
||||
ctx.client_api_format = "OPENAI"
|
||||
ctx.needs_conversion = True
|
||||
ctx.request_id = "req_test"
|
||||
ctx.mapped_model = "gemini-test"
|
||||
|
||||
# Simulate Gemini JSON-array/chunks stream: wrapper lines + two JSON objects.
|
||||
chunk1 = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": [{"text": "Hello"}], "role": "model"},
|
||||
}
|
||||
]
|
||||
}
|
||||
chunk2 = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": [{"text": "Hello world"}], "role": "model"},
|
||||
"finishReason": "STOP",
|
||||
}
|
||||
],
|
||||
"usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 2, "totalTokenCount": 3},
|
||||
}
|
||||
|
||||
upstream_lines = [
|
||||
b"[\n",
|
||||
(json.dumps(chunk1) + ",\n").encode("utf-8"),
|
||||
(json.dumps(chunk2) + "\n").encode("utf-8"),
|
||||
b"]\n",
|
||||
# Simulate Gemini JSON-array/chunks stream: wrapper lines + two JSON objects.
|
||||
chunk1 = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": [{"text": "Hello"}], "role": "model"},
|
||||
}
|
||||
]
|
||||
}
|
||||
chunk2 = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": [{"text": "Hello world"}], "role": "model"},
|
||||
"finishReason": "STOP",
|
||||
}
|
||||
],
|
||||
"usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 2, "totalTokenCount": 3},
|
||||
}
|
||||
|
||||
processor = StreamProcessor(
|
||||
request_id="req_test",
|
||||
default_parser=get_parser_for_format("OPENAI"),
|
||||
)
|
||||
upstream_lines = [
|
||||
b"[\n",
|
||||
(json.dumps(chunk1) + ",\n").encode("utf-8"),
|
||||
(json.dumps(chunk2) + "\n").encode("utf-8"),
|
||||
b"]\n",
|
||||
]
|
||||
|
||||
out = b""
|
||||
async for b in processor.create_response_stream(
|
||||
ctx=ctx,
|
||||
byte_iterator=_iter_bytes(upstream_lines),
|
||||
response_ctx=_DummyResponseCtx(),
|
||||
http_client=_DummyHTTPClient(), # type: ignore[arg-type]
|
||||
prefetched_chunks=[],
|
||||
start_time=None,
|
||||
):
|
||||
out += b
|
||||
processor = StreamProcessor(
|
||||
request_id="req_test",
|
||||
default_parser=get_parser_for_format("OPENAI"),
|
||||
)
|
||||
|
||||
text = out.decode("utf-8", errors="replace")
|
||||
data_lines = [ln for ln in text.splitlines() if ln.startswith("data: ")]
|
||||
out = b""
|
||||
async for b in processor.create_response_stream(
|
||||
ctx=ctx,
|
||||
byte_iterator=_iter_bytes(upstream_lines),
|
||||
response_ctx=_DummyResponseCtx(),
|
||||
http_client=_DummyHTTPClient(), # type: ignore[arg-type]
|
||||
prefetched_chunks=[],
|
||||
start_time=None,
|
||||
):
|
||||
out += b
|
||||
|
||||
# OpenAI termination marker should be present (StreamProcessor will append if upstream doesn't send it).
|
||||
assert "data: [DONE]" in data_lines
|
||||
text = out.decode("utf-8", errors="replace")
|
||||
data_lines = [ln for ln in text.splitlines() if ln.startswith("data: ")]
|
||||
|
||||
# Parse JSON events (excluding [DONE]) and validate we have expected deltas.
|
||||
events = [json.loads(ln[6:]) for ln in data_lines if ln != "data: [DONE]"]
|
||||
# OpenAI termination marker should be present (StreamProcessor will append if upstream doesn't send it).
|
||||
assert "data: [DONE]" in data_lines
|
||||
|
||||
delta_contents: list[str] = []
|
||||
for evt in events:
|
||||
for choice in evt.get("choices", []) or []:
|
||||
delta = choice.get("delta") or {}
|
||||
if "content" in delta and delta["content"]:
|
||||
delta_contents.append(delta["content"])
|
||||
# Parse JSON events (excluding [DONE]) and validate we have expected deltas.
|
||||
events = [json.loads(ln[6:]) for ln in data_lines if ln != "data: [DONE]"]
|
||||
|
||||
assert "Hello" in "".join(delta_contents)
|
||||
assert " world" in "".join(delta_contents)
|
||||
finally:
|
||||
converter_registry._converters = original_converters
|
||||
delta_contents: list[str] = []
|
||||
for evt in events:
|
||||
for choice in evt.get("choices", []) or []:
|
||||
delta = choice.get("delta") or {}
|
||||
if "content" in delta and delta["content"]:
|
||||
delta_contents.append(delta["content"])
|
||||
|
||||
assert "Hello" in "".join(delta_contents)
|
||||
assert " world" in "".join(delta_contents)
|
||||
|
||||
1
tests/core/api_format/conversion/golden_data/.gitkeep
Normal file
1
tests/core/api_format/conversion/golden_data/.gitkeep
Normal file
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
{
|
||||
"contents": [
|
||||
{
|
||||
"parts": [
|
||||
{
|
||||
"text": "hi"
|
||||
}
|
||||
],
|
||||
"role": "user"
|
||||
}
|
||||
],
|
||||
"generation_config": {
|
||||
"max_output_tokens": 12,
|
||||
"stop_sequences": [
|
||||
"A",
|
||||
"B"
|
||||
],
|
||||
"temperature": 0.2,
|
||||
"top_p": 1.0
|
||||
},
|
||||
"model": "claude-3-5-sonnet-latest",
|
||||
"system_instruction": {
|
||||
"parts": [
|
||||
{
|
||||
"text": "sys\n\ndev"
|
||||
}
|
||||
]
|
||||
},
|
||||
"tool_config": {
|
||||
"function_calling_config": {
|
||||
"mode": "AUTO"
|
||||
}
|
||||
},
|
||||
"tools": [
|
||||
{
|
||||
"function_declarations": [
|
||||
{
|
||||
"description": "Get weather",
|
||||
"name": "get_weather",
|
||||
"parameters": {
|
||||
"properties": {
|
||||
"city": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"city"
|
||||
],
|
||||
"type": "object"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
{
|
||||
"max_tokens": 12,
|
||||
"messages": [
|
||||
{
|
||||
"content": "sys\n\ndev",
|
||||
"role": "system"
|
||||
},
|
||||
{
|
||||
"content": "hi",
|
||||
"role": "user"
|
||||
}
|
||||
],
|
||||
"model": "claude-3-5-sonnet-latest",
|
||||
"stop": [
|
||||
"A",
|
||||
"B"
|
||||
],
|
||||
"stream": true,
|
||||
"temperature": 0.2,
|
||||
"tool_choice": "auto",
|
||||
"tools": [
|
||||
{
|
||||
"function": {
|
||||
"description": "Get weather",
|
||||
"name": "get_weather",
|
||||
"parameters": {
|
||||
"properties": {
|
||||
"city": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"city"
|
||||
],
|
||||
"type": "object"
|
||||
}
|
||||
},
|
||||
"type": "function"
|
||||
}
|
||||
],
|
||||
"top_p": 1.0
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
{
|
||||
"max_tokens": 12,
|
||||
"messages": [
|
||||
{
|
||||
"content": "hi",
|
||||
"role": "user"
|
||||
}
|
||||
],
|
||||
"model": "gemini-1.5-flash",
|
||||
"stop_sequences": [
|
||||
"A",
|
||||
"B"
|
||||
],
|
||||
"system": "sys\n\ndev",
|
||||
"temperature": 0.2,
|
||||
"tool_choice": {
|
||||
"type": "auto"
|
||||
},
|
||||
"tools": [
|
||||
{
|
||||
"description": "Get weather",
|
||||
"input_schema": {
|
||||
"properties": {
|
||||
"city": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"city"
|
||||
],
|
||||
"type": "object"
|
||||
},
|
||||
"name": "get_weather"
|
||||
}
|
||||
],
|
||||
"top_p": 1.0
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
{
|
||||
"max_tokens": 12,
|
||||
"messages": [
|
||||
{
|
||||
"content": "sys\n\ndev",
|
||||
"role": "system"
|
||||
},
|
||||
{
|
||||
"content": "hi",
|
||||
"role": "user"
|
||||
}
|
||||
],
|
||||
"model": "gemini-1.5-flash",
|
||||
"stop": [
|
||||
"A",
|
||||
"B"
|
||||
],
|
||||
"temperature": 0.2,
|
||||
"tool_choice": "auto",
|
||||
"tools": [
|
||||
{
|
||||
"function": {
|
||||
"description": "Get weather",
|
||||
"name": "get_weather",
|
||||
"parameters": {
|
||||
"properties": {
|
||||
"city": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"city"
|
||||
],
|
||||
"type": "object"
|
||||
}
|
||||
},
|
||||
"type": "function"
|
||||
}
|
||||
],
|
||||
"top_p": 1.0
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
{
|
||||
"max_tokens": 12,
|
||||
"messages": [
|
||||
{
|
||||
"content": "hi",
|
||||
"role": "user"
|
||||
}
|
||||
],
|
||||
"model": "gpt-4o-mini",
|
||||
"stop_sequences": [
|
||||
"A",
|
||||
"B"
|
||||
],
|
||||
"stream": true,
|
||||
"system": "sys\n\ndev",
|
||||
"temperature": 0.2,
|
||||
"tool_choice": {
|
||||
"type": "auto"
|
||||
},
|
||||
"tools": [
|
||||
{
|
||||
"description": "Get weather",
|
||||
"input_schema": {
|
||||
"properties": {
|
||||
"city": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"city"
|
||||
],
|
||||
"type": "object"
|
||||
},
|
||||
"name": "get_weather"
|
||||
}
|
||||
],
|
||||
"top_p": 1.0
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
{
|
||||
"contents": [
|
||||
{
|
||||
"parts": [
|
||||
{
|
||||
"text": "hi"
|
||||
}
|
||||
],
|
||||
"role": "user"
|
||||
}
|
||||
],
|
||||
"generation_config": {
|
||||
"max_output_tokens": 12,
|
||||
"stop_sequences": [
|
||||
"A",
|
||||
"B"
|
||||
],
|
||||
"temperature": 0.2,
|
||||
"top_p": 1.0
|
||||
},
|
||||
"model": "gpt-4o-mini",
|
||||
"system_instruction": {
|
||||
"parts": [
|
||||
{
|
||||
"text": "sys\n\ndev"
|
||||
}
|
||||
]
|
||||
},
|
||||
"tool_config": {
|
||||
"function_calling_config": {
|
||||
"mode": "AUTO"
|
||||
}
|
||||
},
|
||||
"tools": [
|
||||
{
|
||||
"function_declarations": [
|
||||
{
|
||||
"description": "Get weather",
|
||||
"name": "get_weather",
|
||||
"parameters": {
|
||||
"properties": {
|
||||
"city": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"city"
|
||||
],
|
||||
"type": "object"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [
|
||||
{
|
||||
"text": "hello"
|
||||
}
|
||||
],
|
||||
"role": "model"
|
||||
},
|
||||
"finishReason": "STOP",
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"id": "msg_1",
|
||||
"modelVersion": "claude-3-5-sonnet-latest",
|
||||
"usageMetadata": {
|
||||
"candidatesTokenCount": 7,
|
||||
"promptTokenCount": 5,
|
||||
"totalTokenCount": 12
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
"message": {
|
||||
"content": "hello",
|
||||
"role": "assistant"
|
||||
}
|
||||
}
|
||||
],
|
||||
"id": "msg_1",
|
||||
"model": "claude-3-5-sonnet-latest",
|
||||
"object": "chat.completion",
|
||||
"usage": {
|
||||
"completion_tokens": 7,
|
||||
"prompt_tokens": 5,
|
||||
"total_tokens": 12
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
{
|
||||
"content": [
|
||||
{
|
||||
"text": "hello",
|
||||
"type": "text"
|
||||
}
|
||||
],
|
||||
"id": "msg_unknown",
|
||||
"model": "gemini-1.5-flash",
|
||||
"role": "assistant",
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": null,
|
||||
"type": "message",
|
||||
"usage": {
|
||||
"input_tokens": 5,
|
||||
"output_tokens": 7
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
"message": {
|
||||
"content": "hello",
|
||||
"role": "assistant"
|
||||
}
|
||||
}
|
||||
],
|
||||
"id": "chatcmpl-unknown",
|
||||
"model": "gemini-1.5-flash",
|
||||
"object": "chat.completion",
|
||||
"usage": {
|
||||
"completion_tokens": 7,
|
||||
"prompt_tokens": 5,
|
||||
"total_tokens": 12
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
{
|
||||
"content": [
|
||||
{
|
||||
"text": "hello",
|
||||
"type": "text"
|
||||
}
|
||||
],
|
||||
"id": "msg_chatcmpl_1",
|
||||
"model": "gpt-4o-mini",
|
||||
"role": "assistant",
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": null,
|
||||
"type": "message",
|
||||
"usage": {
|
||||
"input_tokens": 5,
|
||||
"output_tokens": 7
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [
|
||||
{
|
||||
"text": "hello"
|
||||
}
|
||||
],
|
||||
"role": "model"
|
||||
},
|
||||
"finishReason": "STOP",
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"id": "chatcmpl_1",
|
||||
"modelVersion": "gpt-4o-mini",
|
||||
"usageMetadata": {
|
||||
"candidatesTokenCount": 7,
|
||||
"promptTokenCount": 5,
|
||||
"totalTokenCount": 12
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
[
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [
|
||||
{
|
||||
"text": "hello"
|
||||
}
|
||||
],
|
||||
"role": "model"
|
||||
},
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"modelVersion": "claude-3-5-sonnet-latest"
|
||||
},
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [],
|
||||
"role": "model"
|
||||
},
|
||||
"finishReason": "STOP",
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"modelVersion": "claude-3-5-sonnet-latest",
|
||||
"usageMetadata": {
|
||||
"candidatesTokenCount": 7,
|
||||
"promptTokenCount": 5,
|
||||
"totalTokenCount": 12
|
||||
}
|
||||
}
|
||||
]
|
||||
@@ -0,0 +1,42 @@
|
||||
[
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"delta": {
|
||||
"role": "assistant"
|
||||
},
|
||||
"finish_reason": null,
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"id": "msg_1",
|
||||
"model": "claude-3-5-sonnet-latest",
|
||||
"object": "chat.completion.chunk"
|
||||
},
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"delta": {
|
||||
"content": "hello"
|
||||
},
|
||||
"finish_reason": null,
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"id": "msg_1",
|
||||
"model": "claude-3-5-sonnet-latest",
|
||||
"object": "chat.completion.chunk"
|
||||
},
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"delta": {},
|
||||
"finish_reason": "stop",
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"id": "msg_1",
|
||||
"model": "claude-3-5-sonnet-latest",
|
||||
"object": "chat.completion.chunk"
|
||||
}
|
||||
]
|
||||
@@ -0,0 +1,59 @@
|
||||
[
|
||||
{
|
||||
"message": {
|
||||
"content": [],
|
||||
"id": "gemini_1",
|
||||
"model": "gemini-1.5-flash",
|
||||
"role": "assistant",
|
||||
"stop_reason": null,
|
||||
"stop_sequence": null,
|
||||
"type": "message",
|
||||
"usage": {
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0
|
||||
}
|
||||
},
|
||||
"type": "message_start"
|
||||
},
|
||||
{
|
||||
"content_block": {
|
||||
"text": "",
|
||||
"type": "text"
|
||||
},
|
||||
"index": 0,
|
||||
"type": "content_block_start"
|
||||
},
|
||||
{
|
||||
"delta": {
|
||||
"text": "he",
|
||||
"type": "text_delta"
|
||||
},
|
||||
"index": 0,
|
||||
"type": "content_block_delta"
|
||||
},
|
||||
{
|
||||
"delta": {
|
||||
"text": "llo",
|
||||
"type": "text_delta"
|
||||
},
|
||||
"index": 0,
|
||||
"type": "content_block_delta"
|
||||
},
|
||||
{
|
||||
"index": 0,
|
||||
"type": "content_block_stop"
|
||||
},
|
||||
{
|
||||
"delta": {
|
||||
"stop_reason": "end_turn"
|
||||
},
|
||||
"type": "message_delta",
|
||||
"usage": {
|
||||
"input_tokens": 5,
|
||||
"output_tokens": 7
|
||||
}
|
||||
},
|
||||
{
|
||||
"type": "message_stop"
|
||||
}
|
||||
]
|
||||
@@ -0,0 +1,56 @@
|
||||
[
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"delta": {
|
||||
"role": "assistant"
|
||||
},
|
||||
"finish_reason": null,
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"id": "gemini_1",
|
||||
"model": "gemini-1.5-flash",
|
||||
"object": "chat.completion.chunk"
|
||||
},
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"delta": {
|
||||
"content": "he"
|
||||
},
|
||||
"finish_reason": null,
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"id": "gemini_1",
|
||||
"model": "gemini-1.5-flash",
|
||||
"object": "chat.completion.chunk"
|
||||
},
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"delta": {
|
||||
"content": "llo"
|
||||
},
|
||||
"finish_reason": null,
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"id": "gemini_1",
|
||||
"model": "gemini-1.5-flash",
|
||||
"object": "chat.completion.chunk"
|
||||
},
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"delta": {},
|
||||
"finish_reason": "stop",
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"id": "gemini_1",
|
||||
"model": "gemini-1.5-flash",
|
||||
"object": "chat.completion.chunk"
|
||||
}
|
||||
]
|
||||
@@ -0,0 +1,55 @@
|
||||
[
|
||||
{
|
||||
"message": {
|
||||
"content": [],
|
||||
"id": "chatcmpl_1",
|
||||
"model": "gpt-4o-mini",
|
||||
"role": "assistant",
|
||||
"stop_reason": null,
|
||||
"stop_sequence": null,
|
||||
"type": "message",
|
||||
"usage": {
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0
|
||||
}
|
||||
},
|
||||
"type": "message_start"
|
||||
},
|
||||
{
|
||||
"content_block": {
|
||||
"text": "",
|
||||
"type": "text"
|
||||
},
|
||||
"index": 0,
|
||||
"type": "content_block_start"
|
||||
},
|
||||
{
|
||||
"delta": {
|
||||
"text": "He",
|
||||
"type": "text_delta"
|
||||
},
|
||||
"index": 0,
|
||||
"type": "content_block_delta"
|
||||
},
|
||||
{
|
||||
"delta": {
|
||||
"text": "llo",
|
||||
"type": "text_delta"
|
||||
},
|
||||
"index": 0,
|
||||
"type": "content_block_delta"
|
||||
},
|
||||
{
|
||||
"index": 0,
|
||||
"type": "content_block_stop"
|
||||
},
|
||||
{
|
||||
"delta": {
|
||||
"stop_reason": "end_turn"
|
||||
},
|
||||
"type": "message_delta"
|
||||
},
|
||||
{
|
||||
"type": "message_stop"
|
||||
}
|
||||
]
|
||||
@@ -0,0 +1,47 @@
|
||||
[
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [
|
||||
{
|
||||
"text": "He"
|
||||
}
|
||||
],
|
||||
"role": "model"
|
||||
},
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"modelVersion": "gpt-4o-mini"
|
||||
},
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [
|
||||
{
|
||||
"text": "llo"
|
||||
}
|
||||
],
|
||||
"role": "model"
|
||||
},
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"modelVersion": "gpt-4o-mini"
|
||||
},
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [],
|
||||
"role": "model"
|
||||
},
|
||||
"finishReason": "STOP",
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"modelVersion": "gpt-4o-mini"
|
||||
}
|
||||
]
|
||||
@@ -0,0 +1,38 @@
|
||||
{
|
||||
"max_tokens": 12,
|
||||
"messages": [
|
||||
{
|
||||
"content": "hi",
|
||||
"role": "user"
|
||||
}
|
||||
],
|
||||
"model": "claude-3-5-sonnet-latest",
|
||||
"stop_sequences": [
|
||||
"A",
|
||||
"B"
|
||||
],
|
||||
"stream": true,
|
||||
"system": "sys\n\ndev",
|
||||
"temperature": 0.2,
|
||||
"tool_choice": {
|
||||
"type": "auto"
|
||||
},
|
||||
"tools": [
|
||||
{
|
||||
"description": "Get weather",
|
||||
"input_schema": {
|
||||
"properties": {
|
||||
"city": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"city"
|
||||
],
|
||||
"type": "object"
|
||||
},
|
||||
"name": "get_weather"
|
||||
}
|
||||
],
|
||||
"top_p": 1.0
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
{
|
||||
"contents": [
|
||||
{
|
||||
"parts": [
|
||||
{
|
||||
"text": "hi"
|
||||
}
|
||||
],
|
||||
"role": "user"
|
||||
}
|
||||
],
|
||||
"generationConfig": {
|
||||
"maxOutputTokens": 12,
|
||||
"stopSequences": [
|
||||
"A",
|
||||
"B"
|
||||
],
|
||||
"temperature": 0.2,
|
||||
"topP": 1.0
|
||||
},
|
||||
"model": "gemini-1.5-flash",
|
||||
"systemInstruction": {
|
||||
"parts": [
|
||||
{
|
||||
"text": "sys\n\ndev"
|
||||
}
|
||||
]
|
||||
},
|
||||
"toolConfig": {
|
||||
"functionCallingConfig": {
|
||||
"mode": "AUTO"
|
||||
}
|
||||
},
|
||||
"tools": [
|
||||
{
|
||||
"functionDeclarations": [
|
||||
{
|
||||
"description": "Get weather",
|
||||
"name": "get_weather",
|
||||
"parameters": {
|
||||
"properties": {
|
||||
"city": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"city"
|
||||
],
|
||||
"type": "object"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
{
|
||||
"max_tokens": 12,
|
||||
"messages": [
|
||||
{
|
||||
"content": "sys",
|
||||
"role": "system"
|
||||
},
|
||||
{
|
||||
"content": "dev",
|
||||
"role": "developer"
|
||||
},
|
||||
{
|
||||
"content": "hi",
|
||||
"role": "user"
|
||||
}
|
||||
],
|
||||
"model": "gpt-4o-mini",
|
||||
"stop": [
|
||||
"A",
|
||||
"B"
|
||||
],
|
||||
"stream": true,
|
||||
"temperature": 0.2,
|
||||
"tool_choice": "auto",
|
||||
"tools": [
|
||||
{
|
||||
"function": {
|
||||
"description": "Get weather",
|
||||
"name": "get_weather",
|
||||
"parameters": {
|
||||
"properties": {
|
||||
"city": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"city"
|
||||
],
|
||||
"type": "object"
|
||||
}
|
||||
},
|
||||
"type": "function"
|
||||
}
|
||||
],
|
||||
"top_p": 1.0
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
{
|
||||
"content": [
|
||||
{
|
||||
"text": "hello",
|
||||
"type": "text"
|
||||
}
|
||||
],
|
||||
"id": "msg_1",
|
||||
"model": "claude-3-5-sonnet-latest",
|
||||
"role": "assistant",
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": null,
|
||||
"type": "message",
|
||||
"usage": {
|
||||
"input_tokens": 5,
|
||||
"output_tokens": 7
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [
|
||||
{
|
||||
"text": "hello"
|
||||
}
|
||||
],
|
||||
"role": "model"
|
||||
},
|
||||
"finishReason": "STOP",
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"modelVersion": "gemini-1.5-flash",
|
||||
"usageMetadata": {
|
||||
"cachedContentTokenCount": 0,
|
||||
"candidatesTokenCount": 7,
|
||||
"promptTokenCount": 5,
|
||||
"totalTokenCount": 12
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"finish_reason": "stop",
|
||||
"index": 0,
|
||||
"message": {
|
||||
"content": "hello",
|
||||
"role": "assistant"
|
||||
}
|
||||
}
|
||||
],
|
||||
"created": 1,
|
||||
"id": "chatcmpl_1",
|
||||
"model": "gpt-4o-mini",
|
||||
"object": "chat.completion",
|
||||
"usage": {
|
||||
"completion_tokens": 7,
|
||||
"prompt_tokens": 5,
|
||||
"total_tokens": 12
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
[
|
||||
{
|
||||
"message": {
|
||||
"content": [],
|
||||
"id": "msg_1",
|
||||
"model": "claude-3-5-sonnet-latest",
|
||||
"role": "assistant",
|
||||
"stop_reason": null,
|
||||
"stop_sequence": null,
|
||||
"type": "message"
|
||||
},
|
||||
"type": "message_start"
|
||||
},
|
||||
{
|
||||
"content_block": {
|
||||
"text": "",
|
||||
"type": "text"
|
||||
},
|
||||
"index": 0,
|
||||
"type": "content_block_start"
|
||||
},
|
||||
{
|
||||
"delta": {
|
||||
"text": "hello",
|
||||
"type": "text_delta"
|
||||
},
|
||||
"index": 0,
|
||||
"type": "content_block_delta"
|
||||
},
|
||||
{
|
||||
"index": 0,
|
||||
"type": "content_block_stop"
|
||||
},
|
||||
{
|
||||
"delta": {
|
||||
"stop_reason": "end_turn"
|
||||
},
|
||||
"type": "message_delta",
|
||||
"usage": {
|
||||
"input_tokens": 5,
|
||||
"output_tokens": 7
|
||||
}
|
||||
},
|
||||
{
|
||||
"type": "message_stop"
|
||||
}
|
||||
]
|
||||
@@ -0,0 +1,40 @@
|
||||
[
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [
|
||||
{
|
||||
"text": "he"
|
||||
}
|
||||
],
|
||||
"role": "model"
|
||||
},
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"modelVersion": "gemini-1.5-flash"
|
||||
},
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [
|
||||
{
|
||||
"text": "hello"
|
||||
}
|
||||
],
|
||||
"role": "model"
|
||||
},
|
||||
"finishReason": "STOP",
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"modelVersion": "gemini-1.5-flash",
|
||||
"usageMetadata": {
|
||||
"candidatesTokenCount": 7,
|
||||
"promptTokenCount": 5,
|
||||
"totalTokenCount": 12
|
||||
}
|
||||
}
|
||||
]
|
||||
@@ -0,0 +1,33 @@
|
||||
[
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"delta": {
|
||||
"content": "He",
|
||||
"role": "assistant"
|
||||
},
|
||||
"finish_reason": null,
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"created": 1,
|
||||
"id": "chatcmpl_1",
|
||||
"model": "gpt-4o-mini",
|
||||
"object": "chat.completion.chunk"
|
||||
},
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"delta": {
|
||||
"content": "llo"
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
"index": 0
|
||||
}
|
||||
],
|
||||
"created": 1,
|
||||
"id": "chatcmpl_1",
|
||||
"model": "gpt-4o-mini",
|
||||
"object": "chat.completion.chunk"
|
||||
}
|
||||
]
|
||||
352
tests/core/api_format/conversion/test_claude_normalizer.py
Normal file
352
tests/core/api_format/conversion/test_claude_normalizer.py
Normal file
@@ -0,0 +1,352 @@
|
||||
"""
|
||||
ClaudeNormalizer 单元测试
|
||||
|
||||
覆盖重点:
|
||||
- system -> instructions 的提取与还原
|
||||
- tool_use/tool_result 的往返转换
|
||||
- UnknownBlock 内部保留、输出默认丢弃
|
||||
- stop_reason/usage 的映射
|
||||
- streaming event <-> InternalStreamEvent 的基础行为与状态
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List, cast
|
||||
|
||||
from src.core.api_format.conversion.internal import (
|
||||
ErrorType,
|
||||
ImageBlock,
|
||||
StopReason,
|
||||
TextBlock,
|
||||
ToolResultBlock,
|
||||
ToolUseBlock,
|
||||
UnknownBlock,
|
||||
)
|
||||
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer
|
||||
from src.core.api_format.conversion.stream_events import (
|
||||
ContentBlockStartEvent,
|
||||
ContentDeltaEvent,
|
||||
MessageStartEvent,
|
||||
MessageStopEvent,
|
||||
ToolCallDeltaEvent,
|
||||
)
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
def test_claude_request_system_roundtrip() -> None:
|
||||
n = ClaudeNormalizer()
|
||||
|
||||
req = {
|
||||
"model": "claude-3-opus",
|
||||
"system": "sys",
|
||||
"messages": [
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "ok"}]},
|
||||
],
|
||||
"max_tokens": 10,
|
||||
"stop_sequences": ["A"],
|
||||
"stream": True,
|
||||
}
|
||||
|
||||
internal = n.request_to_internal(req)
|
||||
assert internal.model == "claude-3-opus"
|
||||
assert [seg.role.value for seg in internal.instructions] == ["system"]
|
||||
assert internal.instructions[0].text == "sys"
|
||||
assert internal.system == "sys"
|
||||
assert [m.role.value for m in internal.messages] == ["user", "assistant"]
|
||||
assert internal.max_tokens == 10
|
||||
assert internal.stop_sequences == ["A"]
|
||||
assert internal.stream is True
|
||||
|
||||
out = n.request_from_internal(internal)
|
||||
assert out["model"] == "claude-3-opus"
|
||||
assert out["system"] == "sys"
|
||||
|
||||
out_messages: List[Dict[str, Any]] = out["messages"]
|
||||
assert [m["role"] for m in out_messages] == ["user", "assistant"]
|
||||
assert out_messages[0]["content"] == "hi"
|
||||
assert out_messages[1]["content"] == "ok"
|
||||
|
||||
|
||||
def test_claude_request_tool_blocks_roundtrip() -> None:
|
||||
n = ClaudeNormalizer()
|
||||
|
||||
req = {
|
||||
"model": "claude-3-sonnet",
|
||||
"system": "sys",
|
||||
"messages": [
|
||||
{"role": "user", "content": "weather?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "toolu_1",
|
||||
"name": "get_weather",
|
||||
"input": {"city": "SF"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "toolu_1",
|
||||
"content": {"temp_c": 20},
|
||||
"is_error": False,
|
||||
}
|
||||
],
|
||||
},
|
||||
],
|
||||
"tools": [
|
||||
{
|
||||
"name": "get_weather",
|
||||
"description": "Get weather",
|
||||
"input_schema": {"type": "object", "properties": {"city": {"type": "string"}}},
|
||||
}
|
||||
],
|
||||
"tool_choice": {"type": "any"},
|
||||
"max_tokens": 10,
|
||||
}
|
||||
|
||||
internal = n.request_to_internal(req)
|
||||
assert [m.role.value for m in internal.messages] == ["user", "assistant", "user"]
|
||||
|
||||
assistant_msg = internal.messages[1]
|
||||
assert any(isinstance(b, ToolUseBlock) for b in assistant_msg.content)
|
||||
tool_use = next(b for b in assistant_msg.content if isinstance(b, ToolUseBlock))
|
||||
assert tool_use.tool_id == "toolu_1"
|
||||
assert tool_use.tool_name == "get_weather"
|
||||
assert tool_use.tool_input == {"city": "SF"}
|
||||
|
||||
tool_result_msg = internal.messages[2]
|
||||
assert any(isinstance(b, ToolResultBlock) for b in tool_result_msg.content)
|
||||
tool_result = next(b for b in tool_result_msg.content if isinstance(b, ToolResultBlock))
|
||||
assert tool_result.tool_use_id == "toolu_1"
|
||||
assert tool_result.output == {"temp_c": 20}
|
||||
assert tool_result.is_error is False
|
||||
|
||||
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_out = out_messages[1]
|
||||
assert isinstance(assistant_out["content"], list)
|
||||
a_blocks = cast(List[Dict[str, Any]], assistant_out["content"])
|
||||
assert a_blocks[0]["type"] == "tool_use"
|
||||
assert a_blocks[0]["id"] == "toolu_1"
|
||||
assert a_blocks[0]["name"] == "get_weather"
|
||||
|
||||
user_out = out_messages[2]
|
||||
assert isinstance(user_out["content"], list)
|
||||
u_blocks = cast(List[Dict[str, Any]], user_out["content"])
|
||||
assert u_blocks[0]["type"] == "tool_result"
|
||||
assert u_blocks[0]["tool_use_id"] == "toolu_1"
|
||||
assert u_blocks[0]["content"] == {"temp_c": 20}
|
||||
|
||||
|
||||
def test_claude_unknown_block_drop_on_output() -> None:
|
||||
n = ClaudeNormalizer()
|
||||
|
||||
req = {
|
||||
"model": "claude-3-sonnet",
|
||||
"messages": [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "text", "text": "ok"},
|
||||
{"type": "thinking", "text": "secret"},
|
||||
],
|
||||
}
|
||||
],
|
||||
"max_tokens": 10,
|
||||
}
|
||||
|
||||
internal = n.request_to_internal(req)
|
||||
assert len(internal.messages) == 1
|
||||
blocks = internal.messages[0].content
|
||||
assert any(isinstance(b, TextBlock) for b in blocks)
|
||||
assert any(isinstance(b, UnknownBlock) for b in blocks)
|
||||
u = next(b for b in blocks if isinstance(b, UnknownBlock))
|
||||
assert u.raw_type == "thinking"
|
||||
|
||||
out = n.request_from_internal(internal)
|
||||
# Claude 要求以 user 开头,且会做最小修复:插入空 user
|
||||
out_messages: List[Dict[str, Any]] = out["messages"]
|
||||
assert out_messages[0]["role"] == "user"
|
||||
assert out_messages[1]["role"] == "assistant"
|
||||
assert out_messages[1]["content"] == "ok"
|
||||
|
||||
|
||||
def test_claude_response_stop_reason_and_usage_roundtrip() -> None:
|
||||
n = ClaudeNormalizer()
|
||||
|
||||
resp = {
|
||||
"id": "msg_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-3-sonnet",
|
||||
"content": [{"type": "text", "text": "hi"}],
|
||||
"stop_reason": "tool_use",
|
||||
"usage": {
|
||||
"input_tokens": 5,
|
||||
"output_tokens": 7,
|
||||
"cache_read_input_tokens": 1,
|
||||
"cache_creation_input_tokens": 2,
|
||||
},
|
||||
}
|
||||
|
||||
internal = n.response_to_internal(resp)
|
||||
assert internal.id == "msg_1"
|
||||
assert internal.stop_reason == StopReason.TOOL_USE
|
||||
assert internal.usage is not None
|
||||
assert internal.usage.input_tokens == 5
|
||||
assert internal.usage.output_tokens == 7
|
||||
assert internal.usage.total_tokens == 12
|
||||
assert internal.usage.cache_read_tokens == 1
|
||||
assert internal.usage.cache_write_tokens == 2
|
||||
|
||||
out = n.response_from_internal(internal)
|
||||
assert out["id"] == "msg_1"
|
||||
assert out["stop_reason"] == "tool_use"
|
||||
assert out["usage"]["cache_read_input_tokens"] == 1
|
||||
assert out["usage"]["cache_creation_input_tokens"] == 2
|
||||
|
||||
|
||||
def test_claude_stream_chunk_and_event_roundtrip_basic() -> None:
|
||||
n = ClaudeNormalizer()
|
||||
state = StreamState()
|
||||
|
||||
chunks = [
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-3-sonnet",
|
||||
"content": [],
|
||||
"stop_reason": None,
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 1, "output_tokens": 0},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "Hel"}},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "lo"}},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 1,
|
||||
"content_block": {"type": "tool_use", "id": "toolu_1", "name": "get_weather"},
|
||||
},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 1,
|
||||
"delta": {"type": "input_json_delta", "partial_json": '{"city":"SF"}'},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
{"type": "content_block_stop", "index": 1},
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"input_tokens": 1, "output_tokens": 2},
|
||||
},
|
||||
{"type": "message_stop"},
|
||||
]
|
||||
|
||||
events: List[Any] = []
|
||||
for ch in chunks:
|
||||
events.extend(n.stream_chunk_to_internal(ch, state))
|
||||
|
||||
assert any(isinstance(e, MessageStartEvent) for e in events)
|
||||
assert [e.text_delta for e in events if isinstance(e, ContentDeltaEvent)] == ["Hel", "lo"]
|
||||
assert any(isinstance(e, ToolCallDeltaEvent) and e.tool_id == "toolu_1" for e in events)
|
||||
assert any(isinstance(e, MessageStopEvent) and e.stop_reason == StopReason.END_TURN for e in events)
|
||||
|
||||
# internal events -> Claude events
|
||||
state2 = StreamState()
|
||||
out_events: List[Dict[str, Any]] = []
|
||||
for e in events:
|
||||
out_events.extend(n.stream_event_from_internal(e, state2))
|
||||
|
||||
assert out_events[0]["type"] == "message_start"
|
||||
assert out_events[0]["message"]["id"] == "msg_1"
|
||||
|
||||
assert any(ev.get("type") == "content_block_delta" and ev.get("delta", {}).get("type") == "input_json_delta" for ev in out_events)
|
||||
assert out_events[-1]["type"] == "message_stop"
|
||||
|
||||
|
||||
def test_claude_error_conversion() -> None:
|
||||
n = ClaudeNormalizer()
|
||||
err_resp = {
|
||||
"type": "error",
|
||||
"error": {"type": "rate_limit_error", "message": "slow down"},
|
||||
}
|
||||
|
||||
assert n.is_error_response(err_resp) is True
|
||||
internal = n.error_to_internal(err_resp)
|
||||
assert internal.type == ErrorType.RATE_LIMIT
|
||||
assert internal.retryable is True
|
||||
|
||||
out = n.error_from_internal(internal)
|
||||
assert out["type"] == "error"
|
||||
assert out["error"]["type"] == "rate_limit_error"
|
||||
|
||||
|
||||
def test_claude_request_metadata_preserved() -> None:
|
||||
"""测试 Claude 请求中 metadata 字段的保留"""
|
||||
n = ClaudeNormalizer()
|
||||
|
||||
req = {
|
||||
"model": "claude-haiku-4-5-20251001",
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}]}],
|
||||
"system": [
|
||||
{"type": "text", "text": "System prompt 1"},
|
||||
{"type": "text", "text": "System prompt 2"},
|
||||
],
|
||||
"metadata": {
|
||||
"user_id": "user_abc123_session_xyz456"
|
||||
},
|
||||
"max_tokens": 32000,
|
||||
"stream": True,
|
||||
}
|
||||
|
||||
internal = n.request_to_internal(req)
|
||||
|
||||
# system 数组应该被正确合并
|
||||
assert internal.system == "System prompt 1\n\nSystem prompt 2"
|
||||
|
||||
# metadata 应该在 extra 中
|
||||
assert "claude" in internal.extra
|
||||
assert "metadata" in internal.extra["claude"]
|
||||
assert internal.extra["claude"]["metadata"]["user_id"] == "user_abc123_session_xyz456"
|
||||
|
||||
# 往返转换后 metadata 应该被恢复
|
||||
out = n.request_from_internal(internal)
|
||||
assert "metadata" in out
|
||||
assert out["metadata"]["user_id"] == "user_abc123_session_xyz456"
|
||||
|
||||
|
||||
def test_claude_system_array_format() -> None:
|
||||
"""测试 Claude CLI 风格的 system 数组格式"""
|
||||
n = ClaudeNormalizer()
|
||||
|
||||
req = {
|
||||
"model": "claude-3-opus",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"system": [
|
||||
{"type": "text", "text": "x-anthropic-billing-header: cc_version=2.1.19"},
|
||||
{"type": "text", "text": "You are Claude Code."},
|
||||
{"type": "text", "text": "Extract file paths."},
|
||||
],
|
||||
"max_tokens": 4096,
|
||||
}
|
||||
|
||||
internal = n.request_to_internal(req)
|
||||
|
||||
# system 数组中的多个 text 应该用 \n\n 连接
|
||||
assert "x-anthropic-billing-header" in internal.system
|
||||
assert "You are Claude Code" in internal.system
|
||||
assert "Extract file paths" in internal.system
|
||||
assert "\n\n" in internal.system
|
||||
577
tests/core/api_format/conversion/test_cli_conversion.py
Normal file
577
tests/core/api_format/conversion/test_cli_conversion.py
Normal file
@@ -0,0 +1,577 @@
|
||||
"""
|
||||
CLI 格式参与转换的单元测试
|
||||
|
||||
覆盖:
|
||||
- OPENAI_CLI(Responses)与其他格式的 request/response/stream 基础互转
|
||||
- CLAUDE_CLI / GEMINI_CLI 的 registry 注册与互转能力
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List, cast
|
||||
|
||||
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer
|
||||
from src.core.api_format.conversion.normalizers.claude_cli import ClaudeCliNormalizer
|
||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||
from src.core.api_format.conversion.normalizers.gemini_cli import GeminiCliNormalizer
|
||||
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||
from src.core.api_format.conversion.normalizers.openai_cli import OpenAICliNormalizer
|
||||
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
def _make_registry_with_cli() -> FormatConversionRegistry:
|
||||
reg = FormatConversionRegistry()
|
||||
reg.register(OpenAINormalizer())
|
||||
reg.register(OpenAICliNormalizer())
|
||||
reg.register(ClaudeNormalizer())
|
||||
reg.register(ClaudeCliNormalizer())
|
||||
reg.register(GeminiNormalizer())
|
||||
reg.register(GeminiCliNormalizer())
|
||||
return reg
|
||||
|
||||
|
||||
def test_registry_can_convert_full_with_cli_stream() -> None:
|
||||
reg = _make_registry_with_cli()
|
||||
assert reg.can_convert_full("OPENAI_CLI", "OPENAI", require_stream=True) is True
|
||||
assert reg.can_convert_full("OPENAI_CLI", "CLAUDE_CLI", require_stream=True) is True
|
||||
assert reg.can_convert_full("GEMINI_CLI", "CLAUDE", require_stream=True) is True
|
||||
|
||||
|
||||
def test_openai_cli_request_to_claude() -> None:
|
||||
reg = _make_registry_with_cli()
|
||||
|
||||
openai_cli_req = {
|
||||
"model": "gpt-4o-mini",
|
||||
"input": [{"role": "user", "content": [{"type": "input_text", "text": "hi"}]}],
|
||||
"stream": True,
|
||||
"max_output_tokens": 12,
|
||||
}
|
||||
|
||||
claude_req = reg.convert_request(openai_cli_req, "OPENAI_CLI", "CLAUDE")
|
||||
assert claude_req["model"] == "gpt-4o-mini"
|
||||
assert claude_req["stream"] is True
|
||||
assert isinstance(claude_req.get("messages"), list)
|
||||
assert claude_req["messages"][0]["role"] == "user"
|
||||
assert claude_req["messages"][0]["content"] == "hi"
|
||||
|
||||
|
||||
def test_claude_response_to_openai_cli() -> None:
|
||||
reg = _make_registry_with_cli()
|
||||
|
||||
claude_resp = {
|
||||
"id": "msg_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-3-5-sonnet-latest",
|
||||
"content": [{"type": "text", "text": "hello"}],
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 5, "output_tokens": 7},
|
||||
}
|
||||
|
||||
openai_cli_resp = reg.convert_response(claude_resp, "CLAUDE", "OPENAI_CLI")
|
||||
assert openai_cli_resp["object"] == "response"
|
||||
assert isinstance(openai_cli_resp.get("output"), list)
|
||||
msg = cast(Dict[str, Any], openai_cli_resp["output"][0])
|
||||
assert msg["type"] == "message"
|
||||
assert msg["role"] == "assistant"
|
||||
content = cast(List[Dict[str, Any]], msg.get("content") or [])
|
||||
assert content and content[0]["type"] == "output_text"
|
||||
assert content[0]["text"] == "hello"
|
||||
|
||||
|
||||
def test_stream_openai_to_openai_cli_delta() -> None:
|
||||
reg = _make_registry_with_cli()
|
||||
state = StreamState()
|
||||
|
||||
chunk = {
|
||||
"id": "chatcmpl_1",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [{"index": 0, "delta": {"role": "assistant", "content": "hi"}, "finish_reason": None}],
|
||||
}
|
||||
|
||||
out_events = reg.convert_stream_chunk(chunk, "OPENAI", "OPENAI_CLI", state=state)
|
||||
assert isinstance(out_events, list) and out_events
|
||||
assert out_events[0].get("type") == "response.created"
|
||||
assert out_events[1].get("type") == "response.output_text.delta"
|
||||
assert out_events[1].get("delta") == "hi"
|
||||
|
||||
|
||||
def test_stream_openai_cli_to_openai_delta() -> None:
|
||||
reg = _make_registry_with_cli()
|
||||
state = StreamState()
|
||||
|
||||
chunk = {
|
||||
"type": "response.output_text.delta",
|
||||
"delta": "hi",
|
||||
"response": {"id": "resp_1", "model": "gpt-4o-mini"},
|
||||
}
|
||||
|
||||
out_events = reg.convert_stream_chunk(chunk, "OPENAI_CLI", "OPENAI", state=state)
|
||||
assert isinstance(out_events, list) and out_events
|
||||
|
||||
# 第一个 chunk 先补齐 assistant role
|
||||
assert out_events[0].get("object") == "chat.completion.chunk"
|
||||
# 第二个 chunk 才是文本增量
|
||||
choices = out_events[1].get("choices") or []
|
||||
assert isinstance(choices, list) and choices
|
||||
delta = cast(Dict[str, Any], choices[0]).get("delta") or {}
|
||||
assert cast(Dict[str, Any], delta).get("content") == "hi"
|
||||
|
||||
|
||||
def test_openai_cli_function_call_to_claude() -> None:
|
||||
"""测试 OpenAI CLI 的 function_call/function_call_output 转换为 Claude tool_use/tool_result"""
|
||||
reg = _make_registry_with_cli()
|
||||
|
||||
openai_cli_req = {
|
||||
"model": "gpt-5",
|
||||
"instructions": "You are helpful.",
|
||||
"input": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": "列出当前目录"}],
|
||||
},
|
||||
{
|
||||
"type": "function_call",
|
||||
"name": "shell_command",
|
||||
"arguments": '{"command": "ls -la"}',
|
||||
"call_id": "call_abc123",
|
||||
},
|
||||
{
|
||||
"type": "function_call_output",
|
||||
"call_id": "call_abc123",
|
||||
"output": "file1.txt\nfile2.txt",
|
||||
},
|
||||
],
|
||||
"stream": True,
|
||||
}
|
||||
|
||||
claude_req = reg.convert_request(openai_cli_req, "OPENAI_CLI", "CLAUDE")
|
||||
|
||||
messages = claude_req.get("messages", [])
|
||||
assert len(messages) == 3
|
||||
|
||||
# 第一条:user 消息
|
||||
assert messages[0]["role"] == "user"
|
||||
assert messages[0]["content"] == "列出当前目录"
|
||||
|
||||
# 第二条:assistant + tool_use
|
||||
assert messages[1]["role"] == "assistant"
|
||||
content1 = messages[1]["content"]
|
||||
assert isinstance(content1, list) and len(content1) == 1
|
||||
assert content1[0]["type"] == "tool_use"
|
||||
assert content1[0]["name"] == "shell_command"
|
||||
assert content1[0]["id"] == "call_abc123"
|
||||
assert content1[0]["input"] == {"command": "ls -la"}
|
||||
|
||||
# 第三条:user + tool_result
|
||||
assert messages[2]["role"] == "user"
|
||||
content2 = messages[2]["content"]
|
||||
assert isinstance(content2, list) and len(content2) == 1
|
||||
assert content2[0]["type"] == "tool_result"
|
||||
assert content2[0]["tool_use_id"] == "call_abc123"
|
||||
assert content2[0]["content"] == "file1.txt\nfile2.txt"
|
||||
|
||||
|
||||
def test_openai_cli_reasoning_preserved_in_roundtrip() -> None:
|
||||
"""测试 OpenAI CLI 的 reasoning block 在 roundtrip 中被保留"""
|
||||
reg = _make_registry_with_cli()
|
||||
|
||||
openai_cli_req = {
|
||||
"model": "gpt-5",
|
||||
"input": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": "思考一下"}],
|
||||
},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"summary": [{"type": "summary_text", "text": "I am thinking about the problem..."}],
|
||||
"content": None,
|
||||
"encrypted_content": "xxx_encrypted",
|
||||
},
|
||||
{
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "我想好了"}],
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
# 转换到 internal 再转回 OPENAI_CLI
|
||||
converted = reg.convert_request(openai_cli_req, "OPENAI_CLI", "OPENAI_CLI")
|
||||
|
||||
input_items = converted.get("input", [])
|
||||
# 应该有 user message, reasoning, assistant message
|
||||
assert len(input_items) >= 2
|
||||
|
||||
# 找到 reasoning block
|
||||
reasoning_items = [i for i in input_items if isinstance(i, dict) and i.get("type") == "reasoning"]
|
||||
assert len(reasoning_items) == 1
|
||||
assert "summary" in reasoning_items[0]
|
||||
|
||||
|
||||
def test_claude_tool_use_to_openai_cli() -> None:
|
||||
"""测试 Claude tool_use/tool_result 转换为 OpenAI CLI function_call/function_call_output"""
|
||||
reg = _make_registry_with_cli()
|
||||
|
||||
claude_req = {
|
||||
"model": "claude-3",
|
||||
"messages": [
|
||||
{"role": "user", "content": "查看文件"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "tool_123",
|
||||
"name": "read_file",
|
||||
"input": {"path": "/tmp/test.txt"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"tool_use_id": "tool_123",
|
||||
"content": "Hello World",
|
||||
}
|
||||
],
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
openai_cli_req = reg.convert_request(claude_req, "CLAUDE", "OPENAI_CLI")
|
||||
|
||||
input_items = openai_cli_req.get("input", [])
|
||||
assert len(input_items) >= 3
|
||||
|
||||
# 找到 function_call
|
||||
fc_items = [i for i in input_items if isinstance(i, dict) and i.get("type") == "function_call"]
|
||||
assert len(fc_items) == 1
|
||||
assert fc_items[0]["name"] == "read_file"
|
||||
assert fc_items[0]["call_id"] == "tool_123"
|
||||
|
||||
# 找到 function_call_output
|
||||
fco_items = [i for i in input_items if isinstance(i, dict) and i.get("type") == "function_call_output"]
|
||||
assert len(fco_items) == 1
|
||||
assert fco_items[0]["call_id"] == "tool_123"
|
||||
assert fco_items[0]["output"] == "Hello World"
|
||||
|
||||
|
||||
def test_stream_openai_cli_in_progress_event() -> None:
|
||||
"""测试 OpenAI CLI 流式 response.in_progress 事件"""
|
||||
reg = _make_registry_with_cli()
|
||||
state = StreamState()
|
||||
|
||||
# response.created 事件
|
||||
created_chunk = {
|
||||
"type": "response.created",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"object": "response",
|
||||
"model": "gpt-5",
|
||||
"status": "in_progress",
|
||||
},
|
||||
}
|
||||
|
||||
events1 = reg.convert_stream_chunk(created_chunk, "OPENAI_CLI", "CLAUDE", state=state)
|
||||
assert isinstance(events1, list) and events1
|
||||
assert events1[0].get("type") == "message_start"
|
||||
|
||||
# response.in_progress 事件(不应产生内容事件)
|
||||
in_progress_chunk = {
|
||||
"type": "response.in_progress",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"object": "response",
|
||||
"model": "gpt-5",
|
||||
"status": "in_progress",
|
||||
},
|
||||
}
|
||||
|
||||
events2 = reg.convert_stream_chunk(in_progress_chunk, "OPENAI_CLI", "CLAUDE", state=state)
|
||||
# response.in_progress 不应产生任何事件
|
||||
assert events2 == []
|
||||
|
||||
|
||||
def test_stream_openai_cli_function_call_events() -> None:
|
||||
"""测试 OpenAI CLI 流式 function_call 相关事件"""
|
||||
reg = _make_registry_with_cli()
|
||||
state = StreamState()
|
||||
|
||||
# 首先发送 response.created
|
||||
created_chunk = {
|
||||
"type": "response.created",
|
||||
"response": {"id": "resp_456", "model": "gpt-5"},
|
||||
}
|
||||
reg.convert_stream_chunk(created_chunk, "OPENAI_CLI", "CLAUDE", state=state)
|
||||
|
||||
# response.output_item.added (function_call)
|
||||
output_item_chunk = {
|
||||
"type": "response.output_item.added",
|
||||
"item": {
|
||||
"type": "function_call",
|
||||
"call_id": "call_xyz",
|
||||
"name": "get_weather",
|
||||
},
|
||||
}
|
||||
|
||||
events1 = reg.convert_stream_chunk(output_item_chunk, "OPENAI_CLI", "CLAUDE", state=state)
|
||||
assert isinstance(events1, list) and events1
|
||||
assert events1[0].get("type") == "content_block_start"
|
||||
|
||||
# response.function_call_arguments.delta
|
||||
args_delta_chunk = {
|
||||
"type": "response.function_call_arguments.delta",
|
||||
"delta": '{"city":',
|
||||
}
|
||||
|
||||
events2 = reg.convert_stream_chunk(args_delta_chunk, "OPENAI_CLI", "CLAUDE", state=state)
|
||||
assert isinstance(events2, list) and events2
|
||||
# ToolCallDeltaEvent 转换为 Claude 的 content_block_delta
|
||||
assert events2[0].get("type") == "content_block_delta"
|
||||
delta_obj = events2[0].get("delta", {})
|
||||
assert delta_obj.get("type") == "input_json_delta"
|
||||
assert delta_obj.get("partial_json") == '{"city":'
|
||||
|
||||
# response.output_item.done (function_call)
|
||||
output_done_chunk = {
|
||||
"type": "response.output_item.done",
|
||||
"item": {
|
||||
"type": "function_call",
|
||||
"call_id": "call_xyz",
|
||||
"name": "get_weather",
|
||||
"arguments": '{"city": "Beijing"}',
|
||||
},
|
||||
}
|
||||
|
||||
events3 = reg.convert_stream_chunk(output_done_chunk, "OPENAI_CLI", "CLAUDE", state=state)
|
||||
assert isinstance(events3, list) and events3
|
||||
assert events3[0].get("type") == "content_block_stop"
|
||||
|
||||
|
||||
def test_real_claude_cli_stream_response_conversion() -> None:
|
||||
"""测试真实的 Claude CLI 流式响应转换(完整事件序列)
|
||||
|
||||
使用来自 Claude Code 的真实流式响应数据,验证:
|
||||
- message_start, content_block_start, ping, content_block_delta,
|
||||
content_block_stop, message_delta, message_stop 的完整处理链路
|
||||
- 文本增量正确拼接
|
||||
- usage 和 stop_reason 正确提取
|
||||
"""
|
||||
reg = _make_registry_with_cli()
|
||||
state = StreamState()
|
||||
|
||||
# 真实的 Claude CLI 流式响应事件序列
|
||||
chunks = [
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"model": "claude-opus-4-5-20251101",
|
||||
"id": "msg_01JEmQ1u53gZndRGBUBLvVZ9",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"stop_reason": None,
|
||||
"stop_sequence": None,
|
||||
"usage": {
|
||||
"input_tokens": 3,
|
||||
"cache_creation_input_tokens": 32760,
|
||||
"cache_read_input_tokens": 61777,
|
||||
"cache_creation": {
|
||||
"ephemeral_5m_input_tokens": 32760,
|
||||
"ephemeral_1h_input_tokens": 0,
|
||||
},
|
||||
"output_tokens": 2,
|
||||
"service_tier": "standard",
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
},
|
||||
{"type": "ping"},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "明"}},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "白"}},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "了,请"},
|
||||
},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "把"}},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "包"}},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "含 "}},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "tools"},
|
||||
},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": " 具"}},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "体内"},
|
||||
},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "容的 "},
|
||||
},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "Claude"},
|
||||
},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": " CLI"},
|
||||
},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": " 请"}},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "求体"},
|
||||
},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "发给我,"},
|
||||
},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "我会一"},
|
||||
},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "并"}},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "审"}},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "查整"},
|
||||
},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "个改"},
|
||||
},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "动。"},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
|
||||
"usage": {
|
||||
"input_tokens": 3,
|
||||
"cache_creation_input_tokens": 32760,
|
||||
"cache_read_input_tokens": 61777,
|
||||
"output_tokens": 42,
|
||||
},
|
||||
},
|
||||
{"type": "message_stop"},
|
||||
]
|
||||
|
||||
# 收集所有转换后的 OpenAI 格式事件
|
||||
all_openai_events: List[Dict[str, Any]] = []
|
||||
for chunk in chunks:
|
||||
events = reg.convert_stream_chunk(chunk, "CLAUDE_CLI", "OPENAI", state=state)
|
||||
all_openai_events.extend(events)
|
||||
|
||||
# 验证转换结果
|
||||
assert len(all_openai_events) > 0
|
||||
|
||||
# 第一个事件应该是 chat.completion.chunk(来自 message_start)
|
||||
assert all_openai_events[0].get("object") == "chat.completion.chunk"
|
||||
assert all_openai_events[0].get("model") == "claude-opus-4-5-20251101"
|
||||
|
||||
# 收集所有文本增量
|
||||
text_deltas = []
|
||||
for evt in all_openai_events:
|
||||
choices = evt.get("choices") or []
|
||||
if choices:
|
||||
delta = choices[0].get("delta", {})
|
||||
content = delta.get("content")
|
||||
if content:
|
||||
text_deltas.append(content)
|
||||
|
||||
# 验证文本拼接结果
|
||||
full_text = "".join(text_deltas)
|
||||
assert "明白了" in full_text
|
||||
assert "tools" in full_text
|
||||
assert "Claude CLI" in full_text
|
||||
assert "请求体发给我" in full_text
|
||||
|
||||
# 验证最后一个事件有 finish_reason
|
||||
last_with_finish = [
|
||||
e for e in all_openai_events if (e.get("choices") or [{}])[0].get("finish_reason")
|
||||
]
|
||||
assert len(last_with_finish) > 0
|
||||
assert last_with_finish[-1]["choices"][0]["finish_reason"] == "stop"
|
||||
|
||||
|
||||
def test_real_claude_cli_stream_to_openai_cli() -> None:
|
||||
"""测试 Claude CLI 流式响应转换为 OpenAI CLI (Responses API) 格式"""
|
||||
reg = _make_registry_with_cli()
|
||||
state = StreamState()
|
||||
|
||||
# 简化的真实事件序列
|
||||
chunks = [
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"model": "claude-opus-4-5-20251101",
|
||||
"id": "msg_test123",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"stop_reason": None,
|
||||
"usage": {"input_tokens": 10, "output_tokens": 0},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
|
||||
{"type": "ping"},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "Hello"}},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": " World"}},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5},
|
||||
},
|
||||
{"type": "message_stop"},
|
||||
]
|
||||
|
||||
all_events: List[Dict[str, Any]] = []
|
||||
for chunk in chunks:
|
||||
events = reg.convert_stream_chunk(chunk, "CLAUDE_CLI", "OPENAI_CLI", state=state)
|
||||
all_events.extend(events)
|
||||
|
||||
# 验证 OpenAI CLI 格式事件
|
||||
assert len(all_events) > 0
|
||||
|
||||
# 应该有 response.created 事件
|
||||
created_events = [e for e in all_events if e.get("type") == "response.created"]
|
||||
assert len(created_events) == 1
|
||||
|
||||
# 应该有 response.output_text.delta 事件
|
||||
delta_events = [e for e in all_events if e.get("type") == "response.output_text.delta"]
|
||||
assert len(delta_events) >= 2
|
||||
deltas = [e.get("delta") for e in delta_events]
|
||||
assert "Hello" in deltas
|
||||
assert " World" in deltas
|
||||
|
||||
# 应该有 response.completed 或 response.done 事件
|
||||
done_events = [e for e in all_events if e.get("type") in ("response.completed", "response.done")]
|
||||
assert len(done_events) >= 1
|
||||
@@ -3,7 +3,7 @@ is_format_compatible 单元测试
|
||||
|
||||
覆盖:
|
||||
- 同格式透传
|
||||
- CLI 格式禁止转换
|
||||
- CLI 格式允许转换(按 registry 能力)
|
||||
- 全局开关/端点开关/白黑名单
|
||||
- 流式转换开关
|
||||
- 转换器能力校验
|
||||
@@ -28,18 +28,21 @@ def test_same_format_is_compatible() -> None:
|
||||
assert reason is None
|
||||
|
||||
|
||||
def test_cli_format_not_convertible() -> None:
|
||||
def test_cli_format_convertible_when_converter_supports_full() -> None:
|
||||
registry = MagicMock()
|
||||
registry.can_convert_full.return_value = True
|
||||
|
||||
ok, needs_conv, reason = is_format_compatible(
|
||||
"CLAUDE_CLI",
|
||||
"OPENAI",
|
||||
endpoint_format_acceptance_config={"enabled": True},
|
||||
is_stream=False,
|
||||
global_conversion_enabled=True,
|
||||
registry=MagicMock(),
|
||||
registry=registry,
|
||||
)
|
||||
assert ok is False
|
||||
assert needs_conv is False
|
||||
assert reason and "CLI" in reason
|
||||
assert ok is True
|
||||
assert needs_conv is True
|
||||
assert reason is None
|
||||
|
||||
|
||||
def test_global_switch_disabled_blocks_conversion() -> None:
|
||||
@@ -158,4 +161,3 @@ def test_conversion_allowed_when_converter_supports_full() -> None:
|
||||
assert ok is True
|
||||
assert needs_conv is True
|
||||
assert reason is None
|
||||
|
||||
|
||||
71
tests/core/api_format/conversion/test_error_conversion.py
Normal file
71
tests/core/api_format/conversion/test_error_conversion.py
Normal file
@@ -0,0 +1,71 @@
|
||||
"""
|
||||
错误转换单元测试(Canonical)
|
||||
|
||||
重点:
|
||||
- registry_canonical.convert_error_response(_strict) 的基本链路
|
||||
- ErrorEvent 在 stream_event_from_internal 的输出形态
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, cast
|
||||
|
||||
from src.core.api_format.conversion.internal import ErrorType, InternalError
|
||||
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer
|
||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
||||
from src.core.api_format.conversion.stream_events import ErrorEvent
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
def _make_registry() -> FormatConversionRegistry:
|
||||
reg = FormatConversionRegistry()
|
||||
reg.register(OpenAINormalizer())
|
||||
reg.register(ClaudeNormalizer())
|
||||
reg.register(GeminiNormalizer())
|
||||
return reg
|
||||
|
||||
|
||||
def test_error_conversion_openai_to_claude() -> None:
|
||||
reg = _make_registry()
|
||||
|
||||
openai_error = {
|
||||
"error": {"message": "bad request", "type": "invalid_request_error", "code": "bad_request"}
|
||||
}
|
||||
|
||||
out = reg.convert_error_response(openai_error, "OPENAI", "CLAUDE")
|
||||
assert out.get("type") == "error"
|
||||
assert isinstance(out.get("error"), dict)
|
||||
assert out["error"]["message"] == "bad request"
|
||||
|
||||
|
||||
def test_error_conversion_claude_to_openai() -> None:
|
||||
reg = _make_registry()
|
||||
|
||||
claude_error = {"type": "error", "error": {"type": "invalid_request_error", "message": "nope"}}
|
||||
out = reg.convert_error_response(claude_error, "CLAUDE", "OPENAI")
|
||||
assert isinstance(out.get("error"), dict)
|
||||
assert out["error"]["message"] == "nope"
|
||||
|
||||
|
||||
def test_error_event_stream_output_openai() -> None:
|
||||
n = OpenAINormalizer()
|
||||
state = StreamState(model="gpt-4o-mini", message_id="chatcmpl_1")
|
||||
|
||||
internal = InternalError(type=ErrorType.INVALID_REQUEST, message="bad", retryable=False)
|
||||
events = n.stream_event_from_internal(ErrorEvent(error=internal), state)
|
||||
assert events == [{"error": {"message": "bad", "type": "invalid_request_error"}}]
|
||||
|
||||
|
||||
def test_error_event_stream_openai_to_claude_via_registry() -> None:
|
||||
reg = _make_registry()
|
||||
|
||||
# OpenAI 流式错误块
|
||||
chunk = {"error": {"message": "bad", "type": "invalid_request_error"}}
|
||||
out = reg.convert_stream_chunk(chunk, "OPENAI", "CLAUDE", state=StreamState())
|
||||
assert isinstance(out, list) and out
|
||||
evt0 = cast(Dict[str, Any], out[0])
|
||||
assert evt0.get("type") == "error"
|
||||
assert isinstance(evt0.get("error"), dict)
|
||||
assert evt0["error"]["message"] == "bad"
|
||||
273
tests/core/api_format/conversion/test_gemini_normalizer.py
Normal file
273
tests/core/api_format/conversion/test_gemini_normalizer.py
Normal file
@@ -0,0 +1,273 @@
|
||||
"""
|
||||
GeminiNormalizer 单元测试
|
||||
|
||||
覆盖重点:
|
||||
- systemInstruction/system_instruction -> instructions 的提取与还原
|
||||
- parts(text/inline_data/function_call/function_response/unknown)转换
|
||||
- finishReason/usageMetadata 映射
|
||||
- streaming chunk <-> InternalStreamEvent 的基础行为与状态
|
||||
- error <-> InternalError
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, cast
|
||||
|
||||
from src.core.api_format.conversion.internal import (
|
||||
ErrorType,
|
||||
ImageBlock,
|
||||
StopReason,
|
||||
TextBlock,
|
||||
ToolResultBlock,
|
||||
ToolUseBlock,
|
||||
UnknownBlock,
|
||||
)
|
||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||
from src.core.api_format.conversion.stream_events import (
|
||||
ContentDeltaEvent,
|
||||
MessageStartEvent,
|
||||
MessageStopEvent,
|
||||
ToolCallDeltaEvent,
|
||||
)
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
def test_gemini_request_system_and_generation_config_roundtrip() -> None:
|
||||
n = GeminiNormalizer()
|
||||
|
||||
req = {
|
||||
"model": "gemini-1.5",
|
||||
"systemInstruction": {"parts": [{"text": "sys"}]},
|
||||
"contents": [
|
||||
{"role": "user", "parts": [{"text": "hi"}]},
|
||||
{"role": "model", "parts": [{"text": "ok"}]},
|
||||
],
|
||||
"generationConfig": {
|
||||
"maxOutputTokens": 10,
|
||||
"temperature": 0.2,
|
||||
"topP": 0.9,
|
||||
"topK": 1,
|
||||
"stopSequences": ["A", "B"],
|
||||
},
|
||||
"tools": [
|
||||
{
|
||||
"functionDeclarations": [
|
||||
{
|
||||
"name": "get_weather",
|
||||
"description": "Get weather",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"toolConfig": {"functionCallingConfig": {"mode": "ANY"}},
|
||||
}
|
||||
|
||||
internal = n.request_to_internal(req)
|
||||
assert internal.model == "gemini-1.5"
|
||||
assert [seg.role.value for seg in internal.instructions] == ["system"]
|
||||
assert internal.instructions[0].text == "sys"
|
||||
assert internal.system == "sys"
|
||||
assert internal.max_tokens == 10
|
||||
assert internal.temperature == 0.2
|
||||
assert internal.top_p == 0.9
|
||||
assert internal.top_k == 1
|
||||
assert internal.stop_sequences == ["A", "B"]
|
||||
assert internal.tools is not None and internal.tools[0].name == "get_weather"
|
||||
assert internal.tool_choice is not None and internal.tool_choice.type.value == "required"
|
||||
|
||||
out = n.request_from_internal(internal)
|
||||
assert out["system_instruction"]["parts"][0]["text"] == "sys"
|
||||
assert out["generation_config"]["max_output_tokens"] == 10
|
||||
assert out["generation_config"]["stop_sequences"] == ["A", "B"]
|
||||
assert out["tools"][0]["function_declarations"][0]["name"] == "get_weather"
|
||||
assert out["tool_config"]["function_calling_config"]["mode"] == "ANY"
|
||||
|
||||
|
||||
def test_gemini_request_parts_image_tool_and_unknown_drop() -> None:
|
||||
n = GeminiNormalizer()
|
||||
|
||||
req = {
|
||||
"contents": [
|
||||
{
|
||||
"role": "user",
|
||||
"parts": [
|
||||
{"text": "look"},
|
||||
{"inline_data": {"mime_type": "image/png", "data": "AAAA"}},
|
||||
{"foo": 1},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "model",
|
||||
"parts": [
|
||||
{"function_call": {"name": "get_weather", "args": {"city": "SF"}}}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"parts": [
|
||||
{"function_response": {"name": "call_1", "response": {"result": {"temp_c": 20}}}}
|
||||
],
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
internal = n.request_to_internal(req)
|
||||
assert [m.role.value for m in internal.messages] == ["user", "assistant", "user"]
|
||||
|
||||
blocks0 = internal.messages[0].content
|
||||
assert any(isinstance(b, TextBlock) for b in blocks0)
|
||||
assert any(isinstance(b, ImageBlock) for b in blocks0)
|
||||
assert any(isinstance(b, UnknownBlock) for b in blocks0)
|
||||
|
||||
dropped = (internal.extra.get("raw") or {}).get("dropped_blocks") or {}
|
||||
assert dropped.get("gemini_part:foo") == 1
|
||||
|
||||
blocks1 = internal.messages[1].content
|
||||
tool_use = next(b for b in blocks1 if isinstance(b, ToolUseBlock))
|
||||
assert tool_use.tool_name == "get_weather"
|
||||
assert tool_use.tool_input == {"city": "SF"}
|
||||
|
||||
blocks2 = internal.messages[2].content
|
||||
tool_result = next(b for b in blocks2 if isinstance(b, ToolResultBlock))
|
||||
assert tool_result.tool_use_id == "call_1"
|
||||
assert tool_result.output == {"temp_c": 20}
|
||||
|
||||
out = n.request_from_internal(internal)
|
||||
out_contents: List[Dict[str, Any]] = out["contents"]
|
||||
|
||||
# unknown 被丢弃
|
||||
user_parts = cast(List[Dict[str, Any]], out_contents[0]["parts"])
|
||||
assert any(p.get("text") == "look" for p in user_parts)
|
||||
assert any(p.get("inline_data", {}).get("mime_type") == "image/png" for p in user_parts)
|
||||
assert all("foo" not in p for p in user_parts)
|
||||
|
||||
model_parts = cast(List[Dict[str, Any]], out_contents[1]["parts"])
|
||||
assert model_parts[0]["function_call"]["name"] == "get_weather"
|
||||
|
||||
tool_parts = cast(List[Dict[str, Any]], out_contents[2]["parts"])
|
||||
assert tool_parts[0]["function_response"]["name"] == "call_1"
|
||||
assert tool_parts[0]["function_response"]["response"]["result"] == {"temp_c": 20}
|
||||
|
||||
|
||||
def test_gemini_response_finish_reason_and_usage_roundtrip() -> None:
|
||||
n = GeminiNormalizer()
|
||||
|
||||
resp = {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": [{"text": "hi"}], "role": "model"},
|
||||
"finishReason": "MAX_TOKENS",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 5,
|
||||
"candidatesTokenCount": 7,
|
||||
"totalTokenCount": 12,
|
||||
"cachedContentTokenCount": 2,
|
||||
},
|
||||
"modelVersion": "gemini-1.5",
|
||||
}
|
||||
|
||||
internal = n.response_to_internal(resp)
|
||||
assert internal.model == "gemini-1.5"
|
||||
assert internal.stop_reason == StopReason.MAX_TOKENS
|
||||
assert internal.usage is not None
|
||||
assert internal.usage.input_tokens == 5
|
||||
assert internal.usage.output_tokens == 7
|
||||
assert internal.usage.total_tokens == 12
|
||||
assert internal.usage.cache_read_tokens == 2
|
||||
|
||||
out = n.response_from_internal(internal)
|
||||
assert out["candidates"][0]["finishReason"] == "MAX_TOKENS"
|
||||
assert out["usageMetadata"]["cachedContentTokenCount"] == 2
|
||||
|
||||
|
||||
def test_gemini_stream_chunk_and_event_roundtrip_basic() -> None:
|
||||
n = GeminiNormalizer()
|
||||
state = StreamState()
|
||||
|
||||
chunks = [
|
||||
{
|
||||
"candidates": [
|
||||
{"content": {"parts": [{"text": "Hel"}], "role": "model"}, "index": 0}
|
||||
],
|
||||
"modelVersion": "gemini-1.5",
|
||||
},
|
||||
{
|
||||
"candidates": [
|
||||
{"content": {"parts": [{"text": "lo"}], "role": "model"}, "index": 0}
|
||||
]
|
||||
},
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [
|
||||
{"functionCall": {"name": "get_weather", "args": {"city": "SF"}}}
|
||||
],
|
||||
"role": "model",
|
||||
},
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
"modelVersion": "gemini-1.5",
|
||||
},
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"parts": [], "role": "model"},
|
||||
"finishReason": "STOP",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
"usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 2, "totalTokenCount": 3},
|
||||
"modelVersion": "gemini-1.5",
|
||||
},
|
||||
]
|
||||
|
||||
events: List[Any] = []
|
||||
for ch in chunks:
|
||||
events.extend(n.stream_chunk_to_internal(ch, state))
|
||||
|
||||
assert any(isinstance(e, MessageStartEvent) for e in events)
|
||||
assert [e.text_delta for e in events if isinstance(e, ContentDeltaEvent)] == ["Hel", "lo"]
|
||||
assert any(isinstance(e, ToolCallDeltaEvent) and json.loads(e.input_delta) == {"city": "SF"} for e in events)
|
||||
assert any(isinstance(e, MessageStopEvent) and e.stop_reason == StopReason.END_TURN for e in events)
|
||||
|
||||
state2 = StreamState()
|
||||
out_chunks: List[Dict[str, Any]] = []
|
||||
for e in events:
|
||||
out_chunks.extend(n.stream_event_from_internal(e, state2))
|
||||
|
||||
assert any(c["candidates"][0]["content"]["parts"][0].get("text") == "Hel" for c in out_chunks)
|
||||
|
||||
tool_chunk = next(
|
||||
c
|
||||
for c in out_chunks
|
||||
if c["candidates"][0]["content"]["parts"]
|
||||
and "functionCall" in c["candidates"][0]["content"]["parts"][0]
|
||||
)
|
||||
assert tool_chunk["candidates"][0]["content"]["parts"][0]["functionCall"]["name"] == "get_weather"
|
||||
|
||||
assert out_chunks[-1]["candidates"][0]["finishReason"] == "STOP"
|
||||
|
||||
|
||||
def test_gemini_error_conversion() -> None:
|
||||
n = GeminiNormalizer()
|
||||
|
||||
err_resp = {"error": {"code": 429, "message": "slow down", "status": "RESOURCE_EXHAUSTED"}}
|
||||
assert n.is_error_response(err_resp) is True
|
||||
|
||||
internal = n.error_to_internal(err_resp)
|
||||
assert internal.type == ErrorType.RATE_LIMIT
|
||||
assert internal.retryable is True
|
||||
|
||||
out = n.error_from_internal(internal)
|
||||
assert out["error"]["status"] == "RESOURCE_EXHAUSTED"
|
||||
assert out["error"]["message"] == "slow down"
|
||||
116
tests/core/api_format/conversion/test_golden_canonical.py
Normal file
116
tests/core/api_format/conversion/test_golden_canonical.py
Normal file
@@ -0,0 +1,116 @@
|
||||
"""
|
||||
Golden tests(Canonical)
|
||||
|
||||
说明:
|
||||
- 这些 Golden 用于冻结 Canonical registry 的外部输出形态(request/response/stream)。
|
||||
- 文件由 `tools/generate_format_conversion_golden.py` 生成。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer
|
||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
GOLDEN_DIR = Path(__file__).resolve().parent / "golden_data"
|
||||
INPUT_DIR = GOLDEN_DIR / "inputs"
|
||||
EXPECTED_DIR = GOLDEN_DIR / "expected"
|
||||
|
||||
|
||||
def _scrub(obj: Any) -> Any:
|
||||
if isinstance(obj, list):
|
||||
return [_scrub(x) for x in obj]
|
||||
if isinstance(obj, dict):
|
||||
out: Dict[str, Any] = {}
|
||||
for k, v in obj.items():
|
||||
if k in {"created"}:
|
||||
continue
|
||||
if k == "system_fingerprint" and v is None:
|
||||
continue
|
||||
out[k] = _scrub(v)
|
||||
return out
|
||||
return obj
|
||||
|
||||
|
||||
def _load_json(path: Path) -> Any:
|
||||
return json.loads(path.read_text(encoding="utf-8"))
|
||||
|
||||
|
||||
def _make_registry() -> FormatConversionRegistry:
|
||||
reg = FormatConversionRegistry()
|
||||
reg.register(OpenAINormalizer())
|
||||
reg.register(ClaudeNormalizer())
|
||||
reg.register(GeminiNormalizer())
|
||||
return reg
|
||||
|
||||
|
||||
def test_golden_requests() -> None:
|
||||
reg = _make_registry()
|
||||
formats = ["OPENAI", "CLAUDE", "GEMINI"]
|
||||
|
||||
inputs = {
|
||||
"OPENAI": _load_json(INPUT_DIR / "request_openai.json"),
|
||||
"CLAUDE": _load_json(INPUT_DIR / "request_claude.json"),
|
||||
"GEMINI": _load_json(INPUT_DIR / "request_gemini.json"),
|
||||
}
|
||||
|
||||
for source in formats:
|
||||
for target in formats:
|
||||
if source == target:
|
||||
continue
|
||||
expected = _load_json(EXPECTED_DIR / f"request_{source}_to_{target}.json")
|
||||
actual = reg.convert_request(inputs[source], source, target)
|
||||
assert _scrub(actual) == expected
|
||||
|
||||
|
||||
def test_golden_responses() -> None:
|
||||
reg = _make_registry()
|
||||
formats = ["OPENAI", "CLAUDE", "GEMINI"]
|
||||
|
||||
inputs = {
|
||||
"OPENAI": _load_json(INPUT_DIR / "response_openai.json"),
|
||||
"CLAUDE": _load_json(INPUT_DIR / "response_claude.json"),
|
||||
"GEMINI": _load_json(INPUT_DIR / "response_gemini.json"),
|
||||
}
|
||||
|
||||
for source in formats:
|
||||
for target in formats:
|
||||
if source == target:
|
||||
continue
|
||||
expected = _load_json(EXPECTED_DIR / f"response_{source}_to_{target}.json")
|
||||
actual = reg.convert_response(inputs[source], source, target)
|
||||
assert _scrub(actual) == expected
|
||||
|
||||
|
||||
def test_golden_streams() -> None:
|
||||
reg = _make_registry()
|
||||
formats = ["OPENAI", "CLAUDE", "GEMINI"]
|
||||
|
||||
inputs: Dict[str, List[Dict[str, Any]]] = {
|
||||
"OPENAI": _load_json(INPUT_DIR / "stream_openai.json"),
|
||||
"CLAUDE": _load_json(INPUT_DIR / "stream_claude.json"),
|
||||
"GEMINI": _load_json(INPUT_DIR / "stream_gemini.json"),
|
||||
}
|
||||
|
||||
for source in formats:
|
||||
for target in formats:
|
||||
if source == target:
|
||||
continue
|
||||
expected = _load_json(EXPECTED_DIR / f"stream_{source}_to_{target}.json")
|
||||
|
||||
state = StreamState()
|
||||
if source == "GEMINI":
|
||||
state.message_id = "gemini_1"
|
||||
|
||||
out: List[Dict[str, Any]] = []
|
||||
for chunk in inputs[source]:
|
||||
out.extend(reg.convert_stream_chunk(chunk, source, target, state=state))
|
||||
|
||||
assert _scrub(out) == expected
|
||||
643
tests/core/api_format/conversion/test_internal.py
Normal file
643
tests/core/api_format/conversion/test_internal.py
Normal file
@@ -0,0 +1,643 @@
|
||||
"""
|
||||
internal 数据结构单元测试
|
||||
|
||||
目标:
|
||||
- 验证 dataclass/Enum 可正确实例化
|
||||
- 验证 ContentBlock 联合类型可用于运行时判断
|
||||
- 验证 StreamState.substate() 隔离机制
|
||||
- 验证各类型的默认值、字段访问、序列化
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import asdict
|
||||
from typing import get_args
|
||||
|
||||
import pytest
|
||||
|
||||
from src.core.api_format.conversion.internal import (
|
||||
ContentBlock,
|
||||
ContentType,
|
||||
ErrorType,
|
||||
FormatCapabilities,
|
||||
ImageBlock,
|
||||
InstructionSegment,
|
||||
InternalError,
|
||||
InternalMessage,
|
||||
InternalRequest,
|
||||
InternalResponse,
|
||||
Role,
|
||||
StopReason,
|
||||
TextBlock,
|
||||
ToolChoice,
|
||||
ToolChoiceType,
|
||||
ToolDefinition,
|
||||
ToolResultBlock,
|
||||
ToolUseBlock,
|
||||
UnknownBlock,
|
||||
UsageInfo,
|
||||
)
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Enum 类型测试
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestRoleEnum:
|
||||
def test_role_values(self) -> None:
|
||||
assert Role.USER.value == "user"
|
||||
assert Role.ASSISTANT.value == "assistant"
|
||||
assert Role.SYSTEM.value == "system"
|
||||
assert Role.DEVELOPER.value == "developer"
|
||||
assert Role.TOOL.value == "tool"
|
||||
assert Role.UNKNOWN.value == "unknown"
|
||||
|
||||
def test_role_is_str_subclass(self) -> None:
|
||||
assert isinstance(Role.USER, str)
|
||||
assert Role.USER == "user"
|
||||
|
||||
|
||||
class TestContentTypeEnum:
|
||||
def test_content_type_values(self) -> None:
|
||||
assert ContentType.TEXT.value == "text"
|
||||
assert ContentType.IMAGE.value == "image"
|
||||
assert ContentType.TOOL_USE.value == "tool_use"
|
||||
assert ContentType.TOOL_RESULT.value == "tool_result"
|
||||
assert ContentType.UNKNOWN.value == "unknown"
|
||||
|
||||
|
||||
class TestStopReasonEnum:
|
||||
def test_stop_reason_values(self) -> None:
|
||||
assert StopReason.END_TURN.value == "end_turn"
|
||||
assert StopReason.MAX_TOKENS.value == "max_tokens"
|
||||
assert StopReason.STOP_SEQUENCE.value == "stop_sequence"
|
||||
assert StopReason.TOOL_USE.value == "tool_use"
|
||||
assert StopReason.PAUSE_TURN.value == "pause_turn"
|
||||
assert StopReason.REFUSAL.value == "refusal"
|
||||
assert StopReason.CONTENT_FILTERED.value == "content_filtered"
|
||||
assert StopReason.UNKNOWN.value == "unknown"
|
||||
|
||||
|
||||
class TestErrorTypeEnum:
|
||||
def test_error_type_values(self) -> None:
|
||||
assert ErrorType.INVALID_REQUEST.value == "invalid_request"
|
||||
assert ErrorType.AUTHENTICATION.value == "authentication"
|
||||
assert ErrorType.PERMISSION_DENIED.value == "permission_denied"
|
||||
assert ErrorType.NOT_FOUND.value == "not_found"
|
||||
assert ErrorType.RATE_LIMIT.value == "rate_limit"
|
||||
assert ErrorType.OVERLOADED.value == "overloaded"
|
||||
assert ErrorType.SERVER_ERROR.value == "server_error"
|
||||
assert ErrorType.CONTENT_FILTERED.value == "content_filtered"
|
||||
assert ErrorType.CONTEXT_LENGTH_EXCEEDED.value == "context_length_exceeded"
|
||||
assert ErrorType.UNKNOWN.value == "unknown"
|
||||
|
||||
|
||||
class TestToolChoiceTypeEnum:
|
||||
def test_tool_choice_type_values(self) -> None:
|
||||
assert ToolChoiceType.AUTO.value == "auto"
|
||||
assert ToolChoiceType.NONE.value == "none"
|
||||
assert ToolChoiceType.REQUIRED.value == "required"
|
||||
assert ToolChoiceType.TOOL.value == "tool"
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# ContentBlock 类型测试
|
||||
# ============================================================================
|
||||
|
||||
|
||||
def test_content_block_runtime_check() -> None:
|
||||
block_types = get_args(ContentBlock)
|
||||
assert block_types, "typing.get_args(ContentBlock) 应返回可用的类型列表"
|
||||
|
||||
assert isinstance(TextBlock(text="hi"), block_types)
|
||||
assert isinstance(ImageBlock(url="https://example.com/a.png"), block_types)
|
||||
assert isinstance(ToolUseBlock(tool_id="t1", tool_name="x"), block_types)
|
||||
assert isinstance(ToolResultBlock(tool_use_id="t1", output={"ok": True}), block_types)
|
||||
assert isinstance(UnknownBlock(raw_type="weird", payload={"x": 1}), block_types)
|
||||
|
||||
|
||||
def test_internal_error_to_debug_dict() -> None:
|
||||
err = InternalError(
|
||||
type=ErrorType.INVALID_REQUEST,
|
||||
message="bad request",
|
||||
code="bad_request",
|
||||
param="messages",
|
||||
retryable=False,
|
||||
extra={"raw": {"type": "invalid_request_error"}},
|
||||
)
|
||||
d = err.to_debug_dict()
|
||||
assert d["type"] == "invalid_request"
|
||||
assert d["message"] == "bad request"
|
||||
assert d["code"] == "bad_request"
|
||||
assert d["param"] == "messages"
|
||||
assert d["retryable"] is False
|
||||
assert isinstance(d["extra"], dict)
|
||||
|
||||
|
||||
def test_internal_request_debug_dict_and_serialization() -> None:
|
||||
req = InternalRequest(
|
||||
model="m",
|
||||
messages=[],
|
||||
instructions=[
|
||||
InstructionSegment(role=Role.SYSTEM, text="sys1"),
|
||||
InstructionSegment(role=Role.DEVELOPER, text="dev1"),
|
||||
],
|
||||
system="sys1\n\ndev1",
|
||||
stream=True,
|
||||
extra={"raw": {"openai_messages": []}},
|
||||
)
|
||||
d = req.to_debug_dict()
|
||||
assert d["instruction_count"] == 2
|
||||
assert d["message_count"] == 0
|
||||
assert d["has_system"] is True
|
||||
assert d["stream"] is True
|
||||
|
||||
# asdict 只要能跑通即可(Enum 会保留为对象,属于预期)
|
||||
dumped = asdict(req)
|
||||
assert dumped["model"] == "m"
|
||||
assert isinstance(dumped["instructions"], list)
|
||||
|
||||
|
||||
def test_internal_response_debug_dict() -> None:
|
||||
resp = InternalResponse(
|
||||
id="r1",
|
||||
model="m",
|
||||
content=[TextBlock(text="hi")],
|
||||
stop_reason=StopReason.END_TURN,
|
||||
usage=UsageInfo(input_tokens=1, output_tokens=2, total_tokens=3),
|
||||
)
|
||||
d = resp.to_debug_dict()
|
||||
assert d["id"] == "r1"
|
||||
assert d["stop_reason"] == "end_turn"
|
||||
assert d["usage"] == {"input": 1, "output": 2}
|
||||
|
||||
|
||||
def test_stream_state_substate_isolated() -> None:
|
||||
state = StreamState()
|
||||
openai_state = state.substate("openai")
|
||||
claude_state = state.substate("CLAUDE")
|
||||
|
||||
openai_state["x"] = 1
|
||||
assert "x" not in claude_state
|
||||
assert state.substate("OPENAI") is openai_state
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# ContentBlock 各子类型详细测试
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestTextBlock:
|
||||
def test_default_values(self) -> None:
|
||||
block = TextBlock()
|
||||
assert block.type == ContentType.TEXT
|
||||
assert block.text == ""
|
||||
assert block.extra == {}
|
||||
|
||||
def test_with_content(self) -> None:
|
||||
block = TextBlock(text="Hello, world!", extra={"source": "test"})
|
||||
assert block.text == "Hello, world!"
|
||||
assert block.extra == {"source": "test"}
|
||||
|
||||
def test_type_is_readonly(self) -> None:
|
||||
block = TextBlock(text="hi")
|
||||
assert block.type == ContentType.TEXT
|
||||
|
||||
|
||||
class TestImageBlock:
|
||||
def test_with_url(self) -> None:
|
||||
block = ImageBlock(url="https://example.com/img.png")
|
||||
assert block.type == ContentType.IMAGE
|
||||
assert block.url == "https://example.com/img.png"
|
||||
assert block.data is None
|
||||
assert block.media_type is None
|
||||
|
||||
def test_with_base64_data(self) -> None:
|
||||
block = ImageBlock(data="base64encodeddata", media_type="image/png")
|
||||
assert block.data == "base64encodeddata"
|
||||
assert block.media_type == "image/png"
|
||||
assert block.url is None
|
||||
|
||||
|
||||
class TestToolUseBlock:
|
||||
def test_default_values(self) -> None:
|
||||
block = ToolUseBlock()
|
||||
assert block.type == ContentType.TOOL_USE
|
||||
assert block.tool_id == ""
|
||||
assert block.tool_name == ""
|
||||
assert block.tool_input == {}
|
||||
|
||||
def test_with_values(self) -> None:
|
||||
block = ToolUseBlock(
|
||||
tool_id="call_123",
|
||||
tool_name="get_weather",
|
||||
tool_input={"city": "Beijing"},
|
||||
)
|
||||
assert block.tool_id == "call_123"
|
||||
assert block.tool_name == "get_weather"
|
||||
assert block.tool_input == {"city": "Beijing"}
|
||||
|
||||
|
||||
class TestToolResultBlock:
|
||||
def test_default_values(self) -> None:
|
||||
block = ToolResultBlock()
|
||||
assert block.type == ContentType.TOOL_RESULT
|
||||
assert block.tool_use_id == ""
|
||||
assert block.output is None
|
||||
assert block.content_text is None
|
||||
assert block.is_error is False
|
||||
|
||||
def test_with_success_output(self) -> None:
|
||||
block = ToolResultBlock(
|
||||
tool_use_id="call_123",
|
||||
output={"temperature": 25},
|
||||
content_text="Temperature: 25C",
|
||||
)
|
||||
assert block.tool_use_id == "call_123"
|
||||
assert block.output == {"temperature": 25}
|
||||
assert block.is_error is False
|
||||
|
||||
def test_with_error_output(self) -> None:
|
||||
block = ToolResultBlock(
|
||||
tool_use_id="call_456",
|
||||
output="Error: city not found",
|
||||
is_error=True,
|
||||
)
|
||||
assert block.is_error is True
|
||||
|
||||
|
||||
class TestUnknownBlock:
|
||||
def test_default_values(self) -> None:
|
||||
block = UnknownBlock()
|
||||
assert block.type == ContentType.UNKNOWN
|
||||
assert block.raw_type == ""
|
||||
assert block.payload == {}
|
||||
|
||||
def test_with_values(self) -> None:
|
||||
block = UnknownBlock(
|
||||
raw_type="custom_block",
|
||||
payload={"key": "value"},
|
||||
extra={"source": "gemini"},
|
||||
)
|
||||
assert block.raw_type == "custom_block"
|
||||
assert block.payload == {"key": "value"}
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# InternalMessage 测试
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestInternalMessage:
|
||||
def test_simple_user_message(self) -> None:
|
||||
msg = InternalMessage(role=Role.USER, content=[TextBlock(text="Hello")])
|
||||
assert msg.role == Role.USER
|
||||
assert len(msg.content) == 1
|
||||
assert isinstance(msg.content[0], TextBlock)
|
||||
|
||||
def test_assistant_with_tool_use(self) -> None:
|
||||
msg = InternalMessage(
|
||||
role=Role.ASSISTANT,
|
||||
content=[
|
||||
TextBlock(text="Let me check the weather"),
|
||||
ToolUseBlock(tool_id="t1", tool_name="get_weather", tool_input={"city": "Shanghai"}),
|
||||
],
|
||||
)
|
||||
assert msg.role == Role.ASSISTANT
|
||||
assert len(msg.content) == 2
|
||||
assert msg.content[0].type == ContentType.TEXT
|
||||
assert msg.content[1].type == ContentType.TOOL_USE
|
||||
|
||||
def test_tool_message(self) -> None:
|
||||
msg = InternalMessage(
|
||||
role=Role.TOOL,
|
||||
content=[ToolResultBlock(tool_use_id="t1", output="Sunny, 28C")],
|
||||
)
|
||||
assert msg.role == Role.TOOL
|
||||
|
||||
def test_with_extra(self) -> None:
|
||||
msg = InternalMessage(
|
||||
role=Role.USER,
|
||||
content=[],
|
||||
extra={"original_format": "openai"},
|
||||
)
|
||||
assert msg.extra == {"original_format": "openai"}
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# InstructionSegment 测试
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestInstructionSegment:
|
||||
def test_system_instruction(self) -> None:
|
||||
seg = InstructionSegment(role=Role.SYSTEM, text="You are a helpful assistant.")
|
||||
assert seg.role == Role.SYSTEM
|
||||
assert seg.text == "You are a helpful assistant."
|
||||
|
||||
def test_developer_instruction(self) -> None:
|
||||
seg = InstructionSegment(role=Role.DEVELOPER, text="Always respond in JSON format.")
|
||||
assert seg.role == Role.DEVELOPER
|
||||
assert seg.text == "Always respond in JSON format."
|
||||
|
||||
def test_with_extra(self) -> None:
|
||||
seg = InstructionSegment(
|
||||
role=Role.SYSTEM,
|
||||
text="test",
|
||||
extra={"cache_control": {"type": "ephemeral"}},
|
||||
)
|
||||
assert seg.extra == {"cache_control": {"type": "ephemeral"}}
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# ToolDefinition 和 ToolChoice 测试
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestToolDefinition:
|
||||
def test_minimal(self) -> None:
|
||||
tool = ToolDefinition(name="get_time")
|
||||
assert tool.name == "get_time"
|
||||
assert tool.description is None
|
||||
assert tool.parameters is None
|
||||
|
||||
def test_full(self) -> None:
|
||||
tool = ToolDefinition(
|
||||
name="get_weather",
|
||||
description="Get current weather for a city",
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"city": {"type": "string", "description": "City name"},
|
||||
},
|
||||
"required": ["city"],
|
||||
},
|
||||
extra={"strict": True},
|
||||
)
|
||||
assert tool.name == "get_weather"
|
||||
assert tool.description == "Get current weather for a city"
|
||||
assert tool.parameters is not None
|
||||
assert "city" in tool.parameters["properties"]
|
||||
|
||||
|
||||
class TestToolChoice:
|
||||
def test_auto(self) -> None:
|
||||
choice = ToolChoice(type=ToolChoiceType.AUTO)
|
||||
assert choice.type == ToolChoiceType.AUTO
|
||||
assert choice.tool_name is None
|
||||
|
||||
def test_none(self) -> None:
|
||||
choice = ToolChoice(type=ToolChoiceType.NONE)
|
||||
assert choice.type == ToolChoiceType.NONE
|
||||
|
||||
def test_required(self) -> None:
|
||||
choice = ToolChoice(type=ToolChoiceType.REQUIRED)
|
||||
assert choice.type == ToolChoiceType.REQUIRED
|
||||
|
||||
def test_specific_tool(self) -> None:
|
||||
choice = ToolChoice(type=ToolChoiceType.TOOL, tool_name="get_weather")
|
||||
assert choice.type == ToolChoiceType.TOOL
|
||||
assert choice.tool_name == "get_weather"
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# UsageInfo 测试
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestUsageInfo:
|
||||
def test_default_values(self) -> None:
|
||||
usage = UsageInfo()
|
||||
assert usage.input_tokens == 0
|
||||
assert usage.output_tokens == 0
|
||||
assert usage.total_tokens == 0
|
||||
assert usage.cache_read_tokens == 0
|
||||
assert usage.cache_write_tokens == 0
|
||||
|
||||
def test_with_values(self) -> None:
|
||||
usage = UsageInfo(
|
||||
input_tokens=100,
|
||||
output_tokens=50,
|
||||
total_tokens=150,
|
||||
cache_read_tokens=20,
|
||||
cache_write_tokens=10,
|
||||
)
|
||||
assert usage.input_tokens == 100
|
||||
assert usage.output_tokens == 50
|
||||
assert usage.total_tokens == 150
|
||||
assert usage.cache_read_tokens == 20
|
||||
assert usage.cache_write_tokens == 10
|
||||
|
||||
def test_with_extra(self) -> None:
|
||||
usage = UsageInfo(
|
||||
input_tokens=10,
|
||||
output_tokens=5,
|
||||
extra={"reasoning_tokens": 100},
|
||||
)
|
||||
assert usage.extra == {"reasoning_tokens": 100}
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# FormatCapabilities 测试
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestFormatCapabilities:
|
||||
def test_default_values(self) -> None:
|
||||
caps = FormatCapabilities()
|
||||
assert caps.supports_stream is True
|
||||
assert caps.supports_error_conversion is True
|
||||
assert caps.supports_tools is True
|
||||
assert caps.supports_images is False
|
||||
assert caps.supported_features == frozenset()
|
||||
|
||||
def test_custom_values(self) -> None:
|
||||
caps = FormatCapabilities(
|
||||
supports_stream=False,
|
||||
supports_error_conversion=True,
|
||||
supports_tools=False,
|
||||
supports_images=True,
|
||||
supported_features=frozenset({"vision", "function_calling"}),
|
||||
)
|
||||
assert caps.supports_stream is False
|
||||
assert caps.supports_images is True
|
||||
assert "vision" in caps.supported_features
|
||||
|
||||
def test_is_frozen(self) -> None:
|
||||
caps = FormatCapabilities()
|
||||
with pytest.raises(Exception):
|
||||
caps.supports_stream = False # type: ignore
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# InternalRequest 详细测试
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestInternalRequest:
|
||||
def test_minimal(self) -> None:
|
||||
req = InternalRequest(model="gpt-4", messages=[])
|
||||
assert req.model == "gpt-4"
|
||||
assert req.messages == []
|
||||
assert req.stream is False
|
||||
assert req.max_tokens is None
|
||||
assert req.tools is None
|
||||
|
||||
def test_with_messages(self) -> None:
|
||||
req = InternalRequest(
|
||||
model="claude-3",
|
||||
messages=[
|
||||
InternalMessage(role=Role.USER, content=[TextBlock(text="Hi")]),
|
||||
InternalMessage(role=Role.ASSISTANT, content=[TextBlock(text="Hello!")]),
|
||||
],
|
||||
)
|
||||
assert len(req.messages) == 2
|
||||
assert req.messages[0].role == Role.USER
|
||||
assert req.messages[1].role == Role.ASSISTANT
|
||||
|
||||
def test_with_tools(self) -> None:
|
||||
req = InternalRequest(
|
||||
model="gpt-4",
|
||||
messages=[],
|
||||
tools=[ToolDefinition(name="get_time"), ToolDefinition(name="get_weather")],
|
||||
tool_choice=ToolChoice(type=ToolChoiceType.AUTO),
|
||||
)
|
||||
assert req.tools is not None
|
||||
assert len(req.tools) == 2
|
||||
assert req.tool_choice is not None
|
||||
assert req.tool_choice.type == ToolChoiceType.AUTO
|
||||
|
||||
def test_to_debug_dict_with_tools(self) -> None:
|
||||
req = InternalRequest(
|
||||
model="m",
|
||||
messages=[InternalMessage(role=Role.USER, content=[TextBlock(text="test")])],
|
||||
tools=[ToolDefinition(name="tool1")],
|
||||
stream=True,
|
||||
extra={"key": "value"},
|
||||
)
|
||||
d = req.to_debug_dict()
|
||||
assert d["tool_count"] == 1
|
||||
assert d["message_count"] == 1
|
||||
assert d["stream"] is True
|
||||
assert "key" in d["extra_keys"]
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# InternalResponse 详细测试
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestInternalResponse:
|
||||
def test_minimal(self) -> None:
|
||||
resp = InternalResponse(id="r1", model="gpt-4", content=[])
|
||||
assert resp.id == "r1"
|
||||
assert resp.model == "gpt-4"
|
||||
assert resp.content == []
|
||||
assert resp.stop_reason is None
|
||||
assert resp.usage is None
|
||||
|
||||
def test_with_tool_use(self) -> None:
|
||||
resp = InternalResponse(
|
||||
id="r1",
|
||||
model="claude-3",
|
||||
content=[
|
||||
TextBlock(text="I'll check that for you"),
|
||||
ToolUseBlock(tool_id="t1", tool_name="search", tool_input={"q": "test"}),
|
||||
],
|
||||
stop_reason=StopReason.TOOL_USE,
|
||||
)
|
||||
assert resp.stop_reason == StopReason.TOOL_USE
|
||||
assert len(resp.content) == 2
|
||||
|
||||
def test_to_debug_dict_without_usage(self) -> None:
|
||||
resp = InternalResponse(id="r1", model="m", content=[])
|
||||
d = resp.to_debug_dict()
|
||||
assert d["usage"] is None
|
||||
|
||||
def test_to_debug_dict_with_all_stop_reasons(self) -> None:
|
||||
for reason in StopReason:
|
||||
resp = InternalResponse(
|
||||
id="r1",
|
||||
model="m",
|
||||
content=[],
|
||||
stop_reason=reason,
|
||||
)
|
||||
d = resp.to_debug_dict()
|
||||
assert d["stop_reason"] == reason.value
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# InternalError 详细测试
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestInternalError:
|
||||
def test_minimal(self) -> None:
|
||||
err = InternalError(type=ErrorType.UNKNOWN, message="Unknown error")
|
||||
assert err.type == ErrorType.UNKNOWN
|
||||
assert err.message == "Unknown error"
|
||||
assert err.code is None
|
||||
assert err.retryable is False
|
||||
|
||||
def test_retryable_error(self) -> None:
|
||||
err = InternalError(
|
||||
type=ErrorType.RATE_LIMIT,
|
||||
message="Rate limit exceeded",
|
||||
code="rate_limit_exceeded",
|
||||
retryable=True,
|
||||
)
|
||||
assert err.type == ErrorType.RATE_LIMIT
|
||||
assert err.retryable is True
|
||||
|
||||
def test_all_error_types_in_debug_dict(self) -> None:
|
||||
for err_type in ErrorType:
|
||||
err = InternalError(type=err_type, message="test")
|
||||
d = err.to_debug_dict()
|
||||
assert d["type"] == err_type.value
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# asdict 序列化测试
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestSerialization:
|
||||
def test_content_block_asdict(self) -> None:
|
||||
block = TextBlock(text="hello")
|
||||
d = asdict(block)
|
||||
assert d["text"] == "hello"
|
||||
assert d["type"] == ContentType.TEXT
|
||||
|
||||
def test_tool_use_block_asdict(self) -> None:
|
||||
block = ToolUseBlock(
|
||||
tool_id="t1",
|
||||
tool_name="test",
|
||||
tool_input={"a": 1},
|
||||
)
|
||||
d = asdict(block)
|
||||
assert d["tool_id"] == "t1"
|
||||
assert d["tool_input"] == {"a": 1}
|
||||
|
||||
def test_message_asdict(self) -> None:
|
||||
msg = InternalMessage(
|
||||
role=Role.USER,
|
||||
content=[TextBlock(text="hi")],
|
||||
)
|
||||
d = asdict(msg)
|
||||
assert d["role"] == Role.USER
|
||||
assert len(d["content"]) == 1
|
||||
|
||||
def test_response_asdict(self) -> None:
|
||||
resp = InternalResponse(
|
||||
id="r1",
|
||||
model="m",
|
||||
content=[TextBlock(text="response")],
|
||||
stop_reason=StopReason.END_TURN,
|
||||
usage=UsageInfo(input_tokens=10, output_tokens=5),
|
||||
)
|
||||
d = asdict(resp)
|
||||
assert d["id"] == "r1"
|
||||
assert d["stop_reason"] == StopReason.END_TURN
|
||||
assert d["usage"]["input_tokens"] == 10
|
||||
317
tests/core/api_format/conversion/test_openai_normalizer.py
Normal file
317
tests/core/api_format/conversion/test_openai_normalizer.py
Normal file
@@ -0,0 +1,317 @@
|
||||
"""
|
||||
OpenAINormalizer 单元测试
|
||||
|
||||
覆盖重点:
|
||||
- system/developer -> instructions 的提取与还原
|
||||
- tool_calls 与 tool role 的往返转换
|
||||
- content parts(text/image/unknown):UnknownBlock 内部保留、输出默认丢弃
|
||||
- finish_reason/usage 的映射
|
||||
- streaming chunk <-> InternalStreamEvent 的基础行为与状态
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, cast
|
||||
|
||||
from src.core.api_format.conversion.internal import (
|
||||
ContentType,
|
||||
ErrorType,
|
||||
ImageBlock,
|
||||
StopReason,
|
||||
TextBlock,
|
||||
ToolResultBlock,
|
||||
ToolUseBlock,
|
||||
UnknownBlock,
|
||||
)
|
||||
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||
from src.core.api_format.conversion.stream_events import (
|
||||
ContentBlockStartEvent,
|
||||
ContentDeltaEvent,
|
||||
MessageStartEvent,
|
||||
MessageStopEvent,
|
||||
ToolCallDeltaEvent,
|
||||
)
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
def _first_choice_message(response: Dict[str, Any]) -> Dict[str, Any]:
|
||||
choices = response.get("choices") or []
|
||||
assert isinstance(choices, list) and choices
|
||||
c0 = choices[0]
|
||||
assert isinstance(c0, dict)
|
||||
msg = c0.get("message")
|
||||
assert isinstance(msg, dict)
|
||||
return cast(Dict[str, Any], msg)
|
||||
|
||||
|
||||
def test_openai_request_instructions_roundtrip() -> None:
|
||||
n = OpenAINormalizer()
|
||||
|
||||
req = {
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [
|
||||
{"role": "system", "content": "sys"},
|
||||
{"role": "developer", "content": [{"type": "text", "text": "dev"}]},
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": "ok"},
|
||||
],
|
||||
"max_tokens": 12,
|
||||
"temperature": 0.2,
|
||||
"stop": ["A", "B"],
|
||||
"stream": True,
|
||||
}
|
||||
|
||||
internal = n.request_to_internal(req)
|
||||
assert internal.model == "gpt-4o-mini"
|
||||
assert [seg.role.value for seg in internal.instructions] == ["system", "developer"]
|
||||
assert [seg.text for seg in internal.instructions] == ["sys", "dev"]
|
||||
assert internal.system == "sys\n\ndev"
|
||||
assert [m.role.value for m in internal.messages] == ["user", "assistant"]
|
||||
assert internal.stop_sequences == ["A", "B"]
|
||||
assert internal.stream is True
|
||||
|
||||
out = n.request_from_internal(internal)
|
||||
assert out["model"] == "gpt-4o-mini"
|
||||
out_messages = out["messages"]
|
||||
assert [m["role"] for m in out_messages[:2]] == ["system", "developer"]
|
||||
assert [m["content"] for m in out_messages[:2]] == ["sys", "dev"]
|
||||
assert [m["role"] for m in out_messages[2:]] == ["user", "assistant"]
|
||||
assert out_messages[2]["content"] == "hi"
|
||||
assert out_messages[3]["content"] == "ok"
|
||||
|
||||
|
||||
def test_openai_request_tool_calls_and_tool_role_roundtrip() -> None:
|
||||
n = OpenAINormalizer()
|
||||
|
||||
req = {
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [
|
||||
{"role": "user", "content": "weather?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": '{"city":"SF"}'},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_1",
|
||||
"content": '{"temp_c": 20, "unit": "C"}',
|
||||
},
|
||||
{"role": "assistant", "content": "done"},
|
||||
],
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get weather",
|
||||
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}},
|
||||
},
|
||||
}
|
||||
],
|
||||
"tool_choice": "auto",
|
||||
}
|
||||
|
||||
internal = n.request_to_internal(req)
|
||||
assert [m.role.value for m in internal.messages] == ["user", "assistant", "user", "assistant"]
|
||||
|
||||
assistant_msg = internal.messages[1]
|
||||
assert any(isinstance(b, ToolUseBlock) for b in assistant_msg.content)
|
||||
tool_use = next(b for b in assistant_msg.content if isinstance(b, ToolUseBlock))
|
||||
assert tool_use.tool_id == "call_1"
|
||||
assert tool_use.tool_name == "get_weather"
|
||||
assert tool_use.tool_input == {"city": "SF"}
|
||||
|
||||
tool_result_msg = internal.messages[2]
|
||||
assert any(isinstance(b, ToolResultBlock) for b in tool_result_msg.content)
|
||||
tool_result = next(b for b in tool_result_msg.content if isinstance(b, ToolResultBlock))
|
||||
assert tool_result.tool_use_id == "call_1"
|
||||
assert tool_result.output == {"temp_c": 20, "unit": "C"}
|
||||
|
||||
out = n.request_from_internal(internal)
|
||||
out_messages: List[Dict[str, Any]] = out["messages"]
|
||||
|
||||
roles = [m.get("role") for m in out_messages]
|
||||
assert roles == ["user", "assistant", "tool", "assistant"]
|
||||
|
||||
assistant_out = out_messages[1]
|
||||
assert "tool_calls" in assistant_out
|
||||
assert assistant_out["tool_calls"][0]["function"]["name"] == "get_weather"
|
||||
|
||||
tool_out = out_messages[2]
|
||||
assert tool_out["role"] == "tool"
|
||||
assert tool_out["tool_call_id"] == "call_1"
|
||||
assert json.loads(tool_out["content"]) == {"temp_c": 20, "unit": "C"}
|
||||
|
||||
|
||||
def test_openai_request_content_image_and_unknown_drop() -> None:
|
||||
n = OpenAINormalizer()
|
||||
|
||||
req = {
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "look"},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.com/a.png"}},
|
||||
{"type": "foo", "x": 1},
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
internal = n.request_to_internal(req)
|
||||
assert len(internal.messages) == 1
|
||||
blocks = internal.messages[0].content
|
||||
assert any(isinstance(b, TextBlock) for b in blocks)
|
||||
assert any(isinstance(b, ImageBlock) for b in blocks)
|
||||
assert any(isinstance(b, UnknownBlock) for b in blocks)
|
||||
u = next(b for b in blocks if isinstance(b, UnknownBlock))
|
||||
assert u.raw_type == "foo"
|
||||
|
||||
dropped = (internal.extra.get("raw") or {}).get("dropped_blocks") or {}
|
||||
assert dropped.get("openai_part:foo") == 1
|
||||
|
||||
out = n.request_from_internal(internal)
|
||||
out_msg = out["messages"][0]
|
||||
assert out_msg["role"] == "user"
|
||||
assert isinstance(out_msg["content"], list)
|
||||
parts = cast(List[Dict[str, Any]], out_msg["content"])
|
||||
assert any(p.get("type") == "image_url" for p in parts)
|
||||
assert all(p.get("type") != "foo" for p in parts)
|
||||
|
||||
|
||||
def test_openai_response_finish_reason_and_usage_roundtrip() -> None:
|
||||
n = OpenAINormalizer()
|
||||
|
||||
resp = {
|
||||
"id": "chatcmpl_1",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "hi"},
|
||||
"finish_reason": "length",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 5, "completion_tokens": 7, "total_tokens": 12},
|
||||
}
|
||||
|
||||
internal = n.response_to_internal(resp)
|
||||
assert internal.id == "chatcmpl_1"
|
||||
assert internal.stop_reason == StopReason.MAX_TOKENS
|
||||
assert internal.usage is not None
|
||||
assert internal.usage.input_tokens == 5
|
||||
assert internal.usage.output_tokens == 7
|
||||
assert internal.usage.total_tokens == 12
|
||||
|
||||
out = n.response_from_internal(internal)
|
||||
out_msg = _first_choice_message(out)
|
||||
assert out_msg["role"] == "assistant"
|
||||
assert out_msg["content"] == "hi"
|
||||
assert out["choices"][0]["finish_reason"] == "length"
|
||||
assert out["usage"] == {"prompt_tokens": 5, "completion_tokens": 7, "total_tokens": 12}
|
||||
|
||||
|
||||
def test_openai_stream_chunk_and_event_roundtrip_basic() -> None:
|
||||
n = OpenAINormalizer()
|
||||
state = StreamState()
|
||||
|
||||
chunks = [
|
||||
{
|
||||
"id": "chatcmpl_stream_1",
|
||||
"object": "chat.completion.chunk",
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [{"index": 0, "delta": {"content": "Hel"}, "finish_reason": None}],
|
||||
},
|
||||
{"choices": [{"index": 0, "delta": {"content": "lo"}, "finish_reason": None}]},
|
||||
{
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": 0,
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": '{"city":"SF"}'},
|
||||
}
|
||||
]
|
||||
},
|
||||
"finish_reason": None,
|
||||
}
|
||||
]
|
||||
},
|
||||
{"choices": [{"index": 0, "delta": {}, "finish_reason": "tool_calls"}]},
|
||||
]
|
||||
|
||||
events: List[Any] = []
|
||||
for ch in chunks:
|
||||
events.extend(n.stream_chunk_to_internal(ch, state))
|
||||
|
||||
assert any(isinstance(e, MessageStartEvent) for e in events)
|
||||
assert any(isinstance(e, ContentBlockStartEvent) for e in events)
|
||||
assert [e.text_delta for e in events if isinstance(e, ContentDeltaEvent)] == ["Hel", "lo"]
|
||||
assert any(isinstance(e, ToolCallDeltaEvent) for e in events)
|
||||
assert any(isinstance(e, MessageStopEvent) and e.stop_reason == StopReason.TOOL_USE for e in events)
|
||||
|
||||
# internal events -> OpenAI chunks(验证关键字段与 tool_calls index 稳定)
|
||||
state2 = StreamState()
|
||||
out_chunks: List[Dict[str, Any]] = []
|
||||
for e in events:
|
||||
out_chunks.extend(n.stream_event_from_internal(e, state2))
|
||||
|
||||
# 第一个 chunk 应包含 role=assistant
|
||||
assert out_chunks[0]["choices"][0]["delta"].get("role") == "assistant"
|
||||
|
||||
# 至少包含一个 content delta
|
||||
assert any(c["choices"][0]["delta"].get("content") == "Hel" for c in out_chunks)
|
||||
|
||||
# tool_calls start chunk
|
||||
tool_start = next(
|
||||
c for c in out_chunks if c["choices"][0]["delta"].get("tool_calls") and c["choices"][0]["delta"]["tool_calls"][0]["function"].get("name")
|
||||
)
|
||||
assert tool_start["choices"][0]["delta"]["tool_calls"][0]["id"] == "call_1"
|
||||
assert tool_start["choices"][0]["delta"]["tool_calls"][0]["index"] == 0
|
||||
|
||||
# tool_calls delta chunk(arguments 片段)
|
||||
tool_delta = next(
|
||||
c for c in out_chunks if c["choices"][0]["delta"].get("tool_calls") and "arguments" in c["choices"][0]["delta"]["tool_calls"][0]["function"]
|
||||
)
|
||||
assert tool_delta["choices"][0]["delta"]["tool_calls"][0]["id"] == "call_1"
|
||||
assert tool_delta["choices"][0]["delta"]["tool_calls"][0]["index"] == 0
|
||||
|
||||
# 最终 stop chunk finish_reason=tool_calls
|
||||
assert out_chunks[-1]["choices"][0]["finish_reason"] == "tool_calls"
|
||||
|
||||
|
||||
def test_openai_error_conversion() -> None:
|
||||
n = OpenAINormalizer()
|
||||
err_resp = {
|
||||
"error": {
|
||||
"message": "bad request",
|
||||
"type": "invalid_request_error",
|
||||
"code": "bad_request",
|
||||
"param": "messages",
|
||||
}
|
||||
}
|
||||
|
||||
internal = n.error_to_internal(err_resp)
|
||||
assert internal.type == ErrorType.INVALID_REQUEST
|
||||
assert internal.message == "bad request"
|
||||
assert internal.retryable is False
|
||||
|
||||
out = n.error_from_internal(internal)
|
||||
assert out["error"]["type"] == "invalid_request_error"
|
||||
assert out["error"]["message"] == "bad request"
|
||||
@@ -1,406 +0,0 @@
|
||||
"""
|
||||
FormatConverterRegistry 单元测试
|
||||
|
||||
测试覆盖:
|
||||
1. 基本注册和查询
|
||||
2. 能力查询方法
|
||||
3. 严格模式转换
|
||||
4. 流式转换签名适配
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import pytest
|
||||
|
||||
from src.core.api_format.conversion.exceptions import FormatConversionError
|
||||
from src.core.api_format.conversion.registry import FormatConverterRegistry
|
||||
from src.core.api_format.conversion.state import StreamConversionState
|
||||
|
||||
|
||||
# ==================== Mock 转换器 ====================
|
||||
|
||||
|
||||
class MockRequestResponseConverter:
|
||||
"""只支持请求/响应转换的 Mock 转换器"""
|
||||
|
||||
def convert_request(self, request: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return {"converted": True, "original": request}
|
||||
|
||||
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return {"converted": True, "original": response}
|
||||
|
||||
|
||||
class MockNewSignatureStreamConverter:
|
||||
"""使用新签名 (chunk, state) 的流式转换器"""
|
||||
|
||||
def convert_request(self, request: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return request
|
||||
|
||||
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return response
|
||||
|
||||
def convert_stream_chunk(
|
||||
self, chunk: Dict[str, Any], state: StreamConversionState
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""新签名:返回多个事件"""
|
||||
events = []
|
||||
# 模拟一入多出:第一个 chunk 返回 2 个事件
|
||||
if not state.message_started:
|
||||
events.append({"type": "message_start", "model": state.model})
|
||||
state.message_started = True
|
||||
events.append({"type": "content", "data": chunk})
|
||||
return events
|
||||
|
||||
|
||||
class MockUnifiedStreamConverter:
|
||||
"""使用统一签名 (chunk, state) 的流式转换器"""
|
||||
|
||||
def convert_request(self, request: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return request
|
||||
|
||||
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return response
|
||||
|
||||
def convert_stream_chunk(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
state: StreamConversionState,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""统一签名:返回事件列表"""
|
||||
events = []
|
||||
if not state.message_started:
|
||||
events.append({"type": "message_start", "model": state.model, "id": state.message_id})
|
||||
state.message_started = True
|
||||
events.append({"type": "delta", "chunk": chunk})
|
||||
return events
|
||||
|
||||
|
||||
class MockReturnsNoneConverter:
|
||||
"""流式转换返回空列表的转换器"""
|
||||
|
||||
def convert_request(self, request: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return request
|
||||
|
||||
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return response
|
||||
|
||||
def convert_stream_chunk(
|
||||
self,
|
||||
chunk: Dict[str, Any],
|
||||
state: StreamConversionState,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""对于不支持的事件类型返回空列表"""
|
||||
event_type = chunk.get("type")
|
||||
if event_type == "message_start":
|
||||
return [{"choices": [{"delta": {"role": "assistant"}}], "model": state.model}]
|
||||
if event_type == "content_block_delta":
|
||||
return [{"choices": [{"delta": {"content": chunk.get("text", "")}}]}]
|
||||
return []
|
||||
|
||||
|
||||
class MockFailingConverter:
|
||||
"""转换时抛出异常的 Mock 转换器"""
|
||||
|
||||
def convert_request(self, request: Dict[str, Any]) -> Dict[str, Any]:
|
||||
raise ValueError("Request conversion failed")
|
||||
|
||||
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||
raise ValueError("Response conversion failed")
|
||||
|
||||
def convert_stream_chunk(
|
||||
self, chunk: Dict[str, Any], state: Any
|
||||
) -> List[Dict[str, Any]]:
|
||||
raise ValueError("Stream conversion failed")
|
||||
|
||||
|
||||
# ==================== 测试类 ====================
|
||||
|
||||
|
||||
class TestFormatConverterRegistryBasic:
|
||||
"""基本注册和查询测试"""
|
||||
|
||||
def test_register_and_get_converter(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
converter = MockRequestResponseConverter()
|
||||
|
||||
registry.register("OPENAI", "CLAUDE", converter)
|
||||
|
||||
assert registry.get_converter("OPENAI", "CLAUDE") is converter
|
||||
assert registry.get_converter("openai", "claude") is converter # 大小写不敏感
|
||||
|
||||
def test_get_nonexistent_converter_returns_none(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
|
||||
assert registry.get_converter("OPENAI", "CLAUDE") is None
|
||||
|
||||
def test_has_converter(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
converter = MockRequestResponseConverter()
|
||||
|
||||
registry.register("OPENAI", "CLAUDE", converter)
|
||||
|
||||
assert registry.has_converter("OPENAI", "CLAUDE") is True
|
||||
assert registry.has_converter("CLAUDE", "OPENAI") is False
|
||||
|
||||
def test_list_converters(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("OPENAI", "CLAUDE", MockRequestResponseConverter())
|
||||
registry.register("CLAUDE", "OPENAI", MockRequestResponseConverter())
|
||||
|
||||
converters = registry.list_converters()
|
||||
|
||||
assert ("OPENAI", "CLAUDE") in converters
|
||||
assert ("CLAUDE", "OPENAI") in converters
|
||||
assert len(converters) == 2
|
||||
|
||||
|
||||
class TestCapabilityQueries:
|
||||
"""能力查询方法测试"""
|
||||
|
||||
def test_can_convert_request(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("A", "B", MockRequestResponseConverter())
|
||||
|
||||
assert registry.can_convert_request("A", "B") is True
|
||||
assert registry.can_convert_request("B", "A") is False
|
||||
|
||||
def test_can_convert_response(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("A", "B", MockRequestResponseConverter())
|
||||
|
||||
assert registry.can_convert_response("A", "B") is True
|
||||
assert registry.can_convert_response("B", "A") is False
|
||||
|
||||
def test_can_convert_stream_with_convert_stream_chunk(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("A", "B", MockNewSignatureStreamConverter())
|
||||
|
||||
assert registry.can_convert_stream("A", "B") is True
|
||||
assert registry.can_convert_stream("B", "A") is False
|
||||
|
||||
def test_can_convert_stream_unified_signature(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("A", "B", MockUnifiedStreamConverter())
|
||||
|
||||
assert registry.can_convert_stream("A", "B") is True
|
||||
|
||||
def test_can_convert_full_request_response_only(self) -> None:
|
||||
"""只有请求/响应转换器,不支持流式"""
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("OPENAI", "CLAUDE", MockRequestResponseConverter())
|
||||
registry.register("CLAUDE", "OPENAI", MockRequestResponseConverter())
|
||||
|
||||
# 不要求流式:通过
|
||||
assert registry.can_convert_full("OPENAI", "CLAUDE", require_stream=False) is True
|
||||
# 要求流式:失败
|
||||
assert registry.can_convert_full("OPENAI", "CLAUDE", require_stream=True) is False
|
||||
|
||||
def test_can_convert_full_with_stream(self) -> None:
|
||||
"""完整双向转换(含流式)"""
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("OPENAI", "CLAUDE", MockNewSignatureStreamConverter())
|
||||
registry.register("CLAUDE", "OPENAI", MockNewSignatureStreamConverter())
|
||||
|
||||
assert registry.can_convert_full("OPENAI", "CLAUDE", require_stream=True) is True
|
||||
|
||||
def test_can_convert_full_missing_reverse(self) -> None:
|
||||
"""缺少反向转换器"""
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("OPENAI", "CLAUDE", MockNewSignatureStreamConverter())
|
||||
# 没有 CLAUDE -> OPENAI
|
||||
|
||||
assert registry.can_convert_full("OPENAI", "CLAUDE", require_stream=False) is False
|
||||
|
||||
def test_get_supported_targets(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("OPENAI", "CLAUDE", MockRequestResponseConverter())
|
||||
registry.register("OPENAI", "GEMINI", MockRequestResponseConverter())
|
||||
registry.register("CLAUDE", "OPENAI", MockRequestResponseConverter())
|
||||
|
||||
targets = registry.get_supported_targets("OPENAI")
|
||||
|
||||
assert "CLAUDE" in targets
|
||||
assert "GEMINI" in targets
|
||||
assert len(targets) == 2
|
||||
|
||||
|
||||
class TestStrictModeConversion:
|
||||
"""严格模式转换测试"""
|
||||
|
||||
def test_convert_request_strict_success(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("A", "B", MockRequestResponseConverter())
|
||||
|
||||
result = registry.convert_request_strict({"foo": "bar"}, "A", "B")
|
||||
|
||||
assert result["converted"] is True
|
||||
assert result["original"] == {"foo": "bar"}
|
||||
|
||||
def test_convert_request_strict_same_format(self) -> None:
|
||||
"""同格式返回原始请求"""
|
||||
registry = FormatConverterRegistry()
|
||||
|
||||
result = registry.convert_request_strict({"foo": "bar"}, "A", "A")
|
||||
|
||||
assert result == {"foo": "bar"}
|
||||
|
||||
def test_convert_request_strict_no_converter(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
|
||||
with pytest.raises(FormatConversionError) as exc_info:
|
||||
registry.convert_request_strict({"foo": "bar"}, "A", "B")
|
||||
|
||||
assert "未找到转换器" in str(exc_info.value)
|
||||
|
||||
def test_convert_request_strict_converter_fails(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("A", "B", MockFailingConverter())
|
||||
|
||||
with pytest.raises(FormatConversionError) as exc_info:
|
||||
registry.convert_request_strict({"foo": "bar"}, "A", "B")
|
||||
|
||||
assert "Request conversion failed" in str(exc_info.value)
|
||||
|
||||
def test_convert_response_strict_success(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("A", "B", MockRequestResponseConverter())
|
||||
|
||||
result = registry.convert_response_strict({"data": "test"}, "A", "B")
|
||||
|
||||
assert result["converted"] is True
|
||||
|
||||
def test_convert_response_strict_no_converter(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
|
||||
with pytest.raises(FormatConversionError):
|
||||
registry.convert_response_strict({"data": "test"}, "A", "B")
|
||||
|
||||
|
||||
class TestStreamChunkStrictConversion:
|
||||
"""流式转换严格模式测试"""
|
||||
|
||||
def test_stream_chunk_strict_same_format(self) -> None:
|
||||
"""同格式返回原始 chunk 包装在列表中"""
|
||||
registry = FormatConverterRegistry()
|
||||
chunk = {"type": "delta", "text": "hello"}
|
||||
|
||||
result = registry.convert_stream_chunk_strict(chunk, "A", "A")
|
||||
|
||||
assert result == [chunk]
|
||||
|
||||
def test_stream_chunk_strict_no_converter(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
|
||||
with pytest.raises(FormatConversionError) as exc_info:
|
||||
registry.convert_stream_chunk_strict({"type": "delta"}, "A", "B")
|
||||
|
||||
assert "未找到转换器" in str(exc_info.value)
|
||||
|
||||
def test_stream_chunk_strict_new_signature(self) -> None:
|
||||
"""测试新签名 (chunk, state) 转换器"""
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("A", "B", MockNewSignatureStreamConverter())
|
||||
|
||||
state = StreamConversionState(model="test-model", message_id="msg_123")
|
||||
chunk = {"content": "hello"}
|
||||
|
||||
# 第一次调用:应返回 2 个事件(message_start + content)
|
||||
result = registry.convert_stream_chunk_strict(chunk, "A", "B", state)
|
||||
|
||||
assert len(result) == 2
|
||||
assert result[0]["type"] == "message_start"
|
||||
assert result[1]["type"] == "content"
|
||||
assert state.message_started is True
|
||||
|
||||
# 第二次调用:应返回 1 个事件(只有 content)
|
||||
result2 = registry.convert_stream_chunk_strict({"content": "world"}, "A", "B", state)
|
||||
|
||||
assert len(result2) == 1
|
||||
assert result2[0]["type"] == "content"
|
||||
|
||||
def test_stream_chunk_strict_unified_signature(self) -> None:
|
||||
"""测试统一签名 (chunk, state) 转换器"""
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("A", "B", MockUnifiedStreamConverter())
|
||||
|
||||
state = StreamConversionState(
|
||||
model="gpt-4", message_id="chatcmpl_123", message_started=False
|
||||
)
|
||||
chunk = {"text": "hello"}
|
||||
|
||||
result = registry.convert_stream_chunk_strict(chunk, "A", "B", state)
|
||||
|
||||
assert len(result) == 2
|
||||
assert result[0]["type"] == "message_start"
|
||||
assert result[0]["model"] == "gpt-4"
|
||||
assert result[0]["id"] == "chatcmpl_123"
|
||||
assert state.message_started is True
|
||||
|
||||
def test_stream_chunk_strict_returns_empty_list(self) -> None:
|
||||
"""测试转换器返回空列表"""
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("CLAUDE", "OPENAI", MockReturnsNoneConverter())
|
||||
|
||||
state = StreamConversionState(model="gpt-4", message_id="msg_123")
|
||||
event = {"type": "message_start"}
|
||||
|
||||
result = registry.convert_stream_chunk_strict(event, "CLAUDE", "OPENAI", state)
|
||||
|
||||
assert len(result) == 1
|
||||
assert "choices" in result[0]
|
||||
assert result[0]["model"] == "gpt-4"
|
||||
|
||||
def test_stream_chunk_strict_unknown_event_returns_empty(self) -> None:
|
||||
"""不支持的事件类型应返回空列表"""
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("CLAUDE", "OPENAI", MockReturnsNoneConverter())
|
||||
|
||||
state = StreamConversionState(model="gpt-4", message_id="msg_123")
|
||||
event = {"type": "unknown_event"} # 不支持的事件类型
|
||||
|
||||
result = registry.convert_stream_chunk_strict(event, "CLAUDE", "OPENAI", state)
|
||||
|
||||
assert result == []
|
||||
|
||||
def test_stream_chunk_strict_converter_fails(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("A", "B", MockFailingConverter())
|
||||
|
||||
state = StreamConversionState(model="test", message_id="123")
|
||||
|
||||
with pytest.raises(FormatConversionError) as exc_info:
|
||||
registry.convert_stream_chunk_strict({"data": "test"}, "A", "B", state)
|
||||
|
||||
assert "流式块转换失败" in str(exc_info.value)
|
||||
|
||||
|
||||
class TestNonStrictConversion:
|
||||
"""非严格模式转换测试(失败时返回原始数据)"""
|
||||
|
||||
def test_convert_request_fallback_on_failure(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("A", "B", MockFailingConverter())
|
||||
|
||||
original = {"foo": "bar"}
|
||||
result = registry.convert_request(original, "A", "B")
|
||||
|
||||
# 非严格模式:失败时返回原始请求
|
||||
assert result == original
|
||||
|
||||
def test_convert_response_fallback_on_failure(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("A", "B", MockFailingConverter())
|
||||
|
||||
original = {"data": "test"}
|
||||
result = registry.convert_response(original, "A", "B")
|
||||
|
||||
assert result == original
|
||||
|
||||
def test_convert_stream_chunk_fallback_on_failure(self) -> None:
|
||||
registry = FormatConverterRegistry()
|
||||
registry.register("A", "B", MockFailingConverter())
|
||||
|
||||
original = {"chunk": "data"}
|
||||
result = registry.convert_stream_chunk(original, "A", "B")
|
||||
|
||||
assert result == [original]
|
||||
105
tests/core/api_format/conversion/test_registry_canonical.py
Normal file
105
tests/core/api_format/conversion/test_registry_canonical.py
Normal file
@@ -0,0 +1,105 @@
|
||||
"""
|
||||
Canonical Registry 单元测试
|
||||
|
||||
覆盖重点:
|
||||
- request/response/stream 的基本两段式转换(source -> internal -> target)
|
||||
- 严格模式下的可用性(已注册格式可转换)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List, cast
|
||||
|
||||
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer
|
||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
|
||||
|
||||
def _make_registry() -> FormatConversionRegistry:
|
||||
reg = FormatConversionRegistry()
|
||||
reg.register(OpenAINormalizer())
|
||||
reg.register(ClaudeNormalizer())
|
||||
reg.register(GeminiNormalizer())
|
||||
return reg
|
||||
|
||||
|
||||
def _first_openai_choice_message(resp: Dict[str, Any]) -> Dict[str, Any]:
|
||||
choices = resp.get("choices") or []
|
||||
assert isinstance(choices, list) and choices
|
||||
c0 = choices[0]
|
||||
assert isinstance(c0, dict)
|
||||
msg = c0.get("message")
|
||||
assert isinstance(msg, dict)
|
||||
return cast(Dict[str, Any], msg)
|
||||
|
||||
|
||||
def test_registry_canonical_can_convert_full_stream() -> None:
|
||||
reg = _make_registry()
|
||||
assert reg.can_convert_full("OPENAI", "CLAUDE", require_stream=True) is True
|
||||
assert reg.can_convert_full("OPENAI", "GEMINI", require_stream=True) is True
|
||||
assert reg.can_convert_full("CLAUDE", "GEMINI", require_stream=True) is True
|
||||
|
||||
|
||||
def test_registry_canonical_request_openai_to_claude() -> None:
|
||||
reg = _make_registry()
|
||||
|
||||
openai_req = {
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [
|
||||
{"role": "system", "content": "sys"},
|
||||
{"role": "developer", "content": "dev"},
|
||||
{"role": "user", "content": "hi"},
|
||||
],
|
||||
"max_tokens": 12,
|
||||
"temperature": 0.2,
|
||||
"stream": True,
|
||||
}
|
||||
|
||||
claude_req = reg.convert_request(openai_req, "OPENAI", "CLAUDE")
|
||||
assert claude_req["model"] == "gpt-4o-mini"
|
||||
assert claude_req["system"] == "sys\n\ndev"
|
||||
assert claude_req["stream"] is True
|
||||
assert isinstance(claude_req.get("messages"), list)
|
||||
assert claude_req["messages"][0]["role"] == "user"
|
||||
assert claude_req["messages"][0]["content"] == "hi"
|
||||
|
||||
|
||||
def test_registry_canonical_response_claude_to_openai() -> None:
|
||||
reg = _make_registry()
|
||||
|
||||
claude_resp = {
|
||||
"id": "msg_1",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-3-5-sonnet-latest",
|
||||
"content": [{"type": "text", "text": "hello"}],
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 5, "output_tokens": 7},
|
||||
}
|
||||
|
||||
openai_resp = reg.convert_response(claude_resp, "CLAUDE", "OPENAI")
|
||||
assert openai_resp["object"] == "chat.completion"
|
||||
msg = _first_openai_choice_message(openai_resp)
|
||||
assert msg["role"] == "assistant"
|
||||
assert msg["content"] == "hello"
|
||||
|
||||
|
||||
def test_registry_canonical_stream_openai_to_claude() -> None:
|
||||
reg = _make_registry()
|
||||
|
||||
chunk = {
|
||||
"id": "chatcmpl_1",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [{"index": 0, "delta": {"role": "assistant", "content": "hi"}, "finish_reason": None}],
|
||||
}
|
||||
|
||||
state = StreamState()
|
||||
out_events = reg.convert_stream_chunk(chunk, "OPENAI", "CLAUDE", state=state)
|
||||
assert isinstance(out_events, list) and out_events
|
||||
|
||||
types = [cast(Dict[str, Any], e).get("type") for e in cast(List[Dict[str, Any]], out_events)]
|
||||
assert types[:3] == ["message_start", "content_block_start", "content_block_delta"]
|
||||
@@ -1,7 +1,8 @@
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from src.core.api_format import APIFormat, register_all_converters
|
||||
from src.core.api_format import APIFormat
|
||||
from src.core.api_format.conversion import register_default_normalizers
|
||||
from src.services.cache.aware_scheduler import CacheAwareScheduler
|
||||
|
||||
|
||||
@@ -26,7 +27,7 @@ def _mock_endpoint(api_format: str, config: dict | None = None) -> MagicMock:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_candidates_blocks_cross_format_when_global_switch_off() -> None:
|
||||
register_all_converters()
|
||||
register_default_normalizers()
|
||||
|
||||
scheduler = CacheAwareScheduler()
|
||||
scheduler._check_model_support = AsyncMock(return_value=(True, None, None, {"m"})) # type: ignore[attr-defined]
|
||||
@@ -56,7 +57,7 @@ async def test_build_candidates_blocks_cross_format_when_global_switch_off() ->
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_candidates_includes_cross_format_when_enabled() -> None:
|
||||
register_all_converters()
|
||||
register_default_normalizers()
|
||||
|
||||
scheduler = CacheAwareScheduler()
|
||||
scheduler._check_model_support = AsyncMock(return_value=(True, None, None, {"m"})) # type: ignore[attr-defined]
|
||||
@@ -88,7 +89,7 @@ async def test_build_candidates_includes_cross_format_when_enabled() -> None:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_exact_matches_rank_before_convertible() -> None:
|
||||
register_all_converters()
|
||||
register_default_normalizers()
|
||||
|
||||
scheduler = CacheAwareScheduler()
|
||||
scheduler._check_model_support = AsyncMock(return_value=(True, None, None, {"m"})) # type: ignore[attr-defined]
|
||||
|
||||
Reference in New Issue
Block a user