mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
fix: 测试模型中使用请求头和请求体规则修改请求
This commit is contained in:
@@ -474,6 +474,16 @@ async def test_model(
|
||||
"temperature": 0.7,
|
||||
}
|
||||
|
||||
# 获取端点规则(不在此处应用,传递给 check_endpoint 在格式转换后应用)
|
||||
body_rules = getattr(endpoint, "body_rules", None)
|
||||
header_rules = getattr(endpoint, "header_rules", None)
|
||||
extra_headers = endpoint_config.get("extra_headers") or {}
|
||||
|
||||
if body_rules:
|
||||
logger.debug(f"[test-model] 将传递 body_rules 给 check_endpoint: {body_rules}")
|
||||
if header_rules:
|
||||
logger.debug(f"[test-model] 将传递 header_rules 给 check_endpoint: {header_rules}")
|
||||
|
||||
# 发送测试请求
|
||||
async with httpx.AsyncClient(
|
||||
timeout=endpoint_config["timeout"], verify=get_ssl_context()
|
||||
@@ -486,7 +496,10 @@ async def test_model(
|
||||
endpoint_config["base_url"],
|
||||
endpoint_config["api_key"],
|
||||
check_request,
|
||||
endpoint_config.get("extra_headers"),
|
||||
extra_headers if extra_headers else None,
|
||||
# 端点规则(在 check_endpoint 内部格式转换后应用)
|
||||
body_rules=body_rules,
|
||||
header_rules=header_rules,
|
||||
# 用量计算参数(现在强制记录)
|
||||
db=db,
|
||||
user=current_user,
|
||||
|
||||
@@ -630,6 +630,9 @@ class ChatAdapterBase(ApiAdapter):
|
||||
api_key: str,
|
||||
request_data: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
# 端点规则参数
|
||||
body_rules: list[dict[str, Any]] | None = None,
|
||||
header_rules: list[dict[str, Any]] | None = None,
|
||||
# 用量计算参数(现在强制记录)
|
||||
db: Any | None = None,
|
||||
user: Any | None = None,
|
||||
@@ -647,6 +650,8 @@ class ChatAdapterBase(ApiAdapter):
|
||||
api_key: API 密钥(已解密)
|
||||
request_data: 请求数据
|
||||
extra_headers: 端点配置的额外请求头
|
||||
body_rules: 请求体规则(在格式转换后应用)
|
||||
header_rules: 请求头规则(在请求头构建后应用)
|
||||
db: 数据库会话
|
||||
user: 用户对象
|
||||
provider_name: 提供商名称
|
||||
@@ -658,12 +663,30 @@ class ChatAdapterBase(ApiAdapter):
|
||||
测试响应数据
|
||||
"""
|
||||
from src.api.handlers.base.endpoint_checker import run_endpoint_check
|
||||
from src.api.handlers.base.request_builder import apply_body_rules
|
||||
from src.core.api_format.headers import HeaderBuilder
|
||||
|
||||
# 使用子类配置方法构建请求组件
|
||||
url = cls.build_endpoint_url(base_url)
|
||||
headers = cls.build_headers_with_extra(api_key, extra_headers)
|
||||
body = cls.build_request_body(request_data)
|
||||
|
||||
# 应用请求体规则(在格式转换后应用,确保规则效果不被覆盖)
|
||||
if body_rules:
|
||||
body = apply_body_rules(body, body_rules)
|
||||
|
||||
# 应用请求头规则(在请求头构建后应用)
|
||||
if header_rules:
|
||||
# 获取认证头名称,防止被规则覆盖
|
||||
from src.core.api_format import get_auth_config
|
||||
auth_header, _ = get_auth_config(cls._get_api_format())
|
||||
protected_keys = {auth_header.lower(), "content-type"}
|
||||
|
||||
header_builder = HeaderBuilder()
|
||||
header_builder.add_many(headers)
|
||||
header_builder.apply_rules(header_rules, protected_keys)
|
||||
headers = header_builder.build()
|
||||
|
||||
# 使用通用的endpoint checker执行请求
|
||||
return await run_endpoint_check(
|
||||
client=client,
|
||||
|
||||
@@ -598,6 +598,9 @@ class CliAdapterBase(ApiAdapter):
|
||||
api_key: str,
|
||||
request_data: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
# 端点规则参数
|
||||
body_rules: list[dict[str, Any]] | None = None,
|
||||
header_rules: list[dict[str, Any]] | None = None,
|
||||
# 用量计算参数
|
||||
db: Any | None = None,
|
||||
user: Any | None = None,
|
||||
@@ -622,6 +625,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
api_key: API 密钥(已解密)
|
||||
request_data: 请求数据
|
||||
extra_headers: 端点配置的额外请求头
|
||||
body_rules: 请求体规则(在格式转换后应用)
|
||||
header_rules: 请求头规则(在请求头构建后应用)
|
||||
db: 数据库会话
|
||||
user: 用户对象
|
||||
provider_name: 提供商名称
|
||||
@@ -633,6 +638,8 @@ class CliAdapterBase(ApiAdapter):
|
||||
测试响应数据
|
||||
"""
|
||||
from src.api.handlers.base.endpoint_checker import run_endpoint_check
|
||||
from src.api.handlers.base.request_builder import apply_body_rules
|
||||
from src.core.api_format.headers import HeaderBuilder
|
||||
|
||||
# 构建请求组件
|
||||
url = cls.build_endpoint_url(base_url, request_data, model_name)
|
||||
@@ -646,6 +653,22 @@ class CliAdapterBase(ApiAdapter):
|
||||
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None)
|
||||
body = cls.build_request_body(request_data)
|
||||
|
||||
# 应用请求体规则(在格式转换后应用,确保规则效果不被覆盖)
|
||||
if body_rules:
|
||||
body = apply_body_rules(body, body_rules)
|
||||
|
||||
# 应用请求头规则(在请求头构建后应用)
|
||||
if header_rules:
|
||||
# 获取认证头名称,防止被规则覆盖
|
||||
from src.core.api_format import get_auth_config
|
||||
auth_header, _ = get_auth_config(cls._get_api_format())
|
||||
protected_keys = {auth_header.lower(), "content-type"}
|
||||
|
||||
header_builder = HeaderBuilder()
|
||||
header_builder.add_many(headers)
|
||||
header_builder.apply_rules(header_rules, protected_keys)
|
||||
headers = header_builder.build()
|
||||
|
||||
# 获取有效的模型名称
|
||||
effective_model_name = model_name or request_data.get("model")
|
||||
|
||||
|
||||
@@ -252,6 +252,9 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
api_key: str,
|
||||
request_data: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
# 端点规则参数
|
||||
body_rules: list[dict[str, Any]] | None = None,
|
||||
header_rules: list[dict[str, Any]] | None = None,
|
||||
# 用量计算参数
|
||||
db: Any | None = None,
|
||||
user: Any | None = None,
|
||||
@@ -261,6 +264,10 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
model_name: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""测试 Gemini API 模型连接性(非流式)"""
|
||||
from src.api.handlers.base.endpoint_checker import run_endpoint_check
|
||||
from src.api.handlers.base.request_builder import apply_body_rules
|
||||
from src.core.api_format.headers import HeaderBuilder
|
||||
|
||||
# Gemini需要从request_data或model_name参数获取model名称
|
||||
effective_model_name = model_name or request_data.get("model", "")
|
||||
if not effective_model_name:
|
||||
@@ -277,8 +284,21 @@ class GeminiChatAdapter(ChatAdapterBase):
|
||||
headers = cls.build_headers_with_extra(api_key, extra_headers)
|
||||
body = cls.build_request_body(request_data)
|
||||
|
||||
# 使用基类的通用endpoint checker
|
||||
from src.api.handlers.base.endpoint_checker import run_endpoint_check
|
||||
# 应用请求体规则(在格式转换后应用,确保规则效果不被覆盖)
|
||||
if body_rules:
|
||||
body = apply_body_rules(body, body_rules)
|
||||
|
||||
# 应用请求头规则(在请求头构建后应用)
|
||||
if header_rules:
|
||||
# 获取认证头名称,防止被规则覆盖
|
||||
from src.core.api_format import get_auth_config
|
||||
auth_header, _ = get_auth_config(cls._get_api_format())
|
||||
protected_keys = {auth_header.lower(), "content-type"}
|
||||
|
||||
header_builder = HeaderBuilder()
|
||||
header_builder.add_many(headers)
|
||||
header_builder.apply_rules(header_rules, protected_keys)
|
||||
headers = header_builder.build()
|
||||
|
||||
return await run_endpoint_check(
|
||||
client=client,
|
||||
|
||||
Reference in New Issue
Block a user