mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +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:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user