mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
refactor: 统一测试请求构建逻辑,支持多格式测试
- 重构 adapter 基类的 build_request_body 方法,使用 converter_registry 自动处理格式转换 - 后端 test_model 接口增加 endpoint_id 和 api_format 参数支持 - 前端模型映射测试支持根据 Key 和端点配置动态显示可用格式下拉菜单
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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 实现"""
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# 请求构建器
|
||||
# ==============================================================================
|
||||
|
||||
@@ -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]):
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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]:
|
||||
|
||||
Reference in New Issue
Block a user