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:
fawney19
2026-01-21 18:07:10 +08:00
parent 856077a46c
commit 99388bfa33
57 changed files with 2405 additions and 450 deletions

View File

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

View File

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

View File

@@ -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 # 参数保留以符合接口规范

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -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
# 从配置获取路径前缀 # 从配置获取路径前缀

View File

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

View File

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

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

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

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

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

View File

@@ -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 事件

View File

@@ -4,7 +4,14 @@ Gemini 格式转换器
提供 Gemini 与其他 API 格式ClaudeOpenAI之间的转换 提供 Gemini 与其他 API 格式ClaudeOpenAI之间的转换
""" """
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",

View File

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

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

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

View 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=Truetarget -> 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",
]

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

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

View 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"]

View File

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

View File

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

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

View File

@@ -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):
"""用户角色枚举""" """用户角色枚举"""

View File

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

View File

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

View File

@@ -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 = []

View File

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

View File

@@ -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返回能力需求字典

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View 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

View 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

View File

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