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:
fawney19
2026-01-27 02:17:18 +08:00
parent 6cf251b19a
commit add477fece
84 changed files with 9169 additions and 3932 deletions

View File

@@ -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()

View File

@@ -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=Falsedata_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"

View File

@@ -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 可能在首个 chunkmessage_start或最后一个 chunkmessage_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):

View File

@@ -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)
# ==============================================================================

View File

@@ -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:
"""

View File

@@ -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,

View File

@@ -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:
"""

View File

@@ -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]):

View File

@@ -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 格式

View File

@@ -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]:

View File

@@ -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",
]

View File

@@ -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(

View File

@@ -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 格式

View File

@@ -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]:

View File

@@ -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)