mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
refactor: 重构 API 格式相关代码到统一的 api_format 模块
- 新增 src/core/api_format/ 模块,整合所有 API 格式相关功能 - enums.py: APIFormat 枚举定义 - metadata.py: API 格式元数据注册 - headers.py: 请求头处理逻辑 - detection.py: API 格式检测 - conversion/: 格式转换器(Claude/OpenAI/Gemini 互转) - 删除分散在各处的旧实现 - src/core/api_format_metadata.py - src/core/headers.py - src/api/handlers/*/converter.py - src/api/handlers/base/format_converter_registry.py - 更新所有引用路径,使用新的统一模块 - 新增 cli_handler_base 中的格式转换支持
This commit is contained in:
@@ -402,7 +402,7 @@ class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 按 APIFormat 枚举定义的顺序排序 Endpoints
|
# 按 APIFormat 枚举定义的顺序排序 Endpoints
|
||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
|
|
||||||
format_order = {fmt.value: i for i, fmt in enumerate(APIFormat)}
|
format_order = {fmt.value: i for i, fmt in enumerate(APIFormat)}
|
||||||
endpoint_infos.sort(key=lambda e: format_order.get(e.api_format, 999))
|
endpoint_infos.sort(key=lambda e: format_order.get(e.api_format, 999))
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ from sqlalchemy.orm import Session, joinedload
|
|||||||
|
|
||||||
from src.config.constants import TimeoutDefaults
|
from src.config.constants import TimeoutDefaults
|
||||||
from src.core.crypto import crypto_service
|
from src.core.crypto import crypto_service
|
||||||
from src.core.headers import get_extra_headers_from_endpoint
|
from src.core.api_format import get_extra_headers_from_endpoint
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.database.database import get_db
|
from src.database.database import get_db
|
||||||
from src.models.database import Provider, ProviderEndpoint, User
|
from src.models.database import Provider, ProviderEndpoint, User
|
||||||
|
|||||||
@@ -711,8 +711,7 @@ class AdminTriggerCleanupAdapter(AdminApiAdapter):
|
|||||||
class AdminGetApiFormatsAdapter(AdminApiAdapter):
|
class AdminGetApiFormatsAdapter(AdminApiAdapter):
|
||||||
async def handle(self, context): # type: ignore[override]
|
async def handle(self, context): # type: ignore[override]
|
||||||
"""获取所有可用的API格式"""
|
"""获取所有可用的API格式"""
|
||||||
from src.core.api_format_metadata import API_FORMAT_DEFINITIONS
|
from src.core.api_format import API_FORMAT_DEFINITIONS, APIFormat
|
||||||
from src.core.enums import APIFormat
|
|
||||||
|
|
||||||
_ = context # 参数保留以符合接口规范
|
_ = context # 参数保留以符合接口规范
|
||||||
|
|
||||||
|
|||||||
@@ -35,8 +35,7 @@ from fastapi.responses import JSONResponse, StreamingResponse
|
|||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from src.clients.redis_client import get_redis_client_sync
|
from src.clients.redis_client import get_redis_client_sync
|
||||||
from src.core.api_format_metadata import resolve_api_format
|
from src.core.api_format import APIFormat, resolve_api_format
|
||||||
from src.core.enums import APIFormat
|
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.services.orchestration.fallback_orchestrator import FallbackOrchestrator
|
from src.services.orchestration.fallback_orchestrator import FallbackOrchestrator
|
||||||
from src.services.provider.format import normalize_api_format
|
from src.services.provider.format import normalize_api_format
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ from fastapi.responses import JSONResponse
|
|||||||
from src.api.base.adapter import ApiAdapter, ApiMode
|
from src.api.base.adapter import ApiAdapter, ApiMode
|
||||||
from src.api.base.context import ApiRequestContext
|
from src.api.base.context import ApiRequestContext
|
||||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
from src.core.exceptions import (
|
from src.core.exceptions import (
|
||||||
InvalidRequestException,
|
InvalidRequestException,
|
||||||
ModelNotSupportedException,
|
ModelNotSupportedException,
|
||||||
@@ -40,7 +40,7 @@ from src.core.exceptions import (
|
|||||||
QuotaExceededException,
|
QuotaExceededException,
|
||||||
UpstreamClientException,
|
UpstreamClientException,
|
||||||
)
|
)
|
||||||
from src.core.headers import (
|
from src.core.api_format import (
|
||||||
build_adapter_base_headers,
|
build_adapter_base_headers,
|
||||||
build_adapter_headers,
|
build_adapter_headers,
|
||||||
extract_client_api_key,
|
extract_client_api_key,
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ from fastapi.responses import JSONResponse
|
|||||||
from src.api.base.adapter import ApiAdapter, ApiMode
|
from src.api.base.adapter import ApiAdapter, ApiMode
|
||||||
from src.api.base.context import ApiRequestContext
|
from src.api.base.context import ApiRequestContext
|
||||||
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
|
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
|
||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
from src.core.exceptions import (
|
from src.core.exceptions import (
|
||||||
InvalidRequestException,
|
InvalidRequestException,
|
||||||
ModelNotSupportedException,
|
ModelNotSupportedException,
|
||||||
@@ -38,7 +38,7 @@ from src.core.exceptions import (
|
|||||||
QuotaExceededException,
|
QuotaExceededException,
|
||||||
UpstreamClientException,
|
UpstreamClientException,
|
||||||
)
|
)
|
||||||
from src.core.headers import (
|
from src.core.api_format import (
|
||||||
build_adapter_base_headers,
|
build_adapter_base_headers,
|
||||||
build_adapter_headers,
|
build_adapter_headers,
|
||||||
extract_client_api_key,
|
extract_client_api_key,
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ from typing import (
|
|||||||
AsyncGenerator,
|
AsyncGenerator,
|
||||||
Callable,
|
Callable,
|
||||||
Dict,
|
Dict,
|
||||||
|
List,
|
||||||
Optional,
|
Optional,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -671,10 +672,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
# 格式转换或直接透传
|
# 格式转换或直接透传
|
||||||
if needs_conversion:
|
if needs_conversion:
|
||||||
converted_line = self._convert_sse_line(ctx, line, events)
|
converted_lines = self._convert_sse_line(ctx, line, events)
|
||||||
if converted_line:
|
for converted_line in converted_lines:
|
||||||
self._mark_first_output(ctx, output_state)
|
if converted_line:
|
||||||
yield (converted_line + "\n").encode("utf-8")
|
self._mark_first_output(ctx, output_state)
|
||||||
|
yield (converted_line + "\n").encode("utf-8")
|
||||||
else:
|
else:
|
||||||
self._mark_first_output(ctx, output_state)
|
self._mark_first_output(ctx, output_state)
|
||||||
yield (line + "\n").encode("utf-8")
|
yield (line + "\n").encode("utf-8")
|
||||||
@@ -999,10 +1001,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
# 格式转换或直接透传
|
# 格式转换或直接透传
|
||||||
if needs_conversion:
|
if needs_conversion:
|
||||||
converted_line = self._convert_sse_line(ctx, line, events)
|
converted_lines = self._convert_sse_line(ctx, line, events)
|
||||||
if converted_line:
|
for converted_line in converted_lines:
|
||||||
self._mark_first_output(ctx, output_state)
|
if converted_line:
|
||||||
yield (converted_line + "\n").encode("utf-8")
|
self._mark_first_output(ctx, output_state)
|
||||||
|
yield (converted_line + "\n").encode("utf-8")
|
||||||
else:
|
else:
|
||||||
self._mark_first_output(ctx, output_state)
|
self._mark_first_output(ctx, output_state)
|
||||||
yield (line + "\n").encode("utf-8")
|
yield (line + "\n").encode("utf-8")
|
||||||
@@ -1073,10 +1076,11 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
|
|
||||||
# 格式转换或直接透传
|
# 格式转换或直接透传
|
||||||
if needs_conversion:
|
if needs_conversion:
|
||||||
converted_line = self._convert_sse_line(ctx, line, events)
|
converted_lines = self._convert_sse_line(ctx, line, events)
|
||||||
if converted_line:
|
for converted_line in converted_lines:
|
||||||
self._mark_first_output(ctx, output_state)
|
if converted_line:
|
||||||
yield (converted_line + "\n").encode("utf-8")
|
self._mark_first_output(ctx, output_state)
|
||||||
|
yield (converted_line + "\n").encode("utf-8")
|
||||||
else:
|
else:
|
||||||
self._mark_first_output(ctx, output_state)
|
self._mark_first_output(ctx, output_state)
|
||||||
yield (line + "\n").encode("utf-8")
|
yield (line + "\n").encode("utf-8")
|
||||||
@@ -1760,7 +1764,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
and api_format
|
and api_format
|
||||||
and provider_api_format.upper() != api_format.upper()
|
and provider_api_format.upper() != api_format.upper()
|
||||||
):
|
):
|
||||||
from src.api.handlers.base.format_converter_registry import converter_registry
|
from src.core.api_format import converter_registry
|
||||||
|
|
||||||
try:
|
try:
|
||||||
response_json = converter_registry.convert_response(
|
response_json = converter_registry.convert_response(
|
||||||
@@ -1923,28 +1927,32 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
self,
|
self,
|
||||||
ctx: StreamContext,
|
ctx: StreamContext,
|
||||||
line: str,
|
line: str,
|
||||||
events: list,
|
events: list, # noqa: ARG002 - 预留给上下文感知转换
|
||||||
) -> Optional[str]:
|
) -> List[str]:
|
||||||
"""
|
"""
|
||||||
将 SSE 行从 Provider 格式转换为客户端格式
|
将 SSE 行从 Provider 格式转换为客户端格式
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
ctx: 流上下文
|
ctx: 流上下文
|
||||||
line: 原始 SSE 行
|
line: 原始 SSE 行
|
||||||
events: 解析后的事件列表
|
events: 当前累积的事件列表(预留参数,用于未来上下文感知转换如合并相邻事件)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
转换后的 SSE 行,如果无法转换则返回 None
|
转换后的 SSE 行列表(一入多出),空列表表示跳过该行
|
||||||
"""
|
"""
|
||||||
from src.api.handlers.base.format_converter_registry import converter_registry
|
from src.core.api_format import (
|
||||||
|
GeminiStreamConversionState,
|
||||||
|
StreamConversionState,
|
||||||
|
converter_registry,
|
||||||
|
)
|
||||||
|
|
||||||
# 如果是空行或特殊控制行,直接返回
|
# 如果是空行或特殊控制行,直接返回
|
||||||
if not line or line.strip() == "" or line == "data: [DONE]":
|
if not line or line.strip() == "" or line == "data: [DONE]":
|
||||||
return line
|
return [line] if line else []
|
||||||
|
|
||||||
# 如果不是 data 行,直接透传
|
# 如果不是 data 行,直接透传
|
||||||
if not line.startswith("data:"):
|
if not line.startswith("data:"):
|
||||||
return line
|
return [line]
|
||||||
|
|
||||||
# 提取 data 内容
|
# 提取 data 内容
|
||||||
data_content = line[5:].strip() # 去掉 "data:" 前缀
|
data_content = line[5:].strip() # 去掉 "data:" 前缀
|
||||||
@@ -1954,17 +1962,40 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
|||||||
data_obj = json.loads(data_content)
|
data_obj = json.loads(data_content)
|
||||||
except json.JSONDecodeError:
|
except json.JSONDecodeError:
|
||||||
# 无法解析,直接透传
|
# 无法解析,直接透传
|
||||||
return line
|
return [line]
|
||||||
|
|
||||||
# 使用注册表进行格式转换
|
# 类型断言:当 needs_conversion=True 调用此方法时,格式字段必定有值
|
||||||
|
provider_format = ctx.provider_api_format or ""
|
||||||
|
client_format = ctx.client_api_format or ""
|
||||||
|
|
||||||
|
# 初始化流式转换状态(首次调用时,根据 Provider 格式选择状态类)
|
||||||
|
if ctx.stream_conversion_state is None:
|
||||||
|
if provider_format.upper() == "GEMINI":
|
||||||
|
ctx.stream_conversion_state = GeminiStreamConversionState(
|
||||||
|
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 "",
|
||||||
|
)
|
||||||
|
|
||||||
|
# 使用注册表进行格式转换(严格模式,返回 List[Dict])
|
||||||
try:
|
try:
|
||||||
converted_obj = converter_registry.convert_stream_chunk(
|
converted_events = converter_registry.convert_stream_chunk_strict(
|
||||||
data_obj,
|
data_obj,
|
||||||
ctx.provider_api_format,
|
provider_format,
|
||||||
ctx.client_api_format,
|
client_format,
|
||||||
|
state=ctx.stream_conversion_state,
|
||||||
)
|
)
|
||||||
# 重新构建 SSE 行
|
|
||||||
return f"data: {json.dumps(converted_obj, ensure_ascii=False)}"
|
# 转换为 SSE 行列表
|
||||||
|
result = []
|
||||||
|
for evt in converted_events:
|
||||||
|
result.append(f"data: {json.dumps(evt, ensure_ascii=False)}")
|
||||||
|
return result
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"格式转换失败,透传原始数据: {e}")
|
logger.warning(f"格式转换失败,透传原始数据: {e}")
|
||||||
return line
|
return [line]
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ from collections import defaultdict
|
|||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.core.headers import CORE_REDACT_HEADERS, merge_headers_with_protection, redact_headers_for_log
|
from src.core.api_format import CORE_REDACT_HEADERS, merge_headers_with_protection, redact_headers_for_log
|
||||||
from src.utils.ssl_utils import get_ssl_context
|
from src.utils.ssl_utils import get_ssl_context
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,279 +0,0 @@
|
|||||||
"""
|
|
||||||
格式转换器注册表
|
|
||||||
|
|
||||||
自动管理不同 API 格式之间的转换器,支持:
|
|
||||||
- 请求转换:客户端格式 → Provider 格式
|
|
||||||
- 响应转换:Provider 格式 → 客户端格式
|
|
||||||
|
|
||||||
使用方法:
|
|
||||||
1. 实现 Converter 类(需要有 convert_request 和/或 convert_response 方法)
|
|
||||||
2. 调用 registry.register() 注册转换器
|
|
||||||
3. 在 Handler 中调用 registry.convert_request/convert_response
|
|
||||||
|
|
||||||
示例:
|
|
||||||
from src.api.handlers.base.format_converter_registry import converter_registry
|
|
||||||
|
|
||||||
# 注册转换器
|
|
||||||
converter_registry.register("CLAUDE", "GEMINI", ClaudeToGeminiConverter())
|
|
||||||
converter_registry.register("GEMINI", "CLAUDE", GeminiToClaudeConverter())
|
|
||||||
|
|
||||||
# 使用转换器
|
|
||||||
gemini_request = converter_registry.convert_request(claude_request, "CLAUDE", "GEMINI")
|
|
||||||
claude_response = converter_registry.convert_response(gemini_response, "GEMINI", "CLAUDE")
|
|
||||||
"""
|
|
||||||
|
|
||||||
from typing import Any, Dict, Optional, Protocol, Tuple
|
|
||||||
|
|
||||||
from src.core.logger import logger
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class RequestConverter(Protocol):
|
|
||||||
"""请求转换器协议"""
|
|
||||||
|
|
||||||
def convert_request(self, request: Dict[str, Any]) -> Dict[str, Any]: ...
|
|
||||||
|
|
||||||
|
|
||||||
class ResponseConverter(Protocol):
|
|
||||||
"""响应转换器协议"""
|
|
||||||
|
|
||||||
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]: ...
|
|
||||||
|
|
||||||
|
|
||||||
class StreamChunkConverter(Protocol):
|
|
||||||
"""流式响应块转换器协议"""
|
|
||||||
|
|
||||||
def convert_stream_chunk(self, chunk: Dict[str, Any]) -> Dict[str, Any]: ...
|
|
||||||
|
|
||||||
|
|
||||||
class FormatConverterRegistry:
|
|
||||||
"""
|
|
||||||
格式转换器注册表
|
|
||||||
|
|
||||||
管理不同 API 格式之间的双向转换器
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self):
|
|
||||||
# key: (source_format, target_format), value: converter instance
|
|
||||||
self._converters: Dict[Tuple[str, str], Any] = {}
|
|
||||||
|
|
||||||
def register(
|
|
||||||
self,
|
|
||||||
source_format: str,
|
|
||||||
target_format: str,
|
|
||||||
converter: Any,
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
注册格式转换器
|
|
||||||
|
|
||||||
Args:
|
|
||||||
source_format: 源格式(如 "CLAUDE", "OPENAI", "GEMINI")
|
|
||||||
target_format: 目标格式
|
|
||||||
converter: 转换器实例(需要有 convert_request/convert_response 方法)
|
|
||||||
"""
|
|
||||||
key = (source_format.upper(), target_format.upper())
|
|
||||||
self._converters[key] = converter
|
|
||||||
logger.info(f"[ConverterRegistry] 注册转换器: {source_format} -> {target_format}")
|
|
||||||
|
|
||||||
def get_converter(
|
|
||||||
self,
|
|
||||||
source_format: str,
|
|
||||||
target_format: str,
|
|
||||||
) -> Optional[Any]:
|
|
||||||
"""
|
|
||||||
获取转换器
|
|
||||||
|
|
||||||
Args:
|
|
||||||
source_format: 源格式
|
|
||||||
target_format: 目标格式
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
转换器实例,如果不存在返回 None
|
|
||||||
"""
|
|
||||||
key = (source_format.upper(), target_format.upper())
|
|
||||||
return self._converters.get(key)
|
|
||||||
|
|
||||||
def has_converter(
|
|
||||||
self,
|
|
||||||
source_format: str,
|
|
||||||
target_format: str,
|
|
||||||
) -> bool:
|
|
||||||
"""检查是否存在转换器"""
|
|
||||||
key = (source_format.upper(), target_format.upper())
|
|
||||||
return key in self._converters
|
|
||||||
|
|
||||||
def convert_request(
|
|
||||||
self,
|
|
||||||
request: Dict[str, Any],
|
|
||||||
source_format: str,
|
|
||||||
target_format: str,
|
|
||||||
) -> Dict[str, Any]:
|
|
||||||
"""
|
|
||||||
转换请求
|
|
||||||
|
|
||||||
Args:
|
|
||||||
request: 原始请求字典
|
|
||||||
source_format: 源格式(客户端格式)
|
|
||||||
target_format: 目标格式(Provider 格式)
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
转换后的请求字典,如果无需转换或没有转换器则返回原始请求
|
|
||||||
"""
|
|
||||||
# 同格式无需转换
|
|
||||||
if source_format.upper() == target_format.upper():
|
|
||||||
return request
|
|
||||||
|
|
||||||
converter = self.get_converter(source_format, target_format)
|
|
||||||
if converter is None:
|
|
||||||
logger.warning(f"[ConverterRegistry] 未找到请求转换器: {source_format} -> {target_format},返回原始请求")
|
|
||||||
return request
|
|
||||||
|
|
||||||
if not hasattr(converter, "convert_request"):
|
|
||||||
logger.warning(f"[ConverterRegistry] 转换器缺少 convert_request 方法: {source_format} -> {target_format}")
|
|
||||||
return request
|
|
||||||
|
|
||||||
try:
|
|
||||||
converted = converter.convert_request(request)
|
|
||||||
logger.debug(f"[ConverterRegistry] 请求转换成功: {source_format} -> {target_format}")
|
|
||||||
return converted
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"[ConverterRegistry] 请求转换失败: {source_format} -> {target_format}: {e}")
|
|
||||||
return request
|
|
||||||
|
|
||||||
def convert_response(
|
|
||||||
self,
|
|
||||||
response: Dict[str, Any],
|
|
||||||
source_format: str,
|
|
||||||
target_format: str,
|
|
||||||
) -> Dict[str, Any]:
|
|
||||||
"""
|
|
||||||
转换响应
|
|
||||||
|
|
||||||
Args:
|
|
||||||
response: 原始响应字典
|
|
||||||
source_format: 源格式(Provider 格式)
|
|
||||||
target_format: 目标格式(客户端格式)
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
转换后的响应字典,如果无需转换或没有转换器则返回原始响应
|
|
||||||
"""
|
|
||||||
# 同格式无需转换
|
|
||||||
if source_format.upper() == target_format.upper():
|
|
||||||
return response
|
|
||||||
|
|
||||||
converter = self.get_converter(source_format, target_format)
|
|
||||||
if converter is None:
|
|
||||||
logger.warning(f"[ConverterRegistry] 未找到响应转换器: {source_format} -> {target_format},返回原始响应")
|
|
||||||
return response
|
|
||||||
|
|
||||||
if not hasattr(converter, "convert_response"):
|
|
||||||
logger.warning(f"[ConverterRegistry] 转换器缺少 convert_response 方法: {source_format} -> {target_format}")
|
|
||||||
return response
|
|
||||||
|
|
||||||
try:
|
|
||||||
converted = converter.convert_response(response)
|
|
||||||
logger.debug(f"[ConverterRegistry] 响应转换成功: {source_format} -> {target_format}")
|
|
||||||
return converted
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"[ConverterRegistry] 响应转换失败: {source_format} -> {target_format}: {e}")
|
|
||||||
return response
|
|
||||||
|
|
||||||
def convert_stream_chunk(
|
|
||||||
self,
|
|
||||||
chunk: Dict[str, Any],
|
|
||||||
source_format: str,
|
|
||||||
target_format: str,
|
|
||||||
) -> Dict[str, Any]:
|
|
||||||
"""
|
|
||||||
转换流式响应块
|
|
||||||
|
|
||||||
Args:
|
|
||||||
chunk: 原始流式响应块
|
|
||||||
source_format: 源格式(Provider 格式)
|
|
||||||
target_format: 目标格式(客户端格式)
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
转换后的流式响应块
|
|
||||||
"""
|
|
||||||
# 同格式无需转换
|
|
||||||
if source_format.upper() == target_format.upper():
|
|
||||||
return chunk
|
|
||||||
|
|
||||||
converter = self.get_converter(source_format, target_format)
|
|
||||||
if converter is None:
|
|
||||||
return chunk
|
|
||||||
|
|
||||||
# 优先使用专门的流式转换方法
|
|
||||||
if hasattr(converter, "convert_stream_chunk"):
|
|
||||||
try:
|
|
||||||
return converter.convert_stream_chunk(chunk)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"[ConverterRegistry] 流式块转换失败: {source_format} -> {target_format}: {e}")
|
|
||||||
return chunk
|
|
||||||
|
|
||||||
# 降级到普通响应转换
|
|
||||||
if hasattr(converter, "convert_response"):
|
|
||||||
try:
|
|
||||||
return converter.convert_response(chunk)
|
|
||||||
except Exception:
|
|
||||||
return chunk
|
|
||||||
|
|
||||||
return chunk
|
|
||||||
|
|
||||||
def list_converters(self) -> list[Tuple[str, str]]:
|
|
||||||
"""列出所有已注册的转换器"""
|
|
||||||
return list(self._converters.keys())
|
|
||||||
|
|
||||||
|
|
||||||
# 全局单例
|
|
||||||
converter_registry = FormatConverterRegistry()
|
|
||||||
|
|
||||||
|
|
||||||
def register_all_converters():
|
|
||||||
"""
|
|
||||||
注册所有内置的格式转换器
|
|
||||||
|
|
||||||
在应用启动时调用此函数
|
|
||||||
"""
|
|
||||||
# Claude <-> OpenAI
|
|
||||||
try:
|
|
||||||
from src.api.handlers.claude.converter import OpenAIToClaudeConverter
|
|
||||||
from src.api.handlers.openai.converter import ClaudeToOpenAIConverter
|
|
||||||
|
|
||||||
converter_registry.register("OPENAI", "CLAUDE", OpenAIToClaudeConverter())
|
|
||||||
converter_registry.register("CLAUDE", "OPENAI", ClaudeToOpenAIConverter())
|
|
||||||
except ImportError as e:
|
|
||||||
logger.warning(f"[ConverterRegistry] 无法加载 Claude/OpenAI 转换器: {e}")
|
|
||||||
|
|
||||||
# Claude <-> Gemini
|
|
||||||
try:
|
|
||||||
from src.api.handlers.gemini.converter import (
|
|
||||||
ClaudeToGeminiConverter,
|
|
||||||
GeminiToClaudeConverter,
|
|
||||||
)
|
|
||||||
|
|
||||||
converter_registry.register("CLAUDE", "GEMINI", ClaudeToGeminiConverter())
|
|
||||||
converter_registry.register("GEMINI", "CLAUDE", GeminiToClaudeConverter())
|
|
||||||
except ImportError as e:
|
|
||||||
logger.warning(f"[ConverterRegistry] 无法加载 Claude/Gemini 转换器: {e}")
|
|
||||||
|
|
||||||
# OpenAI <-> Gemini
|
|
||||||
try:
|
|
||||||
from src.api.handlers.gemini.converter import (
|
|
||||||
GeminiToOpenAIConverter,
|
|
||||||
OpenAIToGeminiConverter,
|
|
||||||
)
|
|
||||||
|
|
||||||
converter_registry.register("OPENAI", "GEMINI", OpenAIToGeminiConverter())
|
|
||||||
converter_registry.register("GEMINI", "OPENAI", GeminiToOpenAIConverter())
|
|
||||||
except ImportError as e:
|
|
||||||
logger.warning(f"[ConverterRegistry] 无法加载 OpenAI/Gemini 转换器: {e}")
|
|
||||||
|
|
||||||
logger.info(f"[ConverterRegistry] 已注册 {len(converter_registry.list_converters())} 个格式转换器")
|
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"FormatConverterRegistry",
|
|
||||||
"converter_registry",
|
|
||||||
"register_all_converters",
|
|
||||||
]
|
|
||||||
@@ -16,6 +16,9 @@ from src.api.handlers.base.response_parser import (
|
|||||||
)
|
)
|
||||||
from src.api.handlers.base.utils import extract_cache_creation_tokens
|
from src.api.handlers.base.utils import extract_cache_creation_tokens
|
||||||
|
|
||||||
|
# is_cli_format 权威定义在 core 层
|
||||||
|
from src.core.api_format import is_cli_format
|
||||||
|
|
||||||
|
|
||||||
def _check_nested_error(response: Dict[str, Any]) -> Tuple[bool, Optional[Dict[str, Any]]]:
|
def _check_nested_error(response: Dict[str, Any]) -> Tuple[bool, Optional[Dict[str, Any]]]:
|
||||||
"""
|
"""
|
||||||
@@ -571,11 +574,6 @@ def get_parser_for_format(format_id: str) -> ResponseParser:
|
|||||||
return _PARSERS[format_id]()
|
return _PARSERS[format_id]()
|
||||||
|
|
||||||
|
|
||||||
def is_cli_format(format_id: str) -> bool:
|
|
||||||
"""判断是否为 CLI 格式"""
|
|
||||||
return format_id.upper().endswith("_CLI")
|
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"OpenAIResponseParser",
|
"OpenAIResponseParser",
|
||||||
"OpenAICliResponseParser",
|
"OpenAICliResponseParser",
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ from abc import ABC, abstractmethod
|
|||||||
from typing import Any, Dict, FrozenSet, Optional, Tuple
|
from typing import Any, Dict, FrozenSet, Optional, Tuple
|
||||||
|
|
||||||
from src.core.crypto import crypto_service
|
from src.core.crypto import crypto_service
|
||||||
from src.core.headers import HeaderBuilder, UPSTREAM_DROP_HEADERS
|
from src.core.api_format import HeaderBuilder, UPSTREAM_DROP_HEADERS
|
||||||
|
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
# 统一的头部配置常量
|
# 统一的头部配置常量
|
||||||
@@ -125,7 +125,7 @@ class PassthroughRequestBuilder(RequestBuilder):
|
|||||||
"""
|
"""
|
||||||
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
|
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
|
||||||
"""
|
"""
|
||||||
from src.core.api_format_metadata import get_auth_config, resolve_api_format
|
from src.core.api_format import get_auth_config, resolve_api_format
|
||||||
|
|
||||||
# 1. 根据 API 格式自动设置认证头
|
# 1. 根据 API 格式自动设置认证头
|
||||||
decrypted_key = crypto_service.decrypt(key.api_key)
|
decrypted_key = crypto_service.decrypt(key.api_key)
|
||||||
|
|||||||
@@ -10,7 +10,10 @@
|
|||||||
|
|
||||||
import time
|
import time
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from src.core.api_format import GeminiStreamConversionState, StreamConversionState
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -82,6 +85,11 @@ class StreamContext:
|
|||||||
chunk_count: int = 0
|
chunk_count: int = 0
|
||||||
parsed_chunks: List[Dict[str, Any]] = field(default_factory=list)
|
parsed_chunks: List[Dict[str, Any]] = field(default_factory=list)
|
||||||
|
|
||||||
|
# 流式格式转换状态(跨 chunk 追踪)
|
||||||
|
stream_conversion_state: Optional[
|
||||||
|
Union["StreamConversionState", "GeminiStreamConversionState"]
|
||||||
|
] = None
|
||||||
|
|
||||||
def reset_for_retry(self) -> None:
|
def reset_for_retry(self) -> None:
|
||||||
"""
|
"""
|
||||||
重试时重置状态
|
重试时重置状态
|
||||||
@@ -109,6 +117,7 @@ class StreamContext:
|
|||||||
self.response_id = None
|
self.response_id = None
|
||||||
self.final_usage = None
|
self.final_usage = None
|
||||||
self.final_response = None
|
self.final_response = None
|
||||||
|
self.stream_conversion_state = None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def collected_text(self) -> str:
|
def collected_text(self) -> str:
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ import json
|
|||||||
from typing import Any, Dict, Optional
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
from src.core.exceptions import EmbeddedErrorException, ProviderNotAvailableException
|
from src.core.exceptions import EmbeddedErrorException, ProviderNotAvailableException
|
||||||
from src.core.headers import filter_response_headers
|
from src.core.api_format import filter_response_headers
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ from src.api.base.adapter import ApiAdapter, ApiMode
|
|||||||
from src.api.base.context import ApiRequestContext
|
from src.api.base.context import ApiRequestContext
|
||||||
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase, register_adapter
|
from src.api.handlers.base.chat_adapter_base import ChatAdapterBase, register_adapter
|
||||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||||
from src.core.headers import get_header_value
|
from src.core.api_format import get_header_value
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.core.optimization_utils import TokenCounter
|
from src.core.optimization_utils import TokenCounter
|
||||||
from src.models.claude import ClaudeMessagesRequest, ClaudeTokenCountRequest
|
from src.models.claude import ClaudeMessagesRequest, ClaudeTokenCountRequest
|
||||||
|
|||||||
@@ -74,7 +74,7 @@ class ClaudeChatHandler(ChatHandlerBase):
|
|||||||
Returns:
|
Returns:
|
||||||
ClaudeMessagesRequest 对象
|
ClaudeMessagesRequest 对象
|
||||||
"""
|
"""
|
||||||
from src.api.handlers.claude.converter import OpenAIToClaudeConverter
|
from src.core.api_format import OpenAIToClaudeConverter
|
||||||
from src.models.claude import ClaudeMessagesRequest
|
from src.models.claude import ClaudeMessagesRequest
|
||||||
from src.models.openai import OpenAIRequest
|
from src.models.openai import OpenAIRequest
|
||||||
|
|
||||||
|
|||||||
@@ -5,14 +5,14 @@ Gemini API Handler 模块
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from src.api.handlers.gemini.adapter import GeminiChatAdapter, build_gemini_adapter
|
from src.api.handlers.gemini.adapter import GeminiChatAdapter, build_gemini_adapter
|
||||||
from src.api.handlers.gemini.converter import (
|
from src.api.handlers.gemini.handler import GeminiChatHandler
|
||||||
|
from src.api.handlers.gemini.stream_parser import GeminiStreamParser
|
||||||
|
from src.core.api_format import (
|
||||||
ClaudeToGeminiConverter,
|
ClaudeToGeminiConverter,
|
||||||
GeminiToClaudeConverter,
|
GeminiToClaudeConverter,
|
||||||
GeminiToOpenAIConverter,
|
GeminiToOpenAIConverter,
|
||||||
OpenAIToGeminiConverter,
|
OpenAIToGeminiConverter,
|
||||||
)
|
)
|
||||||
from src.api.handlers.gemini.handler import GeminiChatHandler
|
|
||||||
from src.api.handlers.gemini.stream_parser import GeminiStreamParser
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"GeminiChatAdapter",
|
"GeminiChatAdapter",
|
||||||
|
|||||||
@@ -62,7 +62,7 @@ class GeminiChatHandler(ChatHandlerBase):
|
|||||||
Returns:
|
Returns:
|
||||||
GeminiRequest 对象
|
GeminiRequest 对象
|
||||||
"""
|
"""
|
||||||
from src.api.handlers.gemini.converter import (
|
from src.core.api_format import (
|
||||||
ClaudeToGeminiConverter,
|
ClaudeToGeminiConverter,
|
||||||
OpenAIToGeminiConverter,
|
OpenAIToGeminiConverter,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -73,7 +73,7 @@ class OpenAIChatHandler(ChatHandlerBase):
|
|||||||
Returns:
|
Returns:
|
||||||
OpenAIRequest 对象
|
OpenAIRequest 对象
|
||||||
"""
|
"""
|
||||||
from src.api.handlers.openai.converter import ClaudeToOpenAIConverter
|
from src.core.api_format import ClaudeToOpenAIConverter
|
||||||
from src.models.claude import ClaudeMessagesRequest
|
from src.models.claude import ClaudeMessagesRequest
|
||||||
from src.models.openai import OpenAIRequest
|
from src.models.openai import OpenAIRequest
|
||||||
|
|
||||||
|
|||||||
@@ -609,8 +609,7 @@ class PublicApiFormatHealthMonitorAdapter(PublicApiAdapter):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 获取本站入口路径
|
# 获取本站入口路径
|
||||||
from src.core.api_format_metadata import get_local_path
|
from src.core.api_format import APIFormat, get_local_path
|
||||||
from src.core.enums import APIFormat
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
api_format_enum = APIFormat(api_format)
|
api_format_enum = APIFormat(api_format)
|
||||||
|
|||||||
@@ -15,8 +15,7 @@ from src.api.handlers.claude import (
|
|||||||
ClaudeTokenCountAdapter,
|
ClaudeTokenCountAdapter,
|
||||||
build_claude_adapter,
|
build_claude_adapter,
|
||||||
)
|
)
|
||||||
from src.core.api_format_metadata import get_api_format_definition
|
from src.core.api_format import APIFormat, get_api_format_definition
|
||||||
from src.core.enums import APIFormat
|
|
||||||
from src.database import get_db
|
from src.database import get_db
|
||||||
|
|
||||||
_claude_def = get_api_format_definition(APIFormat.CLAUDE)
|
_claude_def = get_api_format_definition(APIFormat.CLAUDE)
|
||||||
|
|||||||
@@ -16,8 +16,7 @@ from sqlalchemy.orm import Session
|
|||||||
from src.api.base.pipeline import ApiRequestPipeline
|
from src.api.base.pipeline import ApiRequestPipeline
|
||||||
from src.api.handlers.gemini import build_gemini_adapter
|
from src.api.handlers.gemini import build_gemini_adapter
|
||||||
from src.api.handlers.gemini_cli import build_gemini_cli_adapter
|
from src.api.handlers.gemini_cli import build_gemini_cli_adapter
|
||||||
from src.core.api_format_metadata import get_api_format_definition
|
from src.core.api_format import APIFormat, get_api_format_definition
|
||||||
from src.core.enums import APIFormat
|
|
||||||
from src.database import get_db
|
from src.database import get_db
|
||||||
|
|
||||||
# 从配置获取路径前缀
|
# 从配置获取路径前缀
|
||||||
|
|||||||
@@ -20,8 +20,12 @@ from src.api.base.models_service import (
|
|||||||
get_available_provider_ids,
|
get_available_provider_ids,
|
||||||
list_available_models,
|
list_available_models,
|
||||||
)
|
)
|
||||||
from src.core.api_format_metadata import API_FORMAT_DEFINITIONS, ApiFormatDefinition
|
from src.core.api_format import (
|
||||||
from src.core.enums import APIFormat
|
API_FORMAT_DEFINITIONS,
|
||||||
|
APIFormat,
|
||||||
|
ApiFormatDefinition,
|
||||||
|
detect_format_and_key_from_starlette,
|
||||||
|
)
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.database import get_db
|
from src.database import get_db
|
||||||
from src.models.database import ApiKey, User
|
from src.models.database import ApiKey, User
|
||||||
@@ -72,26 +76,7 @@ def _detect_api_format_and_key(request: Request) -> Tuple[str, Optional[str]]:
|
|||||||
Returns:
|
Returns:
|
||||||
(api_format, api_key) 元组
|
(api_format, api_key) 元组
|
||||||
"""
|
"""
|
||||||
# Claude: x-api-key + anthropic-version (必须同时存在)
|
return detect_format_and_key_from_starlette(request)
|
||||||
claude_def = API_FORMAT_DEFINITIONS[APIFormat.CLAUDE]
|
|
||||||
claude_key = _extract_api_key_from_request(request, claude_def)
|
|
||||||
if claude_key and request.headers.get("anthropic-version"):
|
|
||||||
return "claude", claude_key
|
|
||||||
|
|
||||||
# Gemini: x-goog-api-key (header 类型) 或 ?key=
|
|
||||||
gemini_def = API_FORMAT_DEFINITIONS[APIFormat.GEMINI]
|
|
||||||
gemini_key = _extract_api_key_from_request(request, gemini_def)
|
|
||||||
if gemini_key:
|
|
||||||
return "gemini", gemini_key
|
|
||||||
|
|
||||||
# OpenAI: Authorization: Bearer (默认)
|
|
||||||
# 注意: 如果只有 x-api-key 但没有 anthropic-version,也走 OpenAI 格式
|
|
||||||
openai_def = API_FORMAT_DEFINITIONS[APIFormat.OPENAI]
|
|
||||||
openai_key = _extract_api_key_from_request(request, openai_def)
|
|
||||||
# 如果 OpenAI 格式没有 key,但有 x-api-key,也用它(兼容)
|
|
||||||
if not openai_key and claude_key:
|
|
||||||
openai_key = claude_key
|
|
||||||
return "openai", openai_key
|
|
||||||
|
|
||||||
|
|
||||||
def _get_formats_for_api(api_format: str) -> list[str]:
|
def _get_formats_for_api(api_format: str) -> list[str]:
|
||||||
|
|||||||
@@ -13,8 +13,7 @@ from sqlalchemy.orm import Session
|
|||||||
from src.api.base.pipeline import ApiRequestPipeline
|
from src.api.base.pipeline import ApiRequestPipeline
|
||||||
from src.api.handlers.openai import OpenAIChatAdapter
|
from src.api.handlers.openai import OpenAIChatAdapter
|
||||||
from src.api.handlers.openai_cli import OpenAICliAdapter
|
from src.api.handlers.openai_cli import OpenAICliAdapter
|
||||||
from src.core.api_format_metadata import get_api_format_definition
|
from src.core.api_format import APIFormat, get_api_format_definition
|
||||||
from src.core.enums import APIFormat
|
|
||||||
from src.database import get_db
|
from src.database import get_db
|
||||||
|
|
||||||
_openai_def = get_api_format_definition(APIFormat.OPENAI)
|
_openai_def = get_api_format_definition(APIFormat.OPENAI)
|
||||||
|
|||||||
154
src/core/api_format/__init__.py
Normal file
154
src/core/api_format/__init__.py
Normal file
@@ -0,0 +1,154 @@
|
|||||||
|
"""
|
||||||
|
API 格式核心模块
|
||||||
|
|
||||||
|
统一管理 API 格式相关的枚举、元数据、工具函数和格式转换功能。
|
||||||
|
|
||||||
|
模块组成:
|
||||||
|
- enums.py: APIFormat 枚举定义
|
||||||
|
- metadata.py: 格式元数据定义(别名、路径、认证等)
|
||||||
|
- headers.py: 请求头处理(构建、过滤、脱敏)
|
||||||
|
- utils.py: 工具函数(is_cli_format, get_base_format 等)
|
||||||
|
- detection.py: 格式检测(从请求头、响应内容检测格式)
|
||||||
|
- conversion/: 格式转换子模块
|
||||||
|
"""
|
||||||
|
|
||||||
|
from src.core.api_format.conversion import (
|
||||||
|
ClaudeToGeminiConverter,
|
||||||
|
ClaudeToOpenAIConverter,
|
||||||
|
FormatConversionError,
|
||||||
|
FormatConverterRegistry,
|
||||||
|
GeminiStreamConversionState,
|
||||||
|
GeminiToClaudeConverter,
|
||||||
|
GeminiToOpenAIConverter,
|
||||||
|
OpenAIToClaudeConverter,
|
||||||
|
OpenAIToGeminiConverter,
|
||||||
|
RequestConverter,
|
||||||
|
ResponseConverter,
|
||||||
|
StreamChunkConverter,
|
||||||
|
StreamConversionState,
|
||||||
|
converter_registry,
|
||||||
|
is_format_compatible,
|
||||||
|
register_all_converters,
|
||||||
|
)
|
||||||
|
from src.core.api_format.detection import (
|
||||||
|
detect_cli_format_from_path,
|
||||||
|
detect_format_and_key_from_starlette,
|
||||||
|
detect_format_from_request,
|
||||||
|
detect_format_from_response,
|
||||||
|
)
|
||||||
|
from src.core.api_format.enums import APIFormat
|
||||||
|
from src.core.api_format.headers import (
|
||||||
|
CORE_REDACT_HEADERS,
|
||||||
|
HOP_BY_HOP_HEADERS,
|
||||||
|
RESPONSE_DROP_HEADERS,
|
||||||
|
SENSITIVE_HEADERS,
|
||||||
|
UPSTREAM_DROP_HEADERS,
|
||||||
|
HeaderBuilder,
|
||||||
|
build_adapter_base_headers,
|
||||||
|
build_adapter_headers,
|
||||||
|
build_upstream_headers,
|
||||||
|
detect_capabilities,
|
||||||
|
extract_client_api_key,
|
||||||
|
extract_set_headers_from_rules,
|
||||||
|
filter_response_headers,
|
||||||
|
get_adapter_protected_keys,
|
||||||
|
get_extra_headers_from_endpoint,
|
||||||
|
get_header_value,
|
||||||
|
merge_headers_with_protection,
|
||||||
|
normalize_headers,
|
||||||
|
redact_headers_for_log,
|
||||||
|
)
|
||||||
|
from src.core.api_format.metadata import (
|
||||||
|
API_FORMAT_DEFINITIONS,
|
||||||
|
ApiFormatDefinition,
|
||||||
|
get_api_format_definition,
|
||||||
|
get_auth_config,
|
||||||
|
get_default_path,
|
||||||
|
get_extra_headers,
|
||||||
|
get_local_path,
|
||||||
|
get_protected_keys,
|
||||||
|
is_cli_api_format,
|
||||||
|
list_api_format_definitions,
|
||||||
|
register_api_format_definition,
|
||||||
|
resolve_api_format,
|
||||||
|
resolve_api_format_alias,
|
||||||
|
)
|
||||||
|
from src.core.api_format.utils import (
|
||||||
|
get_base_format,
|
||||||
|
is_cli_format,
|
||||||
|
is_convertible_format,
|
||||||
|
is_same_format,
|
||||||
|
normalize_format,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
# Enums
|
||||||
|
"APIFormat",
|
||||||
|
# Metadata
|
||||||
|
"ApiFormatDefinition",
|
||||||
|
"API_FORMAT_DEFINITIONS",
|
||||||
|
"get_api_format_definition",
|
||||||
|
"list_api_format_definitions",
|
||||||
|
"resolve_api_format",
|
||||||
|
"resolve_api_format_alias",
|
||||||
|
"register_api_format_definition",
|
||||||
|
"get_default_path",
|
||||||
|
"get_local_path",
|
||||||
|
"get_auth_config",
|
||||||
|
"get_extra_headers",
|
||||||
|
"get_protected_keys",
|
||||||
|
"is_cli_api_format",
|
||||||
|
# Utils
|
||||||
|
"is_cli_format",
|
||||||
|
"get_base_format",
|
||||||
|
"normalize_format",
|
||||||
|
"is_same_format",
|
||||||
|
"is_convertible_format",
|
||||||
|
# Headers
|
||||||
|
"UPSTREAM_DROP_HEADERS",
|
||||||
|
"CORE_REDACT_HEADERS",
|
||||||
|
"HOP_BY_HOP_HEADERS",
|
||||||
|
"RESPONSE_DROP_HEADERS",
|
||||||
|
"SENSITIVE_HEADERS",
|
||||||
|
"normalize_headers",
|
||||||
|
"get_header_value",
|
||||||
|
"extract_client_api_key",
|
||||||
|
"detect_capabilities",
|
||||||
|
"HeaderBuilder",
|
||||||
|
"build_upstream_headers",
|
||||||
|
"merge_headers_with_protection",
|
||||||
|
"filter_response_headers",
|
||||||
|
"redact_headers_for_log",
|
||||||
|
"build_adapter_base_headers",
|
||||||
|
"build_adapter_headers",
|
||||||
|
"get_adapter_protected_keys",
|
||||||
|
"extract_set_headers_from_rules",
|
||||||
|
"get_extra_headers_from_endpoint",
|
||||||
|
# Registry
|
||||||
|
"FormatConverterRegistry",
|
||||||
|
"converter_registry",
|
||||||
|
"register_all_converters",
|
||||||
|
# Protocols
|
||||||
|
"RequestConverter",
|
||||||
|
"ResponseConverter",
|
||||||
|
"StreamChunkConverter",
|
||||||
|
# State
|
||||||
|
"StreamConversionState",
|
||||||
|
"GeminiStreamConversionState",
|
||||||
|
# Exceptions
|
||||||
|
"FormatConversionError",
|
||||||
|
# Compatibility
|
||||||
|
"is_format_compatible",
|
||||||
|
# Detection
|
||||||
|
"detect_format_from_request",
|
||||||
|
"detect_format_and_key_from_starlette",
|
||||||
|
"detect_format_from_response",
|
||||||
|
"detect_cli_format_from_path",
|
||||||
|
# Converters
|
||||||
|
"OpenAIToClaudeConverter",
|
||||||
|
"ClaudeToOpenAIConverter",
|
||||||
|
"ClaudeToGeminiConverter",
|
||||||
|
"GeminiToClaudeConverter",
|
||||||
|
"OpenAIToGeminiConverter",
|
||||||
|
"GeminiToOpenAIConverter",
|
||||||
|
]
|
||||||
86
src/core/api_format/conversion/__init__.py
Normal file
86
src/core/api_format/conversion/__init__.py
Normal file
@@ -0,0 +1,86 @@
|
|||||||
|
"""
|
||||||
|
格式转换核心模块
|
||||||
|
|
||||||
|
该目录用于承载与「API 格式转换」相关的核心能力(不依赖 FastAPI/Handler 层),
|
||||||
|
以便 services/core 可复用并避免出现 services -> api 的反向依赖。
|
||||||
|
|
||||||
|
模块组成:
|
||||||
|
- registry.py: 转换器注册表,管理转换器实例和能力查询
|
||||||
|
- protocols.py: 转换器协议定义(Protocol)
|
||||||
|
- state.py: 流式转换状态类
|
||||||
|
- exceptions.py: 转换异常定义
|
||||||
|
- compatibility.py: 格式兼容性检查函数
|
||||||
|
- converters/: 内置格式转换器实现
|
||||||
|
"""
|
||||||
|
|
||||||
|
from src.core.api_format.conversion.compatibility import is_format_compatible
|
||||||
|
from src.core.api_format.conversion.converters import (
|
||||||
|
ClaudeToGeminiConverter,
|
||||||
|
ClaudeToOpenAIConverter,
|
||||||
|
GeminiToClaudeConverter,
|
||||||
|
GeminiToOpenAIConverter,
|
||||||
|
OpenAIToClaudeConverter,
|
||||||
|
OpenAIToGeminiConverter,
|
||||||
|
)
|
||||||
|
from src.core.api_format.conversion.exceptions import FormatConversionError
|
||||||
|
from src.core.api_format.conversion.protocols import (
|
||||||
|
RequestConverter,
|
||||||
|
ResponseConverter,
|
||||||
|
StreamChunkConverter,
|
||||||
|
)
|
||||||
|
from src.core.api_format.conversion.registry import (
|
||||||
|
FormatConverterRegistry,
|
||||||
|
converter_registry,
|
||||||
|
)
|
||||||
|
from src.core.api_format.conversion.state import (
|
||||||
|
GeminiStreamConversionState,
|
||||||
|
StreamConversionState,
|
||||||
|
)
|
||||||
|
from src.core.logger import logger
|
||||||
|
|
||||||
|
|
||||||
|
def register_all_converters() -> None:
|
||||||
|
"""
|
||||||
|
注册所有内置的格式转换器
|
||||||
|
|
||||||
|
在应用启动时调用此函数
|
||||||
|
"""
|
||||||
|
# Claude <-> OpenAI
|
||||||
|
converter_registry.register("OPENAI", "CLAUDE", OpenAIToClaudeConverter())
|
||||||
|
converter_registry.register("CLAUDE", "OPENAI", ClaudeToOpenAIConverter())
|
||||||
|
|
||||||
|
# Claude <-> Gemini
|
||||||
|
converter_registry.register("CLAUDE", "GEMINI", ClaudeToGeminiConverter())
|
||||||
|
converter_registry.register("GEMINI", "CLAUDE", GeminiToClaudeConverter())
|
||||||
|
|
||||||
|
# OpenAI <-> Gemini
|
||||||
|
converter_registry.register("OPENAI", "GEMINI", OpenAIToGeminiConverter())
|
||||||
|
converter_registry.register("GEMINI", "OPENAI", GeminiToOpenAIConverter())
|
||||||
|
|
||||||
|
logger.info(f"[ConverterRegistry] 已注册 {len(converter_registry.list_converters())} 个格式转换器")
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
# Registry
|
||||||
|
"FormatConverterRegistry",
|
||||||
|
"converter_registry",
|
||||||
|
"register_all_converters",
|
||||||
|
# Protocols
|
||||||
|
"RequestConverter",
|
||||||
|
"ResponseConverter",
|
||||||
|
"StreamChunkConverter",
|
||||||
|
# State
|
||||||
|
"StreamConversionState",
|
||||||
|
"GeminiStreamConversionState",
|
||||||
|
# Exceptions
|
||||||
|
"FormatConversionError",
|
||||||
|
# Compatibility
|
||||||
|
"is_format_compatible",
|
||||||
|
# Converters
|
||||||
|
"OpenAIToClaudeConverter",
|
||||||
|
"ClaudeToOpenAIConverter",
|
||||||
|
"ClaudeToGeminiConverter",
|
||||||
|
"GeminiToClaudeConverter",
|
||||||
|
"OpenAIToGeminiConverter",
|
||||||
|
"GeminiToOpenAIConverter",
|
||||||
|
]
|
||||||
103
src/core/api_format/conversion/compatibility.py
Normal file
103
src/core/api_format/conversion/compatibility.py
Normal file
@@ -0,0 +1,103 @@
|
|||||||
|
"""
|
||||||
|
格式兼容性检查
|
||||||
|
|
||||||
|
用于候选筛选时判断端点是否可以处理客户端请式。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import TYPE_CHECKING, Optional, Tuple
|
||||||
|
|
||||||
|
from src.core.api_format.utils import is_cli_format
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from src.core.api_format.conversion.registry import FormatConverterRegistry
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def is_format_compatible(
|
||||||
|
client_format: str,
|
||||||
|
endpoint_api_format: str,
|
||||||
|
endpoint_format_acceptance_config: Optional[dict],
|
||||||
|
is_stream: bool,
|
||||||
|
global_conversion_enabled: bool,
|
||||||
|
registry: Optional["FormatConverterRegistry"] = None,
|
||||||
|
) -> Tuple[bool, bool, Optional[str]]:
|
||||||
|
"""
|
||||||
|
检查端点是否兼容客户端格式
|
||||||
|
|
||||||
|
Args:
|
||||||
|
client_format: 客户端请求格式
|
||||||
|
endpoint_api_format: 端点的 API 格式
|
||||||
|
endpoint_format_acceptance_config: 端点的格式接受配置
|
||||||
|
is_stream: 是否是流式请求
|
||||||
|
global_conversion_enabled: 全局格式转换开关
|
||||||
|
registry: 转换器注册表(可选,默认使用全局单例)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(is_compatible, needs_conversion, skip_reason)
|
||||||
|
- is_compatible: 是否兼容
|
||||||
|
- needs_conversion: 是否需要转换
|
||||||
|
- skip_reason: 不兼容时的原因
|
||||||
|
"""
|
||||||
|
# 延迟导入避免循环依赖
|
||||||
|
if registry is None:
|
||||||
|
from src.core.api_format.conversion.registry import converter_registry
|
||||||
|
|
||||||
|
registry = converter_registry
|
||||||
|
|
||||||
|
provider_format = endpoint_api_format.upper()
|
||||||
|
client_format_upper = client_format.upper()
|
||||||
|
|
||||||
|
# 1. 格式完全匹配 -> 兼容,无需转换
|
||||||
|
if provider_format == client_format_upper:
|
||||||
|
return True, False, None
|
||||||
|
|
||||||
|
# 2. CLI 格式不参与转换
|
||||||
|
if is_cli_format(client_format_upper):
|
||||||
|
return False, False, "CLI 格式不支持转换"
|
||||||
|
if is_cli_format(provider_format):
|
||||||
|
return False, False, "Provider 为 CLI 格式,不支持转换"
|
||||||
|
|
||||||
|
# 3. 检查全局开关
|
||||||
|
if not global_conversion_enabled:
|
||||||
|
return False, False, "全局格式转换未启用"
|
||||||
|
|
||||||
|
# 4. 检查端点配置
|
||||||
|
if endpoint_format_acceptance_config is None:
|
||||||
|
return False, False, "端点未配置格式转换"
|
||||||
|
|
||||||
|
config = endpoint_format_acceptance_config
|
||||||
|
if not config.get("enabled", False):
|
||||||
|
return False, False, "端点格式转换未启用"
|
||||||
|
|
||||||
|
# 检查 reject_formats(优先)
|
||||||
|
reject_formats = config.get("reject_formats", [])
|
||||||
|
if client_format_upper in [f.upper() for f in reject_formats]:
|
||||||
|
return False, False, f"端点拒绝 {client_format} 格式"
|
||||||
|
|
||||||
|
# 检查 accept_formats
|
||||||
|
accept_formats = config.get("accept_formats", [])
|
||||||
|
if accept_formats and client_format_upper not in [f.upper() for f in accept_formats]:
|
||||||
|
return False, False, f"端点不接受 {client_format} 格式"
|
||||||
|
|
||||||
|
# 检查流式转换
|
||||||
|
if is_stream and not config.get("stream_conversion", True):
|
||||||
|
return False, False, "端点不支持流式格式转换"
|
||||||
|
|
||||||
|
# 5. 检查转换器能力
|
||||||
|
if not registry.can_convert_full(
|
||||||
|
client_format_upper,
|
||||||
|
provider_format,
|
||||||
|
require_stream=is_stream,
|
||||||
|
):
|
||||||
|
return False, False, f"不存在 {client_format} <-> {provider_format} 的完整转换器"
|
||||||
|
|
||||||
|
return True, True, None
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"is_format_compatible",
|
||||||
|
]
|
||||||
26
src/core/api_format/conversion/converters/__init__.py
Normal file
26
src/core/api_format/conversion/converters/__init__.py
Normal file
@@ -0,0 +1,26 @@
|
|||||||
|
"""
|
||||||
|
格式转换器集合
|
||||||
|
|
||||||
|
将 Claude/OpenAI/Gemini 格式之间的转换器统一导出。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from .claude_to_openai import ClaudeToOpenAIConverter
|
||||||
|
from .gemini import (
|
||||||
|
ClaudeToGeminiConverter,
|
||||||
|
GeminiToClaudeConverter,
|
||||||
|
GeminiToOpenAIConverter,
|
||||||
|
OpenAIToGeminiConverter,
|
||||||
|
)
|
||||||
|
from .openai_to_claude import OpenAIToClaudeConverter
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
# OpenAI <-> Claude
|
||||||
|
"OpenAIToClaudeConverter",
|
||||||
|
"ClaudeToOpenAIConverter",
|
||||||
|
# Claude <-> Gemini
|
||||||
|
"ClaudeToGeminiConverter",
|
||||||
|
"GeminiToClaudeConverter",
|
||||||
|
# OpenAI <-> Gemini
|
||||||
|
"OpenAIToGeminiConverter",
|
||||||
|
"GeminiToOpenAIConverter",
|
||||||
|
]
|
||||||
@@ -9,9 +9,10 @@ from __future__ import annotations
|
|||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from typing import Any, Dict, List, Optional, Union
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
from pydantic import BaseModel
|
if TYPE_CHECKING:
|
||||||
|
from src.core.api_format.conversion.state import StreamConversionState
|
||||||
|
|
||||||
|
|
||||||
class ClaudeToOpenAIConverter:
|
class ClaudeToOpenAIConverter:
|
||||||
@@ -47,7 +48,7 @@ class ClaudeToOpenAIConverter:
|
|||||||
|
|
||||||
# ==================== 请求转换 ====================
|
# ==================== 请求转换 ====================
|
||||||
|
|
||||||
def convert_request(self, request: Union[Dict[str, Any], BaseModel]) -> Dict[str, Any]:
|
def convert_request(self, request: Union[Dict[str, Any], Any]) -> Dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
将 Claude 请求转换为 OpenAI 格式
|
将 Claude 请求转换为 OpenAI 格式
|
||||||
|
|
||||||
@@ -269,7 +270,7 @@ class ClaudeToOpenAIConverter:
|
|||||||
|
|
||||||
# 转换停止原因
|
# 转换停止原因
|
||||||
stop_reason = response.get("stop_reason")
|
stop_reason = response.get("stop_reason")
|
||||||
finish_reason = self.STOP_REASON_MAP.get(stop_reason, stop_reason)
|
finish_reason = self.STOP_REASON_MAP.get(stop_reason, stop_reason) if stop_reason else None
|
||||||
|
|
||||||
# 转换 usage
|
# 转换 usage
|
||||||
usage = response.get("usage", {})
|
usage = response.get("usage", {})
|
||||||
@@ -296,14 +297,39 @@ class ClaudeToOpenAIConverter:
|
|||||||
|
|
||||||
# ==================== 流式转换 ====================
|
# ==================== 流式转换 ====================
|
||||||
|
|
||||||
def convert_stream_event(
|
def convert_stream_chunk(
|
||||||
|
self,
|
||||||
|
chunk: Dict[str, Any],
|
||||||
|
state: Optional["StreamConversionState"] = None,
|
||||||
|
) -> List[Dict[str, Any]]:
|
||||||
|
"""
|
||||||
|
将 Claude SSE 事件转换为 OpenAI 格式
|
||||||
|
|
||||||
|
Args:
|
||||||
|
chunk: Claude SSE 事件
|
||||||
|
state: 流式转换状态
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
OpenAI 格式的 SSE chunk 列表
|
||||||
|
"""
|
||||||
|
from src.core.api_format.conversion.state import StreamConversionState
|
||||||
|
|
||||||
|
if state is None:
|
||||||
|
state = StreamConversionState()
|
||||||
|
|
||||||
|
result = self._convert_single_event(chunk, state.model, state.message_id)
|
||||||
|
if result is None:
|
||||||
|
return []
|
||||||
|
return [result]
|
||||||
|
|
||||||
|
def _convert_single_event(
|
||||||
self,
|
self,
|
||||||
event: Dict[str, Any],
|
event: Dict[str, Any],
|
||||||
model: str = "",
|
model: str = "",
|
||||||
message_id: Optional[str] = None,
|
message_id: Optional[str] = None,
|
||||||
) -> Optional[Dict[str, Any]]:
|
) -> Optional[Dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
将 Claude SSE 事件转换为 OpenAI 格式
|
转换单个 Claude SSE 事件为 OpenAI 格式
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
event: Claude SSE 事件
|
event: Claude SSE 事件
|
||||||
@@ -4,7 +4,14 @@ Gemini 格式转换器
|
|||||||
提供 Gemini 与其他 API 格式(Claude、OpenAI)之间的转换
|
提供 Gemini 与其他 API 格式(Claude、OpenAI)之间的转换
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Any, Dict, List, Optional
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from src.core.api_format.conversion.state import GeminiStreamConversionState
|
||||||
|
|
||||||
|
|
||||||
class ClaudeToGeminiConverter:
|
class ClaudeToGeminiConverter:
|
||||||
@@ -269,6 +276,146 @@ class GeminiToClaudeConverter:
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# ==================== 流式转换 ====================
|
||||||
|
|
||||||
|
def convert_stream_chunk(
|
||||||
|
self,
|
||||||
|
chunk: Dict[str, Any],
|
||||||
|
state: Optional["GeminiStreamConversionState"] = None,
|
||||||
|
) -> List[Dict[str, Any]]:
|
||||||
|
"""
|
||||||
|
将 Gemini 流式响应转换为 Claude SSE 事件
|
||||||
|
|
||||||
|
Gemini 流式格式与 Claude/OpenAI 不同:
|
||||||
|
- Gemini 返回完整累积文本,需要计算增量
|
||||||
|
- 需要生成 message_start/content_block_delta/message_delta 事件序列
|
||||||
|
|
||||||
|
Args:
|
||||||
|
chunk: Gemini 流式响应块
|
||||||
|
state: 流式转换状态(跨 chunk 追踪)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Claude SSE 事件列表
|
||||||
|
"""
|
||||||
|
from src.core.api_format.conversion.state import GeminiStreamConversionState
|
||||||
|
|
||||||
|
if state is None:
|
||||||
|
state = GeminiStreamConversionState()
|
||||||
|
|
||||||
|
events: List[Dict[str, Any]] = []
|
||||||
|
candidates = chunk.get("candidates", [])
|
||||||
|
if not candidates:
|
||||||
|
return events
|
||||||
|
|
||||||
|
candidate = candidates[0]
|
||||||
|
content = candidate.get("content", {})
|
||||||
|
parts = content.get("parts", [])
|
||||||
|
|
||||||
|
# 发送 message_start(首次)
|
||||||
|
if not state.message_started:
|
||||||
|
events.append(
|
||||||
|
{
|
||||||
|
"type": "message_start",
|
||||||
|
"message": {
|
||||||
|
"id": state.message_id or "msg_gemini",
|
||||||
|
"type": "message",
|
||||||
|
"role": "assistant",
|
||||||
|
"model": state.model,
|
||||||
|
"content": [],
|
||||||
|
"stop_reason": None,
|
||||||
|
"stop_sequence": None,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
state.message_started = True
|
||||||
|
|
||||||
|
# 处理文本增量
|
||||||
|
for part in parts:
|
||||||
|
if "text" in part:
|
||||||
|
full_text = part["text"]
|
||||||
|
# Gemini 返回累积文本,计算增量
|
||||||
|
delta_text = full_text[len(state.accumulated_text) :]
|
||||||
|
if delta_text:
|
||||||
|
state.accumulated_text = full_text
|
||||||
|
# 首次文本需要发送 content_block_start
|
||||||
|
if not state.content_block_started:
|
||||||
|
events.append(
|
||||||
|
{
|
||||||
|
"type": "content_block_start",
|
||||||
|
"index": state.current_block_index,
|
||||||
|
"content_block": {"type": "text", "text": ""},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
state.content_block_started = True
|
||||||
|
events.append(
|
||||||
|
{
|
||||||
|
"type": "content_block_delta",
|
||||||
|
"index": state.current_block_index,
|
||||||
|
"delta": {"type": "text_delta", "text": delta_text},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
elif "functionCall" in part:
|
||||||
|
# 工具调用
|
||||||
|
func_call = part["functionCall"]
|
||||||
|
tool_id = f"toolu_{func_call.get('name', '')}_{state.tool_call_index}"
|
||||||
|
state.tool_call_index += 1
|
||||||
|
state.current_block_index += 1
|
||||||
|
|
||||||
|
events.append(
|
||||||
|
{
|
||||||
|
"type": "content_block_start",
|
||||||
|
"index": state.current_block_index,
|
||||||
|
"content_block": {
|
||||||
|
"type": "tool_use",
|
||||||
|
"id": tool_id,
|
||||||
|
"name": func_call.get("name", ""),
|
||||||
|
"input": {},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
# 发送完整的 input 作为 JSON delta
|
||||||
|
args = func_call.get("args", {})
|
||||||
|
if args:
|
||||||
|
events.append(
|
||||||
|
{
|
||||||
|
"type": "content_block_delta",
|
||||||
|
"index": state.current_block_index,
|
||||||
|
"delta": {
|
||||||
|
"type": "input_json_delta",
|
||||||
|
"partial_json": json.dumps(args, ensure_ascii=False),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
events.append(
|
||||||
|
{
|
||||||
|
"type": "content_block_stop",
|
||||||
|
"index": state.current_block_index,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
# 处理结束
|
||||||
|
finish_reason = candidate.get("finishReason")
|
||||||
|
if finish_reason:
|
||||||
|
# 关闭当前内容块
|
||||||
|
if state.content_block_started:
|
||||||
|
events.append(
|
||||||
|
{
|
||||||
|
"type": "content_block_stop",
|
||||||
|
"index": state.current_block_index,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
stop_reason = self._convert_finish_reason(finish_reason)
|
||||||
|
events.append(
|
||||||
|
{
|
||||||
|
"type": "message_delta",
|
||||||
|
"delta": {"stop_reason": stop_reason},
|
||||||
|
"usage": {"output_tokens": 0},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
return events
|
||||||
|
|
||||||
|
|
||||||
class OpenAIToGeminiConverter:
|
class OpenAIToGeminiConverter:
|
||||||
"""
|
"""
|
||||||
@@ -334,8 +481,6 @@ class OpenAIToGeminiConverter:
|
|||||||
for tc in tool_calls:
|
for tc in tool_calls:
|
||||||
if tc.get("type") == "function":
|
if tc.get("type") == "function":
|
||||||
func = tc.get("function", {})
|
func = tc.get("function", {})
|
||||||
import json
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
args = json.loads(func.get("arguments", "{}"))
|
args = json.loads(func.get("arguments", "{}"))
|
||||||
except json.JSONDecodeError:
|
except json.JSONDecodeError:
|
||||||
@@ -455,8 +600,6 @@ class GeminiToOpenAIConverter:
|
|||||||
Returns:
|
Returns:
|
||||||
OpenAI 格式的响应字典
|
OpenAI 格式的响应字典
|
||||||
"""
|
"""
|
||||||
import time
|
|
||||||
|
|
||||||
candidates = gemini_response.get("candidates", [])
|
candidates = gemini_response.get("candidates", [])
|
||||||
choices = []
|
choices = []
|
||||||
|
|
||||||
@@ -473,8 +616,6 @@ class GeminiToOpenAIConverter:
|
|||||||
text_parts.append(part["text"])
|
text_parts.append(part["text"])
|
||||||
elif "functionCall" in part:
|
elif "functionCall" in part:
|
||||||
func_call = part["functionCall"]
|
func_call = part["functionCall"]
|
||||||
import json
|
|
||||||
|
|
||||||
tool_calls.append(
|
tool_calls.append(
|
||||||
{
|
{
|
||||||
"id": f"call_{func_call.get('name', '')}_{i}",
|
"id": f"call_{func_call.get('name', '')}_{i}",
|
||||||
@@ -539,6 +680,173 @@ class GeminiToOpenAIConverter:
|
|||||||
return "stop"
|
return "stop"
|
||||||
return mapping.get(finish_reason, "stop")
|
return mapping.get(finish_reason, "stop")
|
||||||
|
|
||||||
|
# ==================== 流式转换 ====================
|
||||||
|
|
||||||
|
def convert_stream_chunk(
|
||||||
|
self,
|
||||||
|
chunk: Dict[str, Any],
|
||||||
|
state: Optional["GeminiStreamConversionState"] = None,
|
||||||
|
) -> List[Dict[str, Any]]:
|
||||||
|
"""
|
||||||
|
将 Gemini 流式响应转换为 OpenAI SSE chunk
|
||||||
|
|
||||||
|
Gemini 流式格式与 OpenAI 不同:
|
||||||
|
- Gemini 返回完整累积文本,需要计算增量
|
||||||
|
- 需要生成 OpenAI chat.completion.chunk 格式
|
||||||
|
|
||||||
|
Args:
|
||||||
|
chunk: Gemini 流式响应块
|
||||||
|
state: 流式转换状态(跨 chunk 追踪)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
OpenAI SSE chunk 列表
|
||||||
|
"""
|
||||||
|
from src.core.api_format.conversion.state import GeminiStreamConversionState
|
||||||
|
|
||||||
|
if state is None:
|
||||||
|
state = GeminiStreamConversionState()
|
||||||
|
|
||||||
|
events: List[Dict[str, Any]] = []
|
||||||
|
candidates = chunk.get("candidates", [])
|
||||||
|
if not candidates:
|
||||||
|
return events
|
||||||
|
|
||||||
|
candidate = candidates[0]
|
||||||
|
content = candidate.get("content", {})
|
||||||
|
parts = content.get("parts", [])
|
||||||
|
finish_reason = candidate.get("finishReason")
|
||||||
|
|
||||||
|
chunk_id = f"chatcmpl-{state.message_id or 'gemini'}"
|
||||||
|
|
||||||
|
# 发送首个 chunk(带 role)
|
||||||
|
if not state.message_started:
|
||||||
|
events.append(
|
||||||
|
{
|
||||||
|
"id": chunk_id,
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": int(time.time()),
|
||||||
|
"model": state.model,
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"delta": {"role": "assistant", "content": ""},
|
||||||
|
"finish_reason": None,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
state.message_started = True
|
||||||
|
|
||||||
|
# 处理文本增量
|
||||||
|
for part in parts:
|
||||||
|
if "text" in part:
|
||||||
|
full_text = part["text"]
|
||||||
|
# Gemini 返回累积文本,计算增量
|
||||||
|
delta_text = full_text[len(state.accumulated_text) :]
|
||||||
|
if delta_text:
|
||||||
|
state.accumulated_text = full_text
|
||||||
|
events.append(
|
||||||
|
{
|
||||||
|
"id": chunk_id,
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": int(time.time()),
|
||||||
|
"model": state.model,
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"delta": {"content": delta_text},
|
||||||
|
"finish_reason": None,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
elif "functionCall" in part:
|
||||||
|
# 工具调用
|
||||||
|
func_call = part["functionCall"]
|
||||||
|
tool_id = f"call_{func_call.get('name', '')}_{state.tool_call_index}"
|
||||||
|
|
||||||
|
# 工具调用开始
|
||||||
|
events.append(
|
||||||
|
{
|
||||||
|
"id": chunk_id,
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": int(time.time()),
|
||||||
|
"model": state.model,
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"delta": {
|
||||||
|
"tool_calls": [
|
||||||
|
{
|
||||||
|
"index": state.tool_call_index,
|
||||||
|
"id": tool_id,
|
||||||
|
"type": "function",
|
||||||
|
"function": {
|
||||||
|
"name": func_call.get("name", ""),
|
||||||
|
"arguments": "",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"finish_reason": None,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
# 工具调用参数
|
||||||
|
args = func_call.get("args", {})
|
||||||
|
if args:
|
||||||
|
events.append(
|
||||||
|
{
|
||||||
|
"id": chunk_id,
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": int(time.time()),
|
||||||
|
"model": state.model,
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"delta": {
|
||||||
|
"tool_calls": [
|
||||||
|
{
|
||||||
|
"index": state.tool_call_index,
|
||||||
|
"function": {
|
||||||
|
"arguments": json.dumps(
|
||||||
|
args, ensure_ascii=False
|
||||||
|
)
|
||||||
|
},
|
||||||
|
}
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"finish_reason": None,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
state.tool_call_index += 1
|
||||||
|
|
||||||
|
# 处理结束
|
||||||
|
if finish_reason:
|
||||||
|
openai_finish_reason = self._convert_finish_reason(finish_reason)
|
||||||
|
events.append(
|
||||||
|
{
|
||||||
|
"id": chunk_id,
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": int(time.time()),
|
||||||
|
"model": state.model,
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"delta": {},
|
||||||
|
"finish_reason": openai_finish_reason,
|
||||||
|
}
|
||||||
|
],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
return events
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"ClaudeToGeminiConverter",
|
"ClaudeToGeminiConverter",
|
||||||
@@ -7,11 +7,11 @@ OpenAI -> Claude 格式转换器
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import time
|
|
||||||
import uuid
|
import uuid
|
||||||
from typing import Any, Dict, List, Optional, Union
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
from pydantic import BaseModel
|
if TYPE_CHECKING:
|
||||||
|
from src.core.api_format.conversion.state import StreamConversionState
|
||||||
|
|
||||||
|
|
||||||
class OpenAIToClaudeConverter:
|
class OpenAIToClaudeConverter:
|
||||||
@@ -48,7 +48,7 @@ class OpenAIToClaudeConverter:
|
|||||||
|
|
||||||
# ==================== 请求转换 ====================
|
# ==================== 请求转换 ====================
|
||||||
|
|
||||||
def convert_request(self, request: Union[Dict[str, Any], BaseModel]) -> Dict[str, Any]:
|
def convert_request(self, request: Union[Dict[str, Any], Any]) -> Dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
将 OpenAI 请求转换为 Claude 格式
|
将 OpenAI 请求转换为 Claude 格式
|
||||||
|
|
||||||
@@ -370,22 +370,23 @@ class OpenAIToClaudeConverter:
|
|||||||
def convert_stream_chunk(
|
def convert_stream_chunk(
|
||||||
self,
|
self,
|
||||||
chunk: Dict[str, Any],
|
chunk: Dict[str, Any],
|
||||||
model: str = "",
|
state: Optional["StreamConversionState"] = None,
|
||||||
message_id: Optional[str] = None,
|
|
||||||
message_started: bool = False,
|
|
||||||
) -> List[Dict[str, Any]]:
|
) -> List[Dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
将 OpenAI SSE chunk 转换为 Claude SSE 事件
|
将 OpenAI SSE chunk 转换为 Claude SSE 事件
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
chunk: OpenAI SSE chunk
|
chunk: OpenAI SSE chunk
|
||||||
model: 模型名称
|
state: 流式转换状态
|
||||||
message_id: 消息 ID
|
|
||||||
message_started: 是否已发送 message_start
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Claude SSE 事件列表
|
Claude SSE 事件列表
|
||||||
"""
|
"""
|
||||||
|
from src.core.api_format.conversion.state import StreamConversionState
|
||||||
|
|
||||||
|
if state is None:
|
||||||
|
state = StreamConversionState()
|
||||||
|
|
||||||
events: List[Dict[str, Any]] = []
|
events: List[Dict[str, Any]] = []
|
||||||
|
|
||||||
choices = chunk.get("choices", [])
|
choices = chunk.get("choices", [])
|
||||||
@@ -398,8 +399,8 @@ class OpenAIToClaudeConverter:
|
|||||||
|
|
||||||
# 处理角色(第一个 chunk)
|
# 处理角色(第一个 chunk)
|
||||||
role = delta.get("role")
|
role = delta.get("role")
|
||||||
if role and not message_started:
|
if role and not state.message_started:
|
||||||
msg_id = message_id or f"msg_{uuid.uuid4().hex[:8]}"
|
msg_id = state.message_id or f"msg_{uuid.uuid4().hex[:8]}"
|
||||||
events.append(
|
events.append(
|
||||||
{
|
{
|
||||||
"type": "message_start",
|
"type": "message_start",
|
||||||
@@ -407,13 +408,14 @@ class OpenAIToClaudeConverter:
|
|||||||
"id": msg_id,
|
"id": msg_id,
|
||||||
"type": "message",
|
"type": "message",
|
||||||
"role": role,
|
"role": role,
|
||||||
"model": model,
|
"model": state.model,
|
||||||
"content": [],
|
"content": [],
|
||||||
"stop_reason": None,
|
"stop_reason": None,
|
||||||
"stop_sequence": None,
|
"stop_sequence": None,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
state.message_started = True
|
||||||
|
|
||||||
# 处理文本内容
|
# 处理文本内容
|
||||||
content_delta = delta.get("content")
|
content_delta = delta.get("content")
|
||||||
27
src/core/api_format/conversion/exceptions.py
Normal file
27
src/core/api_format/conversion/exceptions.py
Normal file
@@ -0,0 +1,27 @@
|
|||||||
|
"""
|
||||||
|
格式转换异常
|
||||||
|
|
||||||
|
用于严格模式转换失败时抛出,让编排器可以尝试下一个候选。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
|
||||||
|
class FormatConversionError(Exception):
|
||||||
|
"""
|
||||||
|
格式转换失败异常
|
||||||
|
|
||||||
|
在严格模式下,转换失败会抛出此异常,
|
||||||
|
让 Orchestrator 可以捕获并尝试下一个候选。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, source_format: str, target_format: str, message: str) -> None:
|
||||||
|
self.source_format = source_format
|
||||||
|
self.target_format = target_format
|
||||||
|
self.message = message
|
||||||
|
super().__init__(f"格式转换失败 ({source_format} -> {target_format}): {message}")
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"FormatConversionError",
|
||||||
|
]
|
||||||
54
src/core/api_format/conversion/protocols.py
Normal file
54
src/core/api_format/conversion/protocols.py
Normal file
@@ -0,0 +1,54 @@
|
|||||||
|
"""
|
||||||
|
转换器协议定义
|
||||||
|
|
||||||
|
定义转换器必须实现的方法签名,用于类型检查和文档说明。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any, Dict, List, Protocol, runtime_checkable
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class RequestConverter(Protocol):
|
||||||
|
"""请求转换器协议"""
|
||||||
|
|
||||||
|
def convert_request(self, request: Dict[str, Any]) -> Dict[str, Any]: ...
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class ResponseConverter(Protocol):
|
||||||
|
"""响应转换器协议"""
|
||||||
|
|
||||||
|
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]: ...
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class StreamChunkConverter(Protocol):
|
||||||
|
"""
|
||||||
|
流式响应块转换器协议
|
||||||
|
|
||||||
|
统一签名:(chunk, state) -> List[Dict]
|
||||||
|
|
||||||
|
说明:
|
||||||
|
- chunk: 单个流式事件/块
|
||||||
|
- state: 跨 chunk 的转换状态(StreamConversionState 或 GeminiStreamConversionState)
|
||||||
|
- 返回: 转换后的事件列表(可能 0-N 个)
|
||||||
|
|
||||||
|
注意:
|
||||||
|
- 所有流式转换器必须实现此签名
|
||||||
|
- state 参数用于跨 chunk 维护状态(如累积文本、消息 ID 等)
|
||||||
|
"""
|
||||||
|
|
||||||
|
def convert_stream_chunk(
|
||||||
|
self,
|
||||||
|
chunk: Dict[str, Any],
|
||||||
|
state: Any,
|
||||||
|
) -> List[Dict[str, Any]]: ...
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"RequestConverter",
|
||||||
|
"ResponseConverter",
|
||||||
|
"StreamChunkConverter",
|
||||||
|
]
|
||||||
390
src/core/api_format/conversion/registry.py
Normal file
390
src/core/api_format/conversion/registry.py
Normal file
@@ -0,0 +1,390 @@
|
|||||||
|
"""
|
||||||
|
格式转换器注册表(核心层)
|
||||||
|
|
||||||
|
自动管理不同 API 格式之间的转换器,支持:
|
||||||
|
- 请求转换:客户端格式 → Provider 格式
|
||||||
|
- 响应转换:Provider 格式 → 客户端格式
|
||||||
|
|
||||||
|
说明:
|
||||||
|
- 该注册表位于 core 层,避免 services 依赖 api/handlers。
|
||||||
|
- 具体转换器的注册(例如 Claude/OpenAI/Gemini)应由应用启动层完成,
|
||||||
|
或由 api 层的 bootstrap 逻辑完成,以保持依赖方向:api -> core。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union
|
||||||
|
|
||||||
|
from src.core.logger import logger
|
||||||
|
|
||||||
|
from .exceptions import FormatConversionError
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from .state import GeminiStreamConversionState, StreamConversionState
|
||||||
|
|
||||||
|
|
||||||
|
class FormatConverterRegistry:
|
||||||
|
"""
|
||||||
|
格式转换器注册表
|
||||||
|
|
||||||
|
管理不同 API 格式之间的双向转换器
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
# key: (source_format, target_format), value: converter instance
|
||||||
|
self._converters: Dict[Tuple[str, str], Any] = {}
|
||||||
|
|
||||||
|
def register(
|
||||||
|
self,
|
||||||
|
source_format: str,
|
||||||
|
target_format: str,
|
||||||
|
converter: Any,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
注册格式转换器
|
||||||
|
|
||||||
|
Args:
|
||||||
|
source_format: 源格式(如 "CLAUDE", "OPENAI", "GEMINI")
|
||||||
|
target_format: 目标格式
|
||||||
|
converter: 转换器实例(需要有 convert_request/convert_response 方法)
|
||||||
|
"""
|
||||||
|
key = (source_format.upper(), target_format.upper())
|
||||||
|
self._converters[key] = converter
|
||||||
|
logger.info(f"[ConverterRegistry] 注册转换器: {source_format} -> {target_format}")
|
||||||
|
|
||||||
|
def get_converter(
|
||||||
|
self,
|
||||||
|
source_format: str,
|
||||||
|
target_format: str,
|
||||||
|
) -> Optional[Any]:
|
||||||
|
"""
|
||||||
|
获取转换器
|
||||||
|
|
||||||
|
Args:
|
||||||
|
source_format: 源格式
|
||||||
|
target_format: 目标格式
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
转换器实例,如果不存在返回 None
|
||||||
|
"""
|
||||||
|
key = (source_format.upper(), target_format.upper())
|
||||||
|
return self._converters.get(key)
|
||||||
|
|
||||||
|
def has_converter(
|
||||||
|
self,
|
||||||
|
source_format: str,
|
||||||
|
target_format: str,
|
||||||
|
) -> bool:
|
||||||
|
"""检查是否存在转换器"""
|
||||||
|
key = (source_format.upper(), target_format.upper())
|
||||||
|
return key in self._converters
|
||||||
|
|
||||||
|
def convert_request(
|
||||||
|
self,
|
||||||
|
request: Dict[str, Any],
|
||||||
|
source_format: str,
|
||||||
|
target_format: str,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""
|
||||||
|
转换请求
|
||||||
|
|
||||||
|
Args:
|
||||||
|
request: 原始请求字典
|
||||||
|
source_format: 源格式(客户端格式)
|
||||||
|
target_format: 目标格式(Provider 格式)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
转换后的请求字典,如果无需转换或没有转换器则返回原始请求
|
||||||
|
"""
|
||||||
|
# 同格式无需转换
|
||||||
|
if source_format.upper() == target_format.upper():
|
||||||
|
return request
|
||||||
|
|
||||||
|
converter = self.get_converter(source_format, target_format)
|
||||||
|
if converter is None:
|
||||||
|
logger.warning(
|
||||||
|
f"[ConverterRegistry] 未找到请求转换器: {source_format} -> {target_format},返回原始请求"
|
||||||
|
)
|
||||||
|
return request
|
||||||
|
|
||||||
|
if not hasattr(converter, "convert_request"):
|
||||||
|
logger.warning(
|
||||||
|
f"[ConverterRegistry] 转换器缺少 convert_request 方法: {source_format} -> {target_format}"
|
||||||
|
)
|
||||||
|
return request
|
||||||
|
|
||||||
|
try:
|
||||||
|
converted: Dict[str, Any] = converter.convert_request(request)
|
||||||
|
logger.debug(f"[ConverterRegistry] 请求转换成功: {source_format} -> {target_format}")
|
||||||
|
return converted
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"[ConverterRegistry] 请求转换失败: {source_format} -> {target_format}: {e}")
|
||||||
|
return request
|
||||||
|
|
||||||
|
def convert_response(
|
||||||
|
self,
|
||||||
|
response: Dict[str, Any],
|
||||||
|
source_format: str,
|
||||||
|
target_format: str,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""
|
||||||
|
转换响应
|
||||||
|
|
||||||
|
Args:
|
||||||
|
response: 原始响应字典
|
||||||
|
source_format: 源格式(Provider 格式)
|
||||||
|
target_format: 目标格式(客户端格式)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
转换后的响应字典,如果无需转换或没有转换器则返回原始响应
|
||||||
|
"""
|
||||||
|
# 同格式无需转换
|
||||||
|
if source_format.upper() == target_format.upper():
|
||||||
|
return response
|
||||||
|
|
||||||
|
converter = self.get_converter(source_format, target_format)
|
||||||
|
if converter is None:
|
||||||
|
logger.warning(
|
||||||
|
f"[ConverterRegistry] 未找到响应转换器: {source_format} -> {target_format},返回原始响应"
|
||||||
|
)
|
||||||
|
return response
|
||||||
|
|
||||||
|
if not hasattr(converter, "convert_response"):
|
||||||
|
logger.warning(
|
||||||
|
f"[ConverterRegistry] 转换器缺少 convert_response 方法: {source_format} -> {target_format}"
|
||||||
|
)
|
||||||
|
return response
|
||||||
|
|
||||||
|
try:
|
||||||
|
converted: Dict[str, Any] = converter.convert_response(response)
|
||||||
|
logger.debug(f"[ConverterRegistry] 响应转换成功: {source_format} -> {target_format}")
|
||||||
|
return converted
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"[ConverterRegistry] 响应转换失败: {source_format} -> {target_format}: {e}")
|
||||||
|
return response
|
||||||
|
|
||||||
|
def convert_stream_chunk(
|
||||||
|
self,
|
||||||
|
chunk: Dict[str, Any],
|
||||||
|
source_format: str,
|
||||||
|
target_format: str,
|
||||||
|
state: Optional[Union["StreamConversionState", "GeminiStreamConversionState"]] = None,
|
||||||
|
) -> list[Dict[str, Any]]:
|
||||||
|
"""
|
||||||
|
转换流式响应块
|
||||||
|
|
||||||
|
Args:
|
||||||
|
chunk: 原始流式响应块
|
||||||
|
source_format: 源格式(Provider 格式)
|
||||||
|
target_format: 目标格式(客户端格式)
|
||||||
|
state: 流式转换状态(StreamConversionState 或 GeminiStreamConversionState)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
转换后的事件列表(可能 0-N 个),失败时返回原始 chunk 的单元素列表
|
||||||
|
"""
|
||||||
|
# 同格式无需转换
|
||||||
|
if source_format.upper() == target_format.upper():
|
||||||
|
return [chunk]
|
||||||
|
|
||||||
|
converter = self.get_converter(source_format, target_format)
|
||||||
|
if converter is None:
|
||||||
|
return [chunk]
|
||||||
|
|
||||||
|
# 使用流式转换方法
|
||||||
|
if hasattr(converter, "convert_stream_chunk"):
|
||||||
|
try:
|
||||||
|
result: list[Dict[str, Any]] = converter.convert_stream_chunk(chunk, state)
|
||||||
|
return result
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"[ConverterRegistry] 流式块转换失败: {source_format} -> {target_format}: {e}")
|
||||||
|
return [chunk]
|
||||||
|
|
||||||
|
# 降级到普通响应转换(作为单个事件返回)
|
||||||
|
if hasattr(converter, "convert_response"):
|
||||||
|
try:
|
||||||
|
result = converter.convert_response(chunk)
|
||||||
|
return [result]
|
||||||
|
except Exception:
|
||||||
|
return [chunk]
|
||||||
|
|
||||||
|
return [chunk]
|
||||||
|
|
||||||
|
def list_converters(self) -> list[Tuple[str, str]]:
|
||||||
|
"""列出所有已注册的转换器"""
|
||||||
|
return list(self._converters.keys())
|
||||||
|
|
||||||
|
# ========== 能力查询方法 ==========
|
||||||
|
|
||||||
|
def can_convert_request(self, source: str, target: str) -> bool:
|
||||||
|
"""检查是否支持请求转换"""
|
||||||
|
converter = self.get_converter(source, target)
|
||||||
|
return converter is not None and hasattr(converter, "convert_request")
|
||||||
|
|
||||||
|
def can_convert_response(self, source: str, target: str) -> bool:
|
||||||
|
"""检查是否支持响应转换"""
|
||||||
|
converter = self.get_converter(source, target)
|
||||||
|
return converter is not None and hasattr(converter, "convert_response")
|
||||||
|
|
||||||
|
def can_convert_stream(self, source: str, target: str) -> bool:
|
||||||
|
"""检查是否支持流式转换"""
|
||||||
|
converter = self.get_converter(source, target)
|
||||||
|
if converter is None:
|
||||||
|
return False
|
||||||
|
return hasattr(converter, "convert_stream_chunk")
|
||||||
|
|
||||||
|
def can_convert_full(
|
||||||
|
self,
|
||||||
|
source: str,
|
||||||
|
target: str,
|
||||||
|
require_stream: bool = False,
|
||||||
|
) -> bool:
|
||||||
|
"""
|
||||||
|
检查是否支持完整的双向转换
|
||||||
|
|
||||||
|
对于跨格式请求,需要:
|
||||||
|
1. 请求转换:source -> target
|
||||||
|
2. 响应转换:target -> source(注意方向相反)
|
||||||
|
3. 流式转换(如果 require_stream=True):target -> source
|
||||||
|
|
||||||
|
Args:
|
||||||
|
source: 客户端格式
|
||||||
|
target: Provider 格式
|
||||||
|
require_stream: 是否要求支持流式转换
|
||||||
|
"""
|
||||||
|
# 请求:client -> provider
|
||||||
|
if not self.can_convert_request(source, target):
|
||||||
|
return False
|
||||||
|
# 响应:provider -> client(方向相反)
|
||||||
|
if not self.can_convert_response(target, source):
|
||||||
|
return False
|
||||||
|
# 流式:provider -> client(方向相反)
|
||||||
|
if require_stream and not self.can_convert_stream(target, source):
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
def get_supported_targets(self, source: str) -> list[str]:
|
||||||
|
"""获取指定源格式支持转换到的目标格式列表"""
|
||||||
|
source_upper = source.upper()
|
||||||
|
return [target for (src, target) in self._converters.keys() if src == source_upper]
|
||||||
|
|
||||||
|
# ========== 严格模式方法 ==========
|
||||||
|
|
||||||
|
def convert_request_strict(
|
||||||
|
self,
|
||||||
|
request: Dict[str, Any],
|
||||||
|
source_format: str,
|
||||||
|
target_format: str,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""
|
||||||
|
严格模式请求转换 - 失败时抛出异常
|
||||||
|
|
||||||
|
用于需要故障转移的场景:转换失败会抛出 FormatConversionError,
|
||||||
|
让 Orchestrator 可以尝试下一个候选。
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
FormatConversionError: 转换失败时抛出
|
||||||
|
"""
|
||||||
|
# 同格式无需转换
|
||||||
|
if source_format.upper() == target_format.upper():
|
||||||
|
return request
|
||||||
|
|
||||||
|
converter = self.get_converter(source_format, target_format)
|
||||||
|
if converter is None:
|
||||||
|
raise FormatConversionError(source_format, target_format, "未找到转换器")
|
||||||
|
|
||||||
|
if not hasattr(converter, "convert_request"):
|
||||||
|
raise FormatConversionError(source_format, target_format, "转换器缺少 convert_request 方法")
|
||||||
|
|
||||||
|
try:
|
||||||
|
converted: Dict[str, Any] = converter.convert_request(request)
|
||||||
|
logger.debug(f"[ConverterRegistry] 请求转换成功: {source_format} -> {target_format}")
|
||||||
|
return converted
|
||||||
|
except FormatConversionError:
|
||||||
|
raise
|
||||||
|
except Exception as e:
|
||||||
|
raise FormatConversionError(source_format, target_format, str(e)) from e
|
||||||
|
|
||||||
|
def convert_response_strict(
|
||||||
|
self,
|
||||||
|
response: Dict[str, Any],
|
||||||
|
source_format: str,
|
||||||
|
target_format: str,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""
|
||||||
|
严格模式响应转换 - 失败时抛出异常
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
FormatConversionError: 转换失败时抛出
|
||||||
|
"""
|
||||||
|
if source_format.upper() == target_format.upper():
|
||||||
|
return response
|
||||||
|
|
||||||
|
converter = self.get_converter(source_format, target_format)
|
||||||
|
if converter is None:
|
||||||
|
raise FormatConversionError(source_format, target_format, "未找到转换器")
|
||||||
|
|
||||||
|
if not hasattr(converter, "convert_response"):
|
||||||
|
raise FormatConversionError(source_format, target_format, "转换器缺少 convert_response 方法")
|
||||||
|
|
||||||
|
try:
|
||||||
|
converted: Dict[str, Any] = converter.convert_response(response)
|
||||||
|
logger.debug(f"[ConverterRegistry] 响应转换成功: {source_format} -> {target_format}")
|
||||||
|
return converted
|
||||||
|
except FormatConversionError:
|
||||||
|
raise
|
||||||
|
except Exception as e:
|
||||||
|
raise FormatConversionError(source_format, target_format, str(e)) from e
|
||||||
|
|
||||||
|
def convert_stream_chunk_strict(
|
||||||
|
self,
|
||||||
|
chunk: Dict[str, Any],
|
||||||
|
source_format: str,
|
||||||
|
target_format: str,
|
||||||
|
state: Optional[Union["StreamConversionState", "GeminiStreamConversionState"]] = None,
|
||||||
|
) -> list[Dict[str, Any]]:
|
||||||
|
"""
|
||||||
|
严格模式流式块转换 - 失败时抛出异常
|
||||||
|
|
||||||
|
Args:
|
||||||
|
chunk: 流式响应块
|
||||||
|
source_format: 源格式
|
||||||
|
target_format: 目标格式
|
||||||
|
state: 流式转换状态(StreamConversionState 或 GeminiStreamConversionState)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
转换后的事件列表(可能 0-N 个)
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
FormatConversionError: 转换失败时抛出
|
||||||
|
"""
|
||||||
|
if source_format.upper() == target_format.upper():
|
||||||
|
return [chunk]
|
||||||
|
|
||||||
|
converter = self.get_converter(source_format, target_format)
|
||||||
|
if converter is None:
|
||||||
|
raise FormatConversionError(source_format, target_format, "未找到转换器")
|
||||||
|
|
||||||
|
if not hasattr(converter, "convert_stream_chunk"):
|
||||||
|
raise FormatConversionError(
|
||||||
|
source_format, target_format, "转换器缺少 convert_stream_chunk 方法"
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
result: list[Dict[str, Any]] = converter.convert_stream_chunk(chunk, state)
|
||||||
|
return result
|
||||||
|
except FormatConversionError:
|
||||||
|
raise
|
||||||
|
except Exception as e:
|
||||||
|
raise FormatConversionError(source_format, target_format, f"流式块转换失败: {e}") from e
|
||||||
|
|
||||||
|
|
||||||
|
# 全局单例
|
||||||
|
converter_registry = FormatConverterRegistry()
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"FormatConverterRegistry",
|
||||||
|
"converter_registry",
|
||||||
|
"FormatConversionError",
|
||||||
|
]
|
||||||
|
|
||||||
69
src/core/api_format/conversion/state.py
Normal file
69
src/core/api_format/conversion/state.py
Normal file
@@ -0,0 +1,69 @@
|
|||||||
|
"""
|
||||||
|
流式转换状态类
|
||||||
|
|
||||||
|
用于在多个 chunk 之间维护转换上下文,例如:
|
||||||
|
- 是否已发送 message_start 事件
|
||||||
|
- 累积的文本(用于计算增量)
|
||||||
|
- 当前内容块索引
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class StreamConversionState:
|
||||||
|
"""
|
||||||
|
SSE 流式转换状态(Claude <-> OpenAI)
|
||||||
|
|
||||||
|
用于跨 chunk 维护转换状态,确保生成正确的事件序列。
|
||||||
|
"""
|
||||||
|
|
||||||
|
message_id: str = ""
|
||||||
|
model: str = ""
|
||||||
|
message_started: bool = False
|
||||||
|
content_block_started: bool = False
|
||||||
|
current_tool_index: int = 0
|
||||||
|
|
||||||
|
def reset(self) -> None:
|
||||||
|
"""重置状态(重试时调用)"""
|
||||||
|
self.message_started = False
|
||||||
|
self.content_block_started = False
|
||||||
|
self.current_tool_index = 0
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class GeminiStreamConversionState:
|
||||||
|
"""
|
||||||
|
Gemini 流式转换状态(JSON 数组格式)
|
||||||
|
|
||||||
|
与 Claude/OpenAI 的 SSE 不同,Gemini 需要追踪额外状态:
|
||||||
|
1. 累积文本(用于计算真正的增量)
|
||||||
|
2. 内容块索引(用于工具调用)
|
||||||
|
3. 是否已发送 message_start
|
||||||
|
"""
|
||||||
|
|
||||||
|
message_id: str = ""
|
||||||
|
model: str = ""
|
||||||
|
accumulated_text: str = "" # 累积的文本(用于计算增量)
|
||||||
|
message_started: bool = False
|
||||||
|
content_block_started: bool = False
|
||||||
|
current_block_index: int = 0
|
||||||
|
tool_call_index: int = 0 # 工具调用计数
|
||||||
|
has_sent_usage: bool = False # 是否已发送 usage
|
||||||
|
|
||||||
|
def reset(self) -> None:
|
||||||
|
"""重置状态(重试时调用)"""
|
||||||
|
self.accumulated_text = ""
|
||||||
|
self.message_started = False
|
||||||
|
self.content_block_started = False
|
||||||
|
self.current_block_index = 0
|
||||||
|
self.tool_call_index = 0
|
||||||
|
self.has_sent_usage = False
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"StreamConversionState",
|
||||||
|
"GeminiStreamConversionState",
|
||||||
|
]
|
||||||
183
src/core/api_format/detection.py
Normal file
183
src/core/api_format/detection.py
Normal file
@@ -0,0 +1,183 @@
|
|||||||
|
"""
|
||||||
|
API 格式检测
|
||||||
|
|
||||||
|
提供从请求头、响应内容等检测 API 格式的函数。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING, Dict, Optional, Tuple
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from starlette.requests import Request
|
||||||
|
|
||||||
|
from src.core.api_format.enums import APIFormat
|
||||||
|
from src.core.api_format.metadata import API_FORMAT_DEFINITIONS, ApiFormatDefinition
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_api_key_by_definition(
|
||||||
|
headers: Dict[str, str],
|
||||||
|
query_params: Optional[Dict[str, str]],
|
||||||
|
definition: ApiFormatDefinition,
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
根据格式定义从请求中提取 API Key
|
||||||
|
|
||||||
|
Args:
|
||||||
|
headers: 请求头字典(key 小写)
|
||||||
|
query_params: 查询参数字典(可选)
|
||||||
|
definition: API 格式定义
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
提取到的 API Key,或 None
|
||||||
|
"""
|
||||||
|
auth_header = definition.auth_header.lower()
|
||||||
|
auth_type = definition.auth_type
|
||||||
|
|
||||||
|
header_value = headers.get(auth_header)
|
||||||
|
if not header_value:
|
||||||
|
# Gemini 还支持 ?key= 参数
|
||||||
|
if definition.api_format in (APIFormat.GEMINI, APIFormat.GEMINI_CLI):
|
||||||
|
return query_params.get("key") if query_params else None
|
||||||
|
return None
|
||||||
|
|
||||||
|
if auth_type == "bearer":
|
||||||
|
# Bearer token: "Bearer xxx"
|
||||||
|
if header_value.lower().startswith("bearer "):
|
||||||
|
return header_value[7:].strip()
|
||||||
|
return None
|
||||||
|
else:
|
||||||
|
# header 类型: 直接使用值
|
||||||
|
return header_value
|
||||||
|
|
||||||
|
|
||||||
|
def detect_format_from_request(
|
||||||
|
headers: Dict[str, str],
|
||||||
|
query_params: Optional[Dict[str, str]] = None,
|
||||||
|
) -> Tuple[APIFormat, Optional[str]]:
|
||||||
|
"""
|
||||||
|
从请求头检测 API 格式和 API Key
|
||||||
|
|
||||||
|
检测优先级:
|
||||||
|
1. x-api-key + anthropic-version -> Claude
|
||||||
|
2. x-goog-api-key 或 ?key= -> Gemini
|
||||||
|
3. Authorization: Bearer -> OpenAI (默认)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
headers: 请求头字典(key 应为小写)
|
||||||
|
query_params: 查询参数字典(可选)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(APIFormat, api_key) 元组
|
||||||
|
"""
|
||||||
|
# Claude: x-api-key + anthropic-version (必须同时存在)
|
||||||
|
claude_def = API_FORMAT_DEFINITIONS[APIFormat.CLAUDE]
|
||||||
|
claude_key = _extract_api_key_by_definition(headers, query_params, claude_def)
|
||||||
|
if claude_key and headers.get("anthropic-version"):
|
||||||
|
return APIFormat.CLAUDE, claude_key
|
||||||
|
|
||||||
|
# Gemini: x-goog-api-key (header 类型) 或 ?key=
|
||||||
|
gemini_def = API_FORMAT_DEFINITIONS[APIFormat.GEMINI]
|
||||||
|
gemini_key = _extract_api_key_by_definition(headers, query_params, gemini_def)
|
||||||
|
if gemini_key:
|
||||||
|
return APIFormat.GEMINI, gemini_key
|
||||||
|
|
||||||
|
# OpenAI: Authorization: Bearer (默认)
|
||||||
|
# 注意: 如果只有 x-api-key 但没有 anthropic-version,也走 OpenAI 格式
|
||||||
|
openai_def = API_FORMAT_DEFINITIONS[APIFormat.OPENAI]
|
||||||
|
openai_key = _extract_api_key_by_definition(headers, query_params, openai_def)
|
||||||
|
# 如果 OpenAI 格式没有 key,但有 x-api-key,也用它(兼容)
|
||||||
|
if not openai_key and claude_key:
|
||||||
|
openai_key = claude_key
|
||||||
|
return APIFormat.OPENAI, openai_key
|
||||||
|
|
||||||
|
|
||||||
|
def detect_format_and_key_from_starlette(
|
||||||
|
request: "Request",
|
||||||
|
) -> Tuple[str, Optional[str]]:
|
||||||
|
"""
|
||||||
|
从 Starlette Request 对象检测 API 格式和 API Key
|
||||||
|
|
||||||
|
这是一个便捷函数,用于直接处理 Starlette/FastAPI 请求对象。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
request: Starlette Request 对象
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(format_name, api_key) 元组,format_name 为小写字符串
|
||||||
|
"""
|
||||||
|
# 规范化 headers 为小写
|
||||||
|
headers = {k.lower(): v for k, v in request.headers.items()}
|
||||||
|
query_params = dict(request.query_params)
|
||||||
|
|
||||||
|
api_format, api_key = detect_format_from_request(headers, query_params)
|
||||||
|
|
||||||
|
# 返回小写格式名
|
||||||
|
format_name = api_format.value.lower()
|
||||||
|
return format_name, api_key
|
||||||
|
|
||||||
|
|
||||||
|
def detect_format_from_response(
|
||||||
|
response_data: dict,
|
||||||
|
) -> Optional[APIFormat]:
|
||||||
|
"""
|
||||||
|
从响应内容检测 API 格式
|
||||||
|
|
||||||
|
Args:
|
||||||
|
response_data: 响应 JSON 字典
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
检测到的格式,或 None
|
||||||
|
"""
|
||||||
|
# Claude: 有 type="message" 或特定的 content 结构
|
||||||
|
if response_data.get("type") == "message":
|
||||||
|
return APIFormat.CLAUDE
|
||||||
|
if "content" in response_data and isinstance(response_data["content"], list):
|
||||||
|
first_content = response_data["content"][0] if response_data["content"] else {}
|
||||||
|
if first_content.get("type") in ("text", "tool_use"):
|
||||||
|
return APIFormat.CLAUDE
|
||||||
|
|
||||||
|
# OpenAI: 有 choices 数组
|
||||||
|
if "choices" in response_data:
|
||||||
|
return APIFormat.OPENAI
|
||||||
|
|
||||||
|
# Gemini: 有 candidates 数组
|
||||||
|
if "candidates" in response_data:
|
||||||
|
return APIFormat.GEMINI
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def detect_cli_format_from_path(
|
||||||
|
path: str,
|
||||||
|
base_format: APIFormat,
|
||||||
|
) -> bool:
|
||||||
|
"""
|
||||||
|
根据请求路径检测是否为 CLI 模式
|
||||||
|
|
||||||
|
CLI 模式的特征:
|
||||||
|
- OpenAI CLI: 请求 /responses 路径
|
||||||
|
- Claude CLI: 有特定的路径模式
|
||||||
|
- Gemini CLI: 有特定的路径模式
|
||||||
|
|
||||||
|
Args:
|
||||||
|
path: 请求路径
|
||||||
|
base_format: 基础格式
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
True 如果是 CLI 模式
|
||||||
|
"""
|
||||||
|
# OpenAI CLI 特征: /v1/responses 路径
|
||||||
|
if base_format == APIFormat.OPENAI and "/responses" in path:
|
||||||
|
return True
|
||||||
|
|
||||||
|
# 其他 CLI 模式通常由 Adapter 层根据具体业务逻辑判断
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"detect_format_from_request",
|
||||||
|
"detect_format_and_key_from_starlette",
|
||||||
|
"detect_format_from_response",
|
||||||
|
"detect_cli_format_from_path",
|
||||||
|
]
|
||||||
21
src/core/api_format/enums.py
Normal file
21
src/core/api_format/enums.py
Normal file
@@ -0,0 +1,21 @@
|
|||||||
|
"""
|
||||||
|
API 格式枚举定义
|
||||||
|
|
||||||
|
定义所有支持的 API 格式,决定请求/响应的处理方式。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
|
||||||
|
class APIFormat(Enum):
|
||||||
|
"""API 格式枚举 - 决定请求/响应的处理方式"""
|
||||||
|
|
||||||
|
CLAUDE = "CLAUDE" # Claude API 格式
|
||||||
|
CLAUDE_CLI = "CLAUDE_CLI" # Claude CLI API 格式(使用 authorization: Bearer)
|
||||||
|
OPENAI = "OPENAI" # OpenAI API 格式
|
||||||
|
OPENAI_CLI = "OPENAI_CLI" # OpenAI CLI/Responses API 格式(用于 Claude Code 等客户端)
|
||||||
|
GEMINI = "GEMINI" # Google Gemini API 格式
|
||||||
|
GEMINI_CLI = "GEMINI_CLI" # Gemini CLI API 格式
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["APIFormat"]
|
||||||
@@ -14,8 +14,8 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from typing import AbstractSet, Any, Dict, FrozenSet, Optional, Set
|
from typing import AbstractSet, Any, Dict, FrozenSet, Optional, Set
|
||||||
|
|
||||||
from .api_format_metadata import get_auth_config, get_extra_headers, get_protected_keys
|
from src.core.api_format.enums import APIFormat
|
||||||
from .enums import APIFormat
|
from src.core.api_format.metadata import get_auth_config, get_extra_headers, get_protected_keys
|
||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
@@ -1,18 +1,16 @@
|
|||||||
"""
|
"""
|
||||||
集中维护 API 格式的元数据,避免新增格式时到处修改常量。
|
API 格式元数据定义
|
||||||
|
|
||||||
此模块与 src/formats/ 的 FormatProtocol 系统配合使用:
|
集中维护 API 格式的元数据,避免新增格式时到处修改常量。
|
||||||
- api_format_metadata: 定义格式的元数据(别名、默认路径)
|
|
||||||
- src/formats/: 定义格式的协议实现(解析、转换、验证)
|
|
||||||
|
|
||||||
使用方式:
|
使用方式:
|
||||||
# 解析格式别名
|
# 解析格式别名
|
||||||
from src.core.api_format_metadata import resolve_api_format
|
from src.core.api_format import resolve_api_format
|
||||||
api_format = resolve_api_format("claude") # -> APIFormat.CLAUDE
|
api_format = resolve_api_format("claude") # -> APIFormat.CLAUDE
|
||||||
|
|
||||||
# 获取格式协议
|
# 获取格式定义
|
||||||
from src.core.api_format_metadata import get_format_protocol
|
from src.core.api_format import get_api_format_definition
|
||||||
protocol = get_format_protocol(APIFormat.CLAUDE) # -> ClaudeProtocol
|
definition = get_api_format_definition(APIFormat.CLAUDE)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -263,7 +261,7 @@ def resolve_api_format(
|
|||||||
return default
|
return default
|
||||||
|
|
||||||
|
|
||||||
def register_api_format_definition(definition: ApiFormatDefinition, *, override: bool = False):
|
def register_api_format_definition(definition: ApiFormatDefinition, *, override: bool = False) -> None:
|
||||||
"""
|
"""
|
||||||
注册或覆盖 API 格式定义,允许运行时扩展。
|
注册或覆盖 API 格式定义,允许运行时扩展。
|
||||||
|
|
||||||
@@ -278,7 +276,7 @@ def register_api_format_definition(definition: ApiFormatDefinition, *, override:
|
|||||||
_refresh_metadata_cache()
|
_refresh_metadata_cache()
|
||||||
|
|
||||||
|
|
||||||
def _refresh_metadata_cache():
|
def _refresh_metadata_cache() -> None:
|
||||||
"""更新别名缓存,供注册函数调用。"""
|
"""更新别名缓存,供注册函数调用。"""
|
||||||
_alias_lookup_cache.cache_clear()
|
_alias_lookup_cache.cache_clear()
|
||||||
|
|
||||||
@@ -293,21 +291,9 @@ def normalize_alias_value(value: str) -> str:
|
|||||||
return text.strip("_")
|
return text.strip("_")
|
||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
# is_cli_format 和 is_cli_api_format 已移至 utils.py
|
||||||
# 格式判断工具
|
# 为保持兼容性,从 utils 重新导出
|
||||||
# =============================================================================
|
from src.core.api_format.utils import is_cli_format # noqa: E402
|
||||||
|
|
||||||
|
# is_cli_api_format 是 is_cli_format 的别名(接受 APIFormat 枚举)
|
||||||
def is_cli_api_format(api_format: APIFormat) -> bool:
|
is_cli_api_format = is_cli_format
|
||||||
"""
|
|
||||||
判断是否为 CLI 透传格式。
|
|
||||||
|
|
||||||
Args:
|
|
||||||
api_format: APIFormat 枚举值
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True 如果是 CLI 格式
|
|
||||||
"""
|
|
||||||
from src.api.handlers.base.parsers import is_cli_format
|
|
||||||
|
|
||||||
return is_cli_format(api_format.value)
|
|
||||||
106
src/core/api_format/utils.py
Normal file
106
src/core/api_format/utils.py
Normal file
@@ -0,0 +1,106 @@
|
|||||||
|
"""
|
||||||
|
API 格式工具函数
|
||||||
|
|
||||||
|
提供格式判断、规范化等工具函数,供整个项目使用。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING, Optional, Union
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from src.core.api_format.enums import APIFormat
|
||||||
|
|
||||||
|
|
||||||
|
def is_cli_format(format_id: Union[str, "APIFormat", None]) -> bool:
|
||||||
|
"""
|
||||||
|
判断是否为 CLI 透传格式
|
||||||
|
|
||||||
|
CLI 格式以 _CLI 结尾,不参与格式转换,请求直接透传。
|
||||||
|
|
||||||
|
Args:
|
||||||
|
format_id: 格式标识符(字符串或 APIFormat 枚举)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
True 如果是 CLI 格式
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
>>> is_cli_format("CLAUDE_CLI")
|
||||||
|
True
|
||||||
|
>>> is_cli_format("CLAUDE")
|
||||||
|
False
|
||||||
|
>>> is_cli_format(APIFormat.OPENAI_CLI)
|
||||||
|
True
|
||||||
|
"""
|
||||||
|
if format_id is None:
|
||||||
|
return False
|
||||||
|
if hasattr(format_id, "value"):
|
||||||
|
format_id = format_id.value
|
||||||
|
return str(format_id).upper().endswith("_CLI")
|
||||||
|
|
||||||
|
|
||||||
|
def get_base_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
获取基础格式(去除 _CLI 后缀)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
format_id: 格式标识符
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
基础格式字符串,或 None
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
>>> get_base_format("CLAUDE_CLI")
|
||||||
|
"CLAUDE"
|
||||||
|
>>> get_base_format("OPENAI")
|
||||||
|
"OPENAI"
|
||||||
|
"""
|
||||||
|
if format_id is None:
|
||||||
|
return None
|
||||||
|
if hasattr(format_id, "value"):
|
||||||
|
format_id = format_id.value
|
||||||
|
format_str = str(format_id).upper()
|
||||||
|
if format_str.endswith("_CLI"):
|
||||||
|
return format_str[:-4]
|
||||||
|
return format_str
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_format(format_id: Union[str, "APIFormat", None]) -> Optional[str]:
|
||||||
|
"""
|
||||||
|
规范化格式标识符
|
||||||
|
|
||||||
|
Args:
|
||||||
|
format_id: 格式标识符(可能是字符串、枚举或 None)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
大写的格式字符串,或 None
|
||||||
|
"""
|
||||||
|
if format_id is None:
|
||||||
|
return None
|
||||||
|
if hasattr(format_id, "value"):
|
||||||
|
return str(format_id.value).upper()
|
||||||
|
return str(format_id).upper()
|
||||||
|
|
||||||
|
|
||||||
|
def is_same_format(
|
||||||
|
format1: Union[str, "APIFormat", None],
|
||||||
|
format2: Union[str, "APIFormat", None],
|
||||||
|
) -> bool:
|
||||||
|
"""
|
||||||
|
判断两个格式是否相同
|
||||||
|
|
||||||
|
忽略大小写和枚举/字符串差异。
|
||||||
|
"""
|
||||||
|
return normalize_format(format1) == normalize_format(format2)
|
||||||
|
|
||||||
|
|
||||||
|
def is_convertible_format(format_id: Union[str, "APIFormat", None]) -> bool:
|
||||||
|
"""
|
||||||
|
判断是否为可转换格式(非 CLI)
|
||||||
|
|
||||||
|
可转换格式可以与其他格式进行双向转换。
|
||||||
|
CLI 格式为透传模式,不参与转换。
|
||||||
|
"""
|
||||||
|
if format_id is None:
|
||||||
|
return False
|
||||||
|
return not is_cli_format(format_id)
|
||||||
@@ -1,22 +1,13 @@
|
|||||||
"""
|
"""
|
||||||
统一的枚举定义
|
统一的枚举定义
|
||||||
避免重复定义造成的不一致
|
避免重复定义造成的不一致
|
||||||
|
|
||||||
|
注意:APIFormat 已移至 src/core/api_format/enums.py
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
|
|
||||||
|
|
||||||
class APIFormat(Enum):
|
|
||||||
"""API格式枚举 - 决定请求/响应的处理方式"""
|
|
||||||
|
|
||||||
CLAUDE = "CLAUDE" # Claude API 格式
|
|
||||||
CLAUDE_CLI = "CLAUDE_CLI" # Claude CLI API 格式(使用 authorization: Bearer)
|
|
||||||
OPENAI = "OPENAI" # OpenAI API 格式
|
|
||||||
OPENAI_CLI = "OPENAI_CLI" # OpenAI CLI/Responses API 格式(用于 Claude Code 等客户端)
|
|
||||||
GEMINI = "GEMINI" # Google Gemini API 格式
|
|
||||||
GEMINI_CLI = "GEMINI_CLI" # Gemini CLI API 格式
|
|
||||||
|
|
||||||
|
|
||||||
class UserRole(Enum):
|
class UserRole(Enum):
|
||||||
"""用户角色枚举"""
|
"""用户角色枚举"""
|
||||||
|
|
||||||
|
|||||||
@@ -156,7 +156,7 @@ async def lifespan(app: FastAPI):
|
|||||||
|
|
||||||
# 注册格式转换器
|
# 注册格式转换器
|
||||||
logger.info("注册格式转换器...")
|
logger.info("注册格式转换器...")
|
||||||
from src.api.handlers.base.format_converter_registry import register_all_converters
|
from src.core.api_format import register_all_converters
|
||||||
|
|
||||||
register_all_converters()
|
register_all_converters()
|
||||||
|
|
||||||
|
|||||||
@@ -10,7 +10,8 @@ from typing import Any, Dict, List, Optional
|
|||||||
|
|
||||||
from pydantic import BaseModel, Field, field_validator, model_validator
|
from pydantic import BaseModel, Field, field_validator, model_validator
|
||||||
|
|
||||||
from src.core.enums import APIFormat, ProviderBillingType
|
from src.core.api_format import APIFormat
|
||||||
|
from src.core.enums import ProviderBillingType
|
||||||
|
|
||||||
|
|
||||||
class ProxyConfig(BaseModel):
|
class ProxyConfig(BaseModel):
|
||||||
|
|||||||
@@ -49,7 +49,7 @@ class ProviderEndpointCreate(BaseModel):
|
|||||||
@classmethod
|
@classmethod
|
||||||
def validate_api_format(cls, v: str) -> str:
|
def validate_api_format(cls, v: str) -> str:
|
||||||
"""验证 API 格式"""
|
"""验证 API 格式"""
|
||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
|
|
||||||
allowed = [fmt.value for fmt in APIFormat]
|
allowed = [fmt.value for fmt in APIFormat]
|
||||||
v_upper = v.upper()
|
v_upper = v.upper()
|
||||||
@@ -200,7 +200,7 @@ class EndpointAPIKeyCreate(BaseModel):
|
|||||||
if v is None:
|
if v is None:
|
||||||
return v
|
return v
|
||||||
|
|
||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
|
|
||||||
allowed = [fmt.value for fmt in APIFormat]
|
allowed = [fmt.value for fmt in APIFormat]
|
||||||
validated = []
|
validated = []
|
||||||
@@ -335,7 +335,7 @@ class EndpointAPIKeyUpdate(BaseModel):
|
|||||||
if v is None:
|
if v is None:
|
||||||
return v
|
return v
|
||||||
|
|
||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
|
|
||||||
allowed = [fmt.value for fmt in APIFormat]
|
allowed = [fmt.value for fmt in APIFormat]
|
||||||
validated = []
|
validated = []
|
||||||
|
|||||||
2
src/services/cache/aware_scheduler.py
vendored
2
src/services/cache/aware_scheduler.py
vendored
@@ -39,7 +39,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
|
|||||||
|
|
||||||
from sqlalchemy.orm import Session, selectinload
|
from sqlalchemy.orm import Session, selectinload
|
||||||
|
|
||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
from src.core.exceptions import ModelNotSupportedException, ProviderNotAvailableException
|
from src.core.exceptions import ModelNotSupportedException, ProviderNotAvailableException
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.models.database import (
|
from src.models.database import (
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ from src.core.key_capabilities import (
|
|||||||
CapabilityConfigMode,
|
CapabilityConfigMode,
|
||||||
get_user_configurable_capabilities,
|
get_user_configurable_capabilities,
|
||||||
)
|
)
|
||||||
from src.core.headers import get_header_value
|
from src.core.api_format import get_header_value
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
|
|
||||||
# Adapter 检测器类型:接受 headers 和可选的 request_body,返回能力需求字典
|
# Adapter 检测器类型:接受 headers 和可选的 request_body,返回能力需求字典
|
||||||
|
|||||||
@@ -11,8 +11,7 @@ from typing import Dict, Optional
|
|||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat, get_extra_headers_from_endpoint
|
||||||
from src.core.headers import get_extra_headers_from_endpoint
|
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.models.database import ProviderEndpoint
|
from src.models.database import ProviderEndpoint
|
||||||
from src.utils.ssl_utils import get_ssl_context
|
from src.utils.ssl_utils import get_ssl_context
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ from typing import Any, Dict, List, Optional, Tuple
|
|||||||
|
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
from src.core.exceptions import ProviderNotAvailableException
|
from src.core.exceptions import ProviderNotAvailableException
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.models.database import ApiKey
|
from src.models.database import ApiKey
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ from typing import Any, Dict, Optional, Tuple, Union
|
|||||||
import httpx
|
import httpx
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
from src.core.exceptions import (
|
from src.core.exceptions import (
|
||||||
ConcurrencyLimitError,
|
ConcurrencyLimitError,
|
||||||
ProviderAuthException,
|
ProviderAuthException,
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ import httpx
|
|||||||
from redis import Redis
|
from redis import Redis
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
from src.core.error_utils import extract_error_message
|
from src.core.error_utils import extract_error_message
|
||||||
from src.core.exceptions import (
|
from src.core.exceptions import (
|
||||||
ConcurrencyLimitError,
|
ConcurrencyLimitError,
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ from typing import Any, Callable, Optional, Tuple
|
|||||||
|
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.models.database import ApiKey
|
from src.models.database import ApiKey
|
||||||
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
|
||||||
|
|||||||
@@ -6,8 +6,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from typing import Optional, Union
|
from typing import Optional, Union
|
||||||
|
|
||||||
from src.core.api_format_metadata import resolve_api_format
|
from src.core.api_format import APIFormat, resolve_api_format
|
||||||
from src.core.enums import APIFormat
|
|
||||||
|
|
||||||
|
|
||||||
def normalize_api_format(
|
def normalize_api_format(
|
||||||
|
|||||||
@@ -8,8 +8,7 @@
|
|||||||
from typing import TYPE_CHECKING, Any, Dict, Optional
|
from typing import TYPE_CHECKING, Any, Dict, Optional
|
||||||
from urllib.parse import urlencode
|
from urllib.parse import urlencode
|
||||||
|
|
||||||
from src.core.api_format_metadata import get_default_path, resolve_api_format
|
from src.core.api_format import APIFormat, get_default_path, resolve_api_format
|
||||||
from src.core.enums import APIFormat
|
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ from typing import Any, Callable, Optional, Union
|
|||||||
|
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
from src.core.exceptions import ConcurrencyLimitError
|
from src.core.exceptions import ConcurrencyLimitError
|
||||||
from src.core.logger import logger
|
from src.core.logger import logger
|
||||||
from src.services.health.monitor import health_monitor
|
from src.services.health.monitor import health_monitor
|
||||||
|
|||||||
281
tests/api/handlers/base/test_cli_handler_convert.py
Normal file
281
tests/api/handlers/base/test_cli_handler_convert.py
Normal file
@@ -0,0 +1,281 @@
|
|||||||
|
"""
|
||||||
|
CliMessageHandlerBase._convert_sse_line 单元测试
|
||||||
|
|
||||||
|
测试覆盖:
|
||||||
|
1. 基本转换(空行、非 data 行、JSON 解析失败)
|
||||||
|
2. 一入多出场景
|
||||||
|
3. 状态追踪
|
||||||
|
4. 错误处理
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.api.handlers.base.stream_context import StreamContext
|
||||||
|
from src.core.api_format import StreamConversionState
|
||||||
|
|
||||||
|
|
||||||
|
# Mock CliMessageHandlerBase 用于测试
|
||||||
|
class MockCliHandler:
|
||||||
|
"""Mock handler for testing _convert_sse_line"""
|
||||||
|
|
||||||
|
def _convert_sse_line(
|
||||||
|
self,
|
||||||
|
ctx: StreamContext,
|
||||||
|
line: str,
|
||||||
|
events: list,
|
||||||
|
) -> List[str]:
|
||||||
|
"""复制自 CliMessageHandlerBase._convert_sse_line"""
|
||||||
|
from src.core.api_format import converter_registry
|
||||||
|
|
||||||
|
# 如果是空行或特殊控制行,直接返回
|
||||||
|
if not line or line.strip() == "" or line == "data: [DONE]":
|
||||||
|
return [line] if line else []
|
||||||
|
|
||||||
|
# 如果不是 data 行,直接透传
|
||||||
|
if not line.startswith("data:"):
|
||||||
|
return [line]
|
||||||
|
|
||||||
|
# 提取 data 内容
|
||||||
|
data_content = line[5:].strip()
|
||||||
|
|
||||||
|
# 尝试解析 JSON
|
||||||
|
try:
|
||||||
|
data_obj = json.loads(data_content)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
return [line]
|
||||||
|
|
||||||
|
# 初始化流式转换状态
|
||||||
|
if ctx.stream_conversion_state is None:
|
||||||
|
ctx.stream_conversion_state = StreamConversionState(
|
||||||
|
model=ctx.mapped_model or ctx.model,
|
||||||
|
message_id=ctx.response_id or ctx.request_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
provider_format = ctx.provider_api_format or ""
|
||||||
|
client_format = ctx.client_api_format or ""
|
||||||
|
|
||||||
|
try:
|
||||||
|
converted_events = converter_registry.convert_stream_chunk_strict(
|
||||||
|
data_obj,
|
||||||
|
provider_format,
|
||||||
|
client_format,
|
||||||
|
state=ctx.stream_conversion_state,
|
||||||
|
)
|
||||||
|
|
||||||
|
if converted_events and not ctx.stream_conversion_state.message_started:
|
||||||
|
for evt in converted_events:
|
||||||
|
if evt.get("type") == "message_start" or evt.get("choices"):
|
||||||
|
ctx.stream_conversion_state.message_started = True
|
||||||
|
break
|
||||||
|
|
||||||
|
result = []
|
||||||
|
for evt in converted_events:
|
||||||
|
result.append(f"data: {json.dumps(evt, ensure_ascii=False)}")
|
||||||
|
return result
|
||||||
|
|
||||||
|
except Exception:
|
||||||
|
return [line]
|
||||||
|
|
||||||
|
|
||||||
|
class TestConvertSseLineBasic:
|
||||||
|
"""基本转换测试"""
|
||||||
|
|
||||||
|
def test_empty_line_returns_empty_list(self) -> None:
|
||||||
|
handler = MockCliHandler()
|
||||||
|
ctx = StreamContext(model="test", api_format="OPENAI")
|
||||||
|
|
||||||
|
result = handler._convert_sse_line(ctx, "", [])
|
||||||
|
|
||||||
|
assert result == []
|
||||||
|
|
||||||
|
def test_whitespace_line_returns_line(self) -> None:
|
||||||
|
handler = MockCliHandler()
|
||||||
|
ctx = StreamContext(model="test", api_format="OPENAI")
|
||||||
|
|
||||||
|
result = handler._convert_sse_line(ctx, " ", [])
|
||||||
|
|
||||||
|
assert result == [" "]
|
||||||
|
|
||||||
|
def test_done_marker_returns_as_is(self) -> None:
|
||||||
|
handler = MockCliHandler()
|
||||||
|
ctx = StreamContext(model="test", api_format="OPENAI")
|
||||||
|
|
||||||
|
result = handler._convert_sse_line(ctx, "data: [DONE]", [])
|
||||||
|
|
||||||
|
assert result == ["data: [DONE]"]
|
||||||
|
|
||||||
|
def test_non_data_line_passthrough(self) -> None:
|
||||||
|
handler = MockCliHandler()
|
||||||
|
ctx = StreamContext(model="test", api_format="OPENAI")
|
||||||
|
|
||||||
|
result = handler._convert_sse_line(ctx, "event: message_start", [])
|
||||||
|
|
||||||
|
assert result == ["event: message_start"]
|
||||||
|
|
||||||
|
def test_invalid_json_passthrough(self) -> None:
|
||||||
|
handler = MockCliHandler()
|
||||||
|
ctx = StreamContext(model="test", api_format="OPENAI")
|
||||||
|
|
||||||
|
result = handler._convert_sse_line(ctx, "data: {invalid json}", [])
|
||||||
|
|
||||||
|
assert result == ["data: {invalid json}"]
|
||||||
|
|
||||||
|
|
||||||
|
class TestConvertSseLineWithMockConverter:
|
||||||
|
"""使用 Mock 转换器的测试"""
|
||||||
|
|
||||||
|
def test_same_format_returns_original(self) -> None:
|
||||||
|
"""同格式无需转换"""
|
||||||
|
handler = MockCliHandler()
|
||||||
|
ctx = StreamContext(model="test", api_format="OPENAI")
|
||||||
|
ctx.provider_api_format = "OPENAI"
|
||||||
|
ctx.client_api_format = "OPENAI"
|
||||||
|
|
||||||
|
chunk = {"choices": [{"delta": {"content": "hello"}}]}
|
||||||
|
line = f"data: {json.dumps(chunk)}"
|
||||||
|
|
||||||
|
result = handler._convert_sse_line(ctx, line, [])
|
||||||
|
|
||||||
|
assert len(result) == 1
|
||||||
|
assert json.loads(result[0][6:]) == chunk
|
||||||
|
|
||||||
|
def test_state_initialization(self) -> None:
|
||||||
|
"""测试状态自动初始化"""
|
||||||
|
handler = MockCliHandler()
|
||||||
|
ctx = StreamContext(model="gpt-4", api_format="OPENAI")
|
||||||
|
ctx.provider_api_format = "OPENAI"
|
||||||
|
ctx.client_api_format = "OPENAI"
|
||||||
|
ctx.mapped_model = "claude-3-5-sonnet"
|
||||||
|
ctx.request_id = "req_123"
|
||||||
|
|
||||||
|
chunk = {"choices": [{"delta": {"content": "test"}}]}
|
||||||
|
line = f"data: {json.dumps(chunk)}"
|
||||||
|
|
||||||
|
handler._convert_sse_line(ctx, line, [])
|
||||||
|
|
||||||
|
# 验证状态已初始化
|
||||||
|
assert ctx.stream_conversion_state is not None
|
||||||
|
assert ctx.stream_conversion_state.model == "claude-3-5-sonnet"
|
||||||
|
assert ctx.stream_conversion_state.message_id == "req_123"
|
||||||
|
|
||||||
|
|
||||||
|
class TestConvertSseLineOneInManyOut:
|
||||||
|
"""一入多出测试(需要注册转换器)"""
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def setup_converters(self):
|
||||||
|
"""注册测试用转换器"""
|
||||||
|
from src.core.api_format import (
|
||||||
|
ClaudeToOpenAIConverter,
|
||||||
|
OpenAIToClaudeConverter,
|
||||||
|
converter_registry,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 保存原始状态
|
||||||
|
original_converters = converter_registry._converters.copy()
|
||||||
|
|
||||||
|
# 注册转换器
|
||||||
|
converter_registry.register("OPENAI", "CLAUDE", OpenAIToClaudeConverter())
|
||||||
|
converter_registry.register("CLAUDE", "OPENAI", ClaudeToOpenAIConverter())
|
||||||
|
|
||||||
|
yield
|
||||||
|
|
||||||
|
# 恢复原始状态
|
||||||
|
converter_registry._converters = original_converters
|
||||||
|
|
||||||
|
def test_openai_to_claude_conversion(self) -> None:
|
||||||
|
"""测试 OpenAI -> Claude 流式转换"""
|
||||||
|
handler = MockCliHandler()
|
||||||
|
ctx = StreamContext(model="gpt-4", api_format="OPENAI")
|
||||||
|
ctx.provider_api_format = "OPENAI"
|
||||||
|
ctx.client_api_format = "CLAUDE"
|
||||||
|
ctx.mapped_model = "claude-3-5-sonnet"
|
||||||
|
ctx.request_id = "req_test"
|
||||||
|
|
||||||
|
# 第一个 chunk:带 role
|
||||||
|
chunk1 = {
|
||||||
|
"id": "chatcmpl-123",
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"choices": [{"index": 0, "delta": {"role": "assistant"}, "finish_reason": None}],
|
||||||
|
}
|
||||||
|
line1 = f"data: {json.dumps(chunk1)}"
|
||||||
|
|
||||||
|
result1 = handler._convert_sse_line(ctx, line1, [])
|
||||||
|
|
||||||
|
# 应返回 message_start 事件
|
||||||
|
assert len(result1) >= 1
|
||||||
|
first_event = json.loads(result1[0][6:])
|
||||||
|
assert first_event.get("type") == "message_start"
|
||||||
|
|
||||||
|
# 状态应更新
|
||||||
|
assert ctx.stream_conversion_state.message_started is True
|
||||||
|
|
||||||
|
def test_claude_to_openai_conversion(self) -> None:
|
||||||
|
"""测试 Claude -> OpenAI 流式转换"""
|
||||||
|
handler = MockCliHandler()
|
||||||
|
ctx = StreamContext(model="claude-3-5-sonnet", api_format="CLAUDE")
|
||||||
|
ctx.provider_api_format = "CLAUDE"
|
||||||
|
ctx.client_api_format = "OPENAI"
|
||||||
|
ctx.mapped_model = "gpt-4"
|
||||||
|
ctx.request_id = "msg_test"
|
||||||
|
|
||||||
|
# Claude message_start 事件
|
||||||
|
event = {"type": "message_start", "message": {"id": "msg_123", "role": "assistant"}}
|
||||||
|
line = f"data: {json.dumps(event)}"
|
||||||
|
|
||||||
|
result = handler._convert_sse_line(ctx, line, [])
|
||||||
|
|
||||||
|
# 应返回 OpenAI 格式的 chunk
|
||||||
|
assert len(result) >= 1
|
||||||
|
chunk = json.loads(result[0][6:])
|
||||||
|
assert "choices" in chunk
|
||||||
|
|
||||||
|
def test_multiple_chunks_state_persistence(self) -> None:
|
||||||
|
"""测试多个 chunk 之间状态持久化"""
|
||||||
|
handler = MockCliHandler()
|
||||||
|
ctx = StreamContext(model="gpt-4", api_format="OPENAI")
|
||||||
|
ctx.provider_api_format = "OPENAI"
|
||||||
|
ctx.client_api_format = "CLAUDE"
|
||||||
|
ctx.mapped_model = "claude-3-5-sonnet"
|
||||||
|
|
||||||
|
# 第一个 chunk
|
||||||
|
chunk1 = {"choices": [{"delta": {"role": "assistant"}}]}
|
||||||
|
handler._convert_sse_line(ctx, f"data: {json.dumps(chunk1)}", [])
|
||||||
|
|
||||||
|
state_after_first = ctx.stream_conversion_state
|
||||||
|
message_started_after_first = state_after_first.message_started
|
||||||
|
|
||||||
|
# 第二个 chunk
|
||||||
|
chunk2 = {"choices": [{"delta": {"content": "hello"}}]}
|
||||||
|
handler._convert_sse_line(ctx, f"data: {json.dumps(chunk2)}", [])
|
||||||
|
|
||||||
|
# 状态应该是同一个对象
|
||||||
|
assert ctx.stream_conversion_state is state_after_first
|
||||||
|
# message_started 应该保持 True
|
||||||
|
assert ctx.stream_conversion_state.message_started is True
|
||||||
|
|
||||||
|
|
||||||
|
class TestStreamContextIntegration:
|
||||||
|
"""StreamContext 集成测试"""
|
||||||
|
|
||||||
|
def test_stream_conversion_state_reset_on_retry(self) -> None:
|
||||||
|
"""测试重试时重置流式转换状态"""
|
||||||
|
ctx = StreamContext(model="test", api_format="OPENAI")
|
||||||
|
ctx.stream_conversion_state = StreamConversionState(
|
||||||
|
model="test", message_id="123", message_started=True
|
||||||
|
)
|
||||||
|
|
||||||
|
ctx.reset_for_retry()
|
||||||
|
|
||||||
|
assert ctx.stream_conversion_state is None
|
||||||
|
|
||||||
|
def test_stream_conversion_state_field_exists(self) -> None:
|
||||||
|
"""测试 StreamContext 有 stream_conversion_state 字段"""
|
||||||
|
ctx = StreamContext(model="test", api_format="OPENAI")
|
||||||
|
|
||||||
|
assert hasattr(ctx, "stream_conversion_state")
|
||||||
|
assert ctx.stream_conversion_state is None
|
||||||
0
tests/core/api_format/conversion/__init__.py
Normal file
0
tests/core/api_format/conversion/__init__.py
Normal file
406
tests/core/api_format/conversion/test_registry.py
Normal file
406
tests/core/api_format/conversion/test_registry.py
Normal file
@@ -0,0 +1,406 @@
|
|||||||
|
"""
|
||||||
|
FormatConverterRegistry 单元测试
|
||||||
|
|
||||||
|
测试覆盖:
|
||||||
|
1. 基本注册和查询
|
||||||
|
2. 能力查询方法
|
||||||
|
3. 严格模式转换
|
||||||
|
4. 流式转换签名适配
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any, Dict, List
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.core.api_format.conversion.exceptions import FormatConversionError
|
||||||
|
from src.core.api_format.conversion.registry import FormatConverterRegistry
|
||||||
|
from src.core.api_format.conversion.state import StreamConversionState
|
||||||
|
|
||||||
|
|
||||||
|
# ==================== Mock 转换器 ====================
|
||||||
|
|
||||||
|
|
||||||
|
class MockRequestResponseConverter:
|
||||||
|
"""只支持请求/响应转换的 Mock 转换器"""
|
||||||
|
|
||||||
|
def convert_request(self, request: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
return {"converted": True, "original": request}
|
||||||
|
|
||||||
|
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
return {"converted": True, "original": response}
|
||||||
|
|
||||||
|
|
||||||
|
class MockNewSignatureStreamConverter:
|
||||||
|
"""使用新签名 (chunk, state) 的流式转换器"""
|
||||||
|
|
||||||
|
def convert_request(self, request: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
return request
|
||||||
|
|
||||||
|
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
return response
|
||||||
|
|
||||||
|
def convert_stream_chunk(
|
||||||
|
self, chunk: Dict[str, Any], state: StreamConversionState
|
||||||
|
) -> List[Dict[str, Any]]:
|
||||||
|
"""新签名:返回多个事件"""
|
||||||
|
events = []
|
||||||
|
# 模拟一入多出:第一个 chunk 返回 2 个事件
|
||||||
|
if not state.message_started:
|
||||||
|
events.append({"type": "message_start", "model": state.model})
|
||||||
|
state.message_started = True
|
||||||
|
events.append({"type": "content", "data": chunk})
|
||||||
|
return events
|
||||||
|
|
||||||
|
|
||||||
|
class MockUnifiedStreamConverter:
|
||||||
|
"""使用统一签名 (chunk, state) 的流式转换器"""
|
||||||
|
|
||||||
|
def convert_request(self, request: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
return request
|
||||||
|
|
||||||
|
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
return response
|
||||||
|
|
||||||
|
def convert_stream_chunk(
|
||||||
|
self,
|
||||||
|
chunk: Dict[str, Any],
|
||||||
|
state: StreamConversionState,
|
||||||
|
) -> List[Dict[str, Any]]:
|
||||||
|
"""统一签名:返回事件列表"""
|
||||||
|
events = []
|
||||||
|
if not state.message_started:
|
||||||
|
events.append({"type": "message_start", "model": state.model, "id": state.message_id})
|
||||||
|
state.message_started = True
|
||||||
|
events.append({"type": "delta", "chunk": chunk})
|
||||||
|
return events
|
||||||
|
|
||||||
|
|
||||||
|
class MockReturnsNoneConverter:
|
||||||
|
"""流式转换返回空列表的转换器"""
|
||||||
|
|
||||||
|
def convert_request(self, request: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
return request
|
||||||
|
|
||||||
|
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
return response
|
||||||
|
|
||||||
|
def convert_stream_chunk(
|
||||||
|
self,
|
||||||
|
chunk: Dict[str, Any],
|
||||||
|
state: StreamConversionState,
|
||||||
|
) -> List[Dict[str, Any]]:
|
||||||
|
"""对于不支持的事件类型返回空列表"""
|
||||||
|
event_type = chunk.get("type")
|
||||||
|
if event_type == "message_start":
|
||||||
|
return [{"choices": [{"delta": {"role": "assistant"}}], "model": state.model}]
|
||||||
|
if event_type == "content_block_delta":
|
||||||
|
return [{"choices": [{"delta": {"content": chunk.get("text", "")}}]}]
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
class MockFailingConverter:
|
||||||
|
"""转换时抛出异常的 Mock 转换器"""
|
||||||
|
|
||||||
|
def convert_request(self, request: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
raise ValueError("Request conversion failed")
|
||||||
|
|
||||||
|
def convert_response(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
raise ValueError("Response conversion failed")
|
||||||
|
|
||||||
|
def convert_stream_chunk(
|
||||||
|
self, chunk: Dict[str, Any], state: Any
|
||||||
|
) -> List[Dict[str, Any]]:
|
||||||
|
raise ValueError("Stream conversion failed")
|
||||||
|
|
||||||
|
|
||||||
|
# ==================== 测试类 ====================
|
||||||
|
|
||||||
|
|
||||||
|
class TestFormatConverterRegistryBasic:
|
||||||
|
"""基本注册和查询测试"""
|
||||||
|
|
||||||
|
def test_register_and_get_converter(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
converter = MockRequestResponseConverter()
|
||||||
|
|
||||||
|
registry.register("OPENAI", "CLAUDE", converter)
|
||||||
|
|
||||||
|
assert registry.get_converter("OPENAI", "CLAUDE") is converter
|
||||||
|
assert registry.get_converter("openai", "claude") is converter # 大小写不敏感
|
||||||
|
|
||||||
|
def test_get_nonexistent_converter_returns_none(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
|
||||||
|
assert registry.get_converter("OPENAI", "CLAUDE") is None
|
||||||
|
|
||||||
|
def test_has_converter(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
converter = MockRequestResponseConverter()
|
||||||
|
|
||||||
|
registry.register("OPENAI", "CLAUDE", converter)
|
||||||
|
|
||||||
|
assert registry.has_converter("OPENAI", "CLAUDE") is True
|
||||||
|
assert registry.has_converter("CLAUDE", "OPENAI") is False
|
||||||
|
|
||||||
|
def test_list_converters(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("OPENAI", "CLAUDE", MockRequestResponseConverter())
|
||||||
|
registry.register("CLAUDE", "OPENAI", MockRequestResponseConverter())
|
||||||
|
|
||||||
|
converters = registry.list_converters()
|
||||||
|
|
||||||
|
assert ("OPENAI", "CLAUDE") in converters
|
||||||
|
assert ("CLAUDE", "OPENAI") in converters
|
||||||
|
assert len(converters) == 2
|
||||||
|
|
||||||
|
|
||||||
|
class TestCapabilityQueries:
|
||||||
|
"""能力查询方法测试"""
|
||||||
|
|
||||||
|
def test_can_convert_request(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("A", "B", MockRequestResponseConverter())
|
||||||
|
|
||||||
|
assert registry.can_convert_request("A", "B") is True
|
||||||
|
assert registry.can_convert_request("B", "A") is False
|
||||||
|
|
||||||
|
def test_can_convert_response(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("A", "B", MockRequestResponseConverter())
|
||||||
|
|
||||||
|
assert registry.can_convert_response("A", "B") is True
|
||||||
|
assert registry.can_convert_response("B", "A") is False
|
||||||
|
|
||||||
|
def test_can_convert_stream_with_convert_stream_chunk(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("A", "B", MockNewSignatureStreamConverter())
|
||||||
|
|
||||||
|
assert registry.can_convert_stream("A", "B") is True
|
||||||
|
assert registry.can_convert_stream("B", "A") is False
|
||||||
|
|
||||||
|
def test_can_convert_stream_unified_signature(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("A", "B", MockUnifiedStreamConverter())
|
||||||
|
|
||||||
|
assert registry.can_convert_stream("A", "B") is True
|
||||||
|
|
||||||
|
def test_can_convert_full_request_response_only(self) -> None:
|
||||||
|
"""只有请求/响应转换器,不支持流式"""
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("OPENAI", "CLAUDE", MockRequestResponseConverter())
|
||||||
|
registry.register("CLAUDE", "OPENAI", MockRequestResponseConverter())
|
||||||
|
|
||||||
|
# 不要求流式:通过
|
||||||
|
assert registry.can_convert_full("OPENAI", "CLAUDE", require_stream=False) is True
|
||||||
|
# 要求流式:失败
|
||||||
|
assert registry.can_convert_full("OPENAI", "CLAUDE", require_stream=True) is False
|
||||||
|
|
||||||
|
def test_can_convert_full_with_stream(self) -> None:
|
||||||
|
"""完整双向转换(含流式)"""
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("OPENAI", "CLAUDE", MockNewSignatureStreamConverter())
|
||||||
|
registry.register("CLAUDE", "OPENAI", MockNewSignatureStreamConverter())
|
||||||
|
|
||||||
|
assert registry.can_convert_full("OPENAI", "CLAUDE", require_stream=True) is True
|
||||||
|
|
||||||
|
def test_can_convert_full_missing_reverse(self) -> None:
|
||||||
|
"""缺少反向转换器"""
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("OPENAI", "CLAUDE", MockNewSignatureStreamConverter())
|
||||||
|
# 没有 CLAUDE -> OPENAI
|
||||||
|
|
||||||
|
assert registry.can_convert_full("OPENAI", "CLAUDE", require_stream=False) is False
|
||||||
|
|
||||||
|
def test_get_supported_targets(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("OPENAI", "CLAUDE", MockRequestResponseConverter())
|
||||||
|
registry.register("OPENAI", "GEMINI", MockRequestResponseConverter())
|
||||||
|
registry.register("CLAUDE", "OPENAI", MockRequestResponseConverter())
|
||||||
|
|
||||||
|
targets = registry.get_supported_targets("OPENAI")
|
||||||
|
|
||||||
|
assert "CLAUDE" in targets
|
||||||
|
assert "GEMINI" in targets
|
||||||
|
assert len(targets) == 2
|
||||||
|
|
||||||
|
|
||||||
|
class TestStrictModeConversion:
|
||||||
|
"""严格模式转换测试"""
|
||||||
|
|
||||||
|
def test_convert_request_strict_success(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("A", "B", MockRequestResponseConverter())
|
||||||
|
|
||||||
|
result = registry.convert_request_strict({"foo": "bar"}, "A", "B")
|
||||||
|
|
||||||
|
assert result["converted"] is True
|
||||||
|
assert result["original"] == {"foo": "bar"}
|
||||||
|
|
||||||
|
def test_convert_request_strict_same_format(self) -> None:
|
||||||
|
"""同格式返回原始请求"""
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
|
||||||
|
result = registry.convert_request_strict({"foo": "bar"}, "A", "A")
|
||||||
|
|
||||||
|
assert result == {"foo": "bar"}
|
||||||
|
|
||||||
|
def test_convert_request_strict_no_converter(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
|
||||||
|
with pytest.raises(FormatConversionError) as exc_info:
|
||||||
|
registry.convert_request_strict({"foo": "bar"}, "A", "B")
|
||||||
|
|
||||||
|
assert "未找到转换器" in str(exc_info.value)
|
||||||
|
|
||||||
|
def test_convert_request_strict_converter_fails(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("A", "B", MockFailingConverter())
|
||||||
|
|
||||||
|
with pytest.raises(FormatConversionError) as exc_info:
|
||||||
|
registry.convert_request_strict({"foo": "bar"}, "A", "B")
|
||||||
|
|
||||||
|
assert "Request conversion failed" in str(exc_info.value)
|
||||||
|
|
||||||
|
def test_convert_response_strict_success(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("A", "B", MockRequestResponseConverter())
|
||||||
|
|
||||||
|
result = registry.convert_response_strict({"data": "test"}, "A", "B")
|
||||||
|
|
||||||
|
assert result["converted"] is True
|
||||||
|
|
||||||
|
def test_convert_response_strict_no_converter(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
|
||||||
|
with pytest.raises(FormatConversionError):
|
||||||
|
registry.convert_response_strict({"data": "test"}, "A", "B")
|
||||||
|
|
||||||
|
|
||||||
|
class TestStreamChunkStrictConversion:
|
||||||
|
"""流式转换严格模式测试"""
|
||||||
|
|
||||||
|
def test_stream_chunk_strict_same_format(self) -> None:
|
||||||
|
"""同格式返回原始 chunk 包装在列表中"""
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
chunk = {"type": "delta", "text": "hello"}
|
||||||
|
|
||||||
|
result = registry.convert_stream_chunk_strict(chunk, "A", "A")
|
||||||
|
|
||||||
|
assert result == [chunk]
|
||||||
|
|
||||||
|
def test_stream_chunk_strict_no_converter(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
|
||||||
|
with pytest.raises(FormatConversionError) as exc_info:
|
||||||
|
registry.convert_stream_chunk_strict({"type": "delta"}, "A", "B")
|
||||||
|
|
||||||
|
assert "未找到转换器" in str(exc_info.value)
|
||||||
|
|
||||||
|
def test_stream_chunk_strict_new_signature(self) -> None:
|
||||||
|
"""测试新签名 (chunk, state) 转换器"""
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("A", "B", MockNewSignatureStreamConverter())
|
||||||
|
|
||||||
|
state = StreamConversionState(model="test-model", message_id="msg_123")
|
||||||
|
chunk = {"content": "hello"}
|
||||||
|
|
||||||
|
# 第一次调用:应返回 2 个事件(message_start + content)
|
||||||
|
result = registry.convert_stream_chunk_strict(chunk, "A", "B", state)
|
||||||
|
|
||||||
|
assert len(result) == 2
|
||||||
|
assert result[0]["type"] == "message_start"
|
||||||
|
assert result[1]["type"] == "content"
|
||||||
|
assert state.message_started is True
|
||||||
|
|
||||||
|
# 第二次调用:应返回 1 个事件(只有 content)
|
||||||
|
result2 = registry.convert_stream_chunk_strict({"content": "world"}, "A", "B", state)
|
||||||
|
|
||||||
|
assert len(result2) == 1
|
||||||
|
assert result2[0]["type"] == "content"
|
||||||
|
|
||||||
|
def test_stream_chunk_strict_unified_signature(self) -> None:
|
||||||
|
"""测试统一签名 (chunk, state) 转换器"""
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("A", "B", MockUnifiedStreamConverter())
|
||||||
|
|
||||||
|
state = StreamConversionState(
|
||||||
|
model="gpt-4", message_id="chatcmpl_123", message_started=False
|
||||||
|
)
|
||||||
|
chunk = {"text": "hello"}
|
||||||
|
|
||||||
|
result = registry.convert_stream_chunk_strict(chunk, "A", "B", state)
|
||||||
|
|
||||||
|
assert len(result) == 2
|
||||||
|
assert result[0]["type"] == "message_start"
|
||||||
|
assert result[0]["model"] == "gpt-4"
|
||||||
|
assert result[0]["id"] == "chatcmpl_123"
|
||||||
|
assert state.message_started is True
|
||||||
|
|
||||||
|
def test_stream_chunk_strict_returns_empty_list(self) -> None:
|
||||||
|
"""测试转换器返回空列表"""
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("CLAUDE", "OPENAI", MockReturnsNoneConverter())
|
||||||
|
|
||||||
|
state = StreamConversionState(model="gpt-4", message_id="msg_123")
|
||||||
|
event = {"type": "message_start"}
|
||||||
|
|
||||||
|
result = registry.convert_stream_chunk_strict(event, "CLAUDE", "OPENAI", state)
|
||||||
|
|
||||||
|
assert len(result) == 1
|
||||||
|
assert "choices" in result[0]
|
||||||
|
assert result[0]["model"] == "gpt-4"
|
||||||
|
|
||||||
|
def test_stream_chunk_strict_unknown_event_returns_empty(self) -> None:
|
||||||
|
"""不支持的事件类型应返回空列表"""
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("CLAUDE", "OPENAI", MockReturnsNoneConverter())
|
||||||
|
|
||||||
|
state = StreamConversionState(model="gpt-4", message_id="msg_123")
|
||||||
|
event = {"type": "unknown_event"} # 不支持的事件类型
|
||||||
|
|
||||||
|
result = registry.convert_stream_chunk_strict(event, "CLAUDE", "OPENAI", state)
|
||||||
|
|
||||||
|
assert result == []
|
||||||
|
|
||||||
|
def test_stream_chunk_strict_converter_fails(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("A", "B", MockFailingConverter())
|
||||||
|
|
||||||
|
state = StreamConversionState(model="test", message_id="123")
|
||||||
|
|
||||||
|
with pytest.raises(FormatConversionError) as exc_info:
|
||||||
|
registry.convert_stream_chunk_strict({"data": "test"}, "A", "B", state)
|
||||||
|
|
||||||
|
assert "流式块转换失败" in str(exc_info.value)
|
||||||
|
|
||||||
|
|
||||||
|
class TestNonStrictConversion:
|
||||||
|
"""非严格模式转换测试(失败时返回原始数据)"""
|
||||||
|
|
||||||
|
def test_convert_request_fallback_on_failure(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("A", "B", MockFailingConverter())
|
||||||
|
|
||||||
|
original = {"foo": "bar"}
|
||||||
|
result = registry.convert_request(original, "A", "B")
|
||||||
|
|
||||||
|
# 非严格模式:失败时返回原始请求
|
||||||
|
assert result == original
|
||||||
|
|
||||||
|
def test_convert_response_fallback_on_failure(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("A", "B", MockFailingConverter())
|
||||||
|
|
||||||
|
original = {"data": "test"}
|
||||||
|
result = registry.convert_response(original, "A", "B")
|
||||||
|
|
||||||
|
assert result == original
|
||||||
|
|
||||||
|
def test_convert_stream_chunk_fallback_on_failure(self) -> None:
|
||||||
|
registry = FormatConverterRegistry()
|
||||||
|
registry.register("A", "B", MockFailingConverter())
|
||||||
|
|
||||||
|
original = {"chunk": "data"}
|
||||||
|
result = registry.convert_stream_chunk(original, "A", "B")
|
||||||
|
|
||||||
|
assert result == original
|
||||||
@@ -1,5 +1,5 @@
|
|||||||
from src.core.enums import APIFormat
|
from src.core.api_format import APIFormat
|
||||||
from src.core.headers import (
|
from src.core.api_format import (
|
||||||
CORE_REDACT_HEADERS,
|
CORE_REDACT_HEADERS,
|
||||||
HeaderBuilder,
|
HeaderBuilder,
|
||||||
build_upstream_headers,
|
build_upstream_headers,
|
||||||
|
|||||||
Reference in New Issue
Block a user