refactor: 统一测试请求构建逻辑,支持多格式测试

- 重构 adapter 基类的 build_request_body 方法,使用 converter_registry 自动处理格式转换
- 后端 test_model 接口增加 endpoint_id 和 api_format 参数支持
- 前端模型映射测试支持根据 Key 和端点配置动态显示可用格式下拉菜单
This commit is contained in:
fawney19
2026-01-22 23:25:51 +08:00
parent 6c7b40764f
commit 1da70d0518
12 changed files with 264 additions and 85 deletions

View File

@@ -45,6 +45,7 @@ class TestModelRequest(BaseModel):
provider_id: str
model_name: str
api_key_id: Optional[str] = None
endpoint_id: Optional[str] = None # 指定使用的端点ID
stream: bool = False
message: Optional[str] = "你好"
api_format: Optional[str] = None # 指定使用的API格式如果不指定则使用端点的默认格式
@@ -200,17 +201,73 @@ async def test_model(
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
# 构建 api_format -> endpoint 映射
# 构建 api_format -> endpoint 映射 和 id -> endpoint 映射
format_to_endpoint: dict[str, ProviderEndpoint] = {}
id_to_endpoint: dict[str, ProviderEndpoint] = {}
for ep in provider.endpoints:
if ep.is_active:
format_to_endpoint[ep.api_format] = ep
id_to_endpoint[ep.id] = ep
# 找到合适的端点和 API Key
endpoint = None
api_key = None
if request.api_key_id:
# 优先级: api_format > endpoint_id > api_key_id > 自动选择
# 如果指定了 api_format优先使用该格式对应的 endpoint
if request.api_format:
endpoint = format_to_endpoint.get(request.api_format)
if not endpoint:
raise HTTPException(
status_code=404,
detail=f"No active endpoint found for API format: {request.api_format}"
)
if request.api_key_id:
# 使用指定的 Key但需要校验是否支持该格式
api_key = next(
(key for key in provider.api_keys if key.id == request.api_key_id and key.is_active),
None
)
if api_key and request.api_format not in (api_key.api_formats or []):
raise HTTPException(
status_code=400,
detail=f"API Key does not support format: {request.api_format}"
)
else:
# 找支持该格式的第一个可用 Key
for key in provider.api_keys:
if not key.is_active:
continue
if request.api_format in (key.api_formats or []):
api_key = key
break
elif request.endpoint_id:
# 使用指定的端点
endpoint = id_to_endpoint.get(request.endpoint_id)
if not endpoint:
raise HTTPException(status_code=404, detail="Endpoint not found or not active")
if request.api_key_id:
# 同时指定了 Key需要校验是否支持该端点格式
api_key = next(
(key for key in provider.api_keys if key.id == request.api_key_id and key.is_active),
None
)
if api_key and endpoint.api_format not in (api_key.api_formats or []):
raise HTTPException(
status_code=400,
detail=f"API Key does not support endpoint format: {endpoint.api_format}"
)
else:
# 找支持该端点格式的第一个可用 Key
for key in provider.api_keys:
if not key.is_active:
continue
if endpoint.api_format in (key.api_formats or []):
api_key = key
break
elif request.api_key_id:
# 使用指定的 API Key
api_key = next(
(key for key in provider.api_keys if key.id == request.api_key_id and key.is_active),
@@ -274,24 +331,6 @@ async def test_model(
logger.debug(f"[test-model] 使用 Adapter: {adapter_class.__name__}")
logger.debug(f"[test-model] 端点 API Format: {endpoint.api_format}")
# 如果请求指定了 api_format优先使用它
target_api_format = request.api_format or endpoint.api_format
if request.api_format and request.api_format != endpoint.api_format:
logger.debug(f"[test-model] 请求指定 API Format: {request.api_format}")
# 重新获取适配器
adapter_class = _get_adapter_for_format(request.api_format)
if not adapter_class:
return {
"success": False,
"error": f"Unknown API format: {request.api_format}",
"provider": {
"id": provider.id,
"name": provider.name,
},
"model": request.model_name,
}
logger.debug(f"[test-model] 重新选择 Adapter: {adapter_class.__name__}")
# 准备测试请求数据
check_request = {
"model": request.model_name,

View File

@@ -107,10 +107,18 @@ class ChatAdapterBase(ApiAdapter):
return build_adapter_headers(cls._get_api_format(), api_key, extra_headers)
@classmethod
def build_request_body(cls, request_data: Dict[str, Any]) -> Dict[str, Any]:
"""构建请求体,子类可以覆盖以自定义请求格式转换"""
# 默认实现:直接使用请求数据
return request_data.copy()
def build_request_body(cls, request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""构建测试请求体,使用转换器注册表自动处理格式转换
Args:
request_data: 可选的请求数据,会与默认测试请求合并
Returns:
转换为目标 API 格式的请求体
"""
from src.api.handlers.base.request_builder import build_test_request_body
return build_test_request_body(cls.FORMAT_ID, request_data)
def extract_api_key(self, request: Request) -> Optional[str]:
"""从请求中提取 API 密钥,使用统一的 headers.py 实现"""

View File

@@ -684,17 +684,18 @@ class CliAdapterBase(ApiAdapter):
raise NotImplementedError(f"{cls.FORMAT_ID} adapter must implement build_endpoint_url")
@classmethod
def build_request_body(cls, request_data: Dict[str, Any]) -> Dict[str, Any]:
"""
构建CLI API请求体 - 子类应覆盖
def build_request_body(cls, request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""构建测试请求体,使用转换器注册表自动处理格式转换
Args:
request_data: 请求数据
request_data: 可选的请求数据,会与默认测试请求合并
Returns:
请求体字典
转换为目标 API 格式的请求体
"""
raise NotImplementedError(f"{cls.FORMAT_ID} adapter must implement build_request_body")
from src.api.handlers.base.request_builder import build_test_request_body
return build_test_request_body(cls.FORMAT_ID, request_data)
@classmethod
def get_cli_user_agent(cls) -> Optional[str]:

View File

@@ -27,6 +27,67 @@ from src.core.api_format import HeaderBuilder, UPSTREAM_DROP_HEADERS
SENSITIVE_HEADERS: FrozenSet[str] = UPSTREAM_DROP_HEADERS
# ==============================================================================
# 测试请求常量与辅助函数
# ==============================================================================
# 标准测试请求体OpenAI 格式)
# 用于 check_endpoint 等测试场景,使用简单安全的消息内容避免触发安全过滤
DEFAULT_TEST_REQUEST: Dict[str, Any] = {
"messages": [{"role": "user", "content": "Hi"}],
"max_tokens": 5,
"temperature": 0,
}
def get_test_request_data(request_data: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""获取测试请求数据
如果传入 request_data则合并到默认测试请求中
否则使用默认测试请求。
Args:
request_data: 用户提供的请求数据(会覆盖默认值)
Returns:
合并后的测试请求数据OpenAI 格式)
"""
if request_data:
merged = DEFAULT_TEST_REQUEST.copy()
merged.update(request_data)
return merged
return DEFAULT_TEST_REQUEST.copy()
def build_test_request_body(
format_id: str,
request_data: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
"""构建测试请求体,自动处理格式转换
使用 converter_registry 将 OpenAI 格式的测试请求转换为目标格式。
Args:
format_id: 目标 API 格式 ID"CLAUDE", "GEMINI", "OPENAI_CLI"
request_data: 可选的请求数据,会与默认测试请求合并
Returns:
转换为目标 API 格式的请求体
"""
from src.core.api_format.conversion import converter_registry
from src.core.api_format.utils import get_base_format
# 获取测试请求数据OpenAI 格式)
source_data = get_test_request_data(request_data)
# CLI 格式使用基础格式进行转换CLAUDE_CLI -> CLAUDE
# 因为 converter_registry 只注册了基础格式之间的转换器
target_format = get_base_format(format_id) or format_id
# 使用注册表进行格式转换 (OPENAI -> 目标基础格式)
return converter_registry.convert_request(source_data, "OPENAI", target_format)
# ==============================================================================
# 请求构建器
# ==============================================================================

View File

@@ -198,14 +198,7 @@ class ClaudeChatAdapter(ChatAdapterBase):
else:
return f"{base_url}/v1/messages"
@classmethod
def build_request_body(cls, request_data: Dict[str, Any]) -> Dict[str, Any]:
"""构建Claude API请求体"""
return {
"model": request_data.get("model"),
"max_tokens": request_data.get("max_tokens", 100),
"messages": request_data.get("messages", []),
}
# build_request_body 使用基类实现,通过 converter_registry 自动转换 OPENAI -> CLAUDE
def build_claude_adapter(x_app_header: Optional[str]):

View File

@@ -128,14 +128,7 @@ class ClaudeCliAdapter(CliAdapterBase):
else:
return f"{base_url}/v1/messages"
@classmethod
def build_request_body(cls, request_data: Dict[str, Any]) -> Dict[str, Any]:
"""构建Claude CLI API请求体"""
return {
"model": request_data.get("model"),
"max_tokens": request_data.get("max_tokens", 100),
"messages": request_data.get("messages", []),
}
# build_request_body 使用基类实现,通过 converter_registry 自动转换 OPENAI -> CLAUDE_CLI
@classmethod
def get_cli_user_agent(cls) -> Optional[str]:

View File

@@ -223,19 +223,7 @@ class GeminiChatAdapter(ChatAdapterBase):
else:
return f"{base_url}/v1beta"
@classmethod
def build_request_body(cls, request_data: Dict[str, Any]) -> Dict[str, Any]:
"""构建Gemini API请求体"""
return {
"contents": request_data.get("messages", []),
"generationConfig": {
"maxOutputTokens": request_data.get("max_tokens", 100),
"temperature": request_data.get("temperature", 0.7),
},
"safetySettings": [
{"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"}
],
}
# build_request_body 使用基类实现,通过 converter_registry 自动转换 OPENAI -> GEMINI
@classmethod
async def check_endpoint(

View File

@@ -149,19 +149,7 @@ class GeminiCliAdapter(CliAdapterBase):
prefix = f"{base_url}/v1beta"
return f"{prefix}/models/{effective_model_name}:generateContent"
@classmethod
def build_request_body(cls, request_data: Dict[str, Any]) -> Dict[str, Any]:
"""构建Gemini CLI API请求体"""
return {
"contents": request_data.get("messages", []),
"generationConfig": {
"maxOutputTokens": request_data.get("max_tokens", 100),
"temperature": request_data.get("temperature", 0.7),
},
"safetySettings": [
{"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"}
],
}
# build_request_body 使用基类实现,通过 converter_registry 自动转换 OPENAI -> GEMINI_CLI
@classmethod
def get_cli_user_agent(cls) -> Optional[str]:

View File

@@ -70,10 +70,8 @@ class OpenAICliAdapter(CliAdapterBase):
else:
return f"{base_url}/v1/chat/completions"
@classmethod
def build_request_body(cls, request_data: Dict[str, Any]) -> Dict[str, Any]:
"""构建OpenAI CLI API请求体"""
return request_data.copy()
# build_request_body 使用基类实现
# OPENAI -> OPENAI_CLI 无转换器,会直接透传原始请求
@classmethod
def get_cli_user_agent(cls) -> Optional[str]: