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