mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 11:19:50 +08:00
fix: 覆写规则条件可使用映射前请求体
新增 rules_original_body 参数贯穿请求构建链路,确保 body_rules/header_rules 条件评估使用模型映射前的原始请求体;附带将 handlers __init__ 改为延迟导入。 Closes #255 Co-authored-by: hemo94931 <[email protected]>
This commit is contained in:
@@ -14,11 +14,9 @@ API Handlers - 请求处理器
|
||||
注意:Handler 基类和具体 Handler 使用延迟导入以避免循环依赖。
|
||||
"""
|
||||
|
||||
# Adapter 基类(不会引起循环导入,可以直接导入)
|
||||
from src.api.handlers.base import (
|
||||
ChatAdapterBase,
|
||||
CliAdapterBase,
|
||||
)
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
__all__ = [
|
||||
# Adapter 基类
|
||||
@@ -49,6 +47,9 @@ __all__ = [
|
||||
|
||||
# 延迟导入映射表
|
||||
_LAZY_IMPORTS = {
|
||||
# Adapter 基类
|
||||
"ChatAdapterBase": ("src.api.handlers.base.chat_adapter_base", "ChatAdapterBase"),
|
||||
"CliAdapterBase": ("src.api.handlers.base.cli_adapter_base", "CliAdapterBase"),
|
||||
# Handler 基类
|
||||
"ChatHandlerBase": ("src.api.handlers.base.chat_handler_base", "ChatHandlerBase"),
|
||||
"CliMessageHandlerBase": (
|
||||
@@ -88,7 +89,7 @@ _LAZY_IMPORTS = {
|
||||
}
|
||||
|
||||
|
||||
def __getattr__(name: str) -> None:
|
||||
def __getattr__(name: str) -> Any:
|
||||
"""延迟导入以避免循环依赖"""
|
||||
if name in _LAZY_IMPORTS:
|
||||
module_path, attr_name = _LAZY_IMPORTS[name]
|
||||
|
||||
@@ -8,39 +8,9 @@ Handler 基类模块
|
||||
会形成循环导入。请直接从具体模块导入 Handler 基类。
|
||||
"""
|
||||
|
||||
# Chat Adapter 基类(不会引起循环导入)
|
||||
from src.api.handlers.base.chat_adapter_base import (
|
||||
ChatAdapterBase,
|
||||
get_adapter_class,
|
||||
get_adapter_instance,
|
||||
list_registered_formats,
|
||||
register_adapter,
|
||||
)
|
||||
from __future__ import annotations
|
||||
|
||||
# CLI Adapter 基类
|
||||
from src.api.handlers.base.cli_adapter_base import (
|
||||
CliAdapterBase,
|
||||
get_cli_adapter_class,
|
||||
get_cli_adapter_instance,
|
||||
list_registered_cli_formats,
|
||||
register_cli_adapter,
|
||||
)
|
||||
|
||||
# 请求构建器
|
||||
from src.api.handlers.base.request_builder import (
|
||||
SENSITIVE_HEADERS,
|
||||
PassthroughRequestBuilder,
|
||||
RequestBuilder,
|
||||
build_passthrough_request,
|
||||
)
|
||||
|
||||
# 响应解析器
|
||||
from src.api.handlers.base.response_parser import (
|
||||
ParsedChunk,
|
||||
ParsedResponse,
|
||||
ResponseParser,
|
||||
StreamStats,
|
||||
)
|
||||
from typing import Any
|
||||
|
||||
__all__ = [
|
||||
# Chat Adapter
|
||||
@@ -66,3 +36,60 @@ __all__ = [
|
||||
"ParsedResponse",
|
||||
"StreamStats",
|
||||
]
|
||||
|
||||
_LAZY_IMPORTS: dict[str, tuple[str, str]] = {
|
||||
# Chat Adapter
|
||||
"ChatAdapterBase": ("src.api.handlers.base.chat_adapter_base", "ChatAdapterBase"),
|
||||
"register_adapter": ("src.api.handlers.base.chat_adapter_base", "register_adapter"),
|
||||
"get_adapter_class": ("src.api.handlers.base.chat_adapter_base", "get_adapter_class"),
|
||||
"get_adapter_instance": ("src.api.handlers.base.chat_adapter_base", "get_adapter_instance"),
|
||||
"list_registered_formats": (
|
||||
"src.api.handlers.base.chat_adapter_base",
|
||||
"list_registered_formats",
|
||||
),
|
||||
# CLI Adapter
|
||||
"CliAdapterBase": ("src.api.handlers.base.cli_adapter_base", "CliAdapterBase"),
|
||||
"register_cli_adapter": (
|
||||
"src.api.handlers.base.cli_adapter_base",
|
||||
"register_cli_adapter",
|
||||
),
|
||||
"get_cli_adapter_class": (
|
||||
"src.api.handlers.base.cli_adapter_base",
|
||||
"get_cli_adapter_class",
|
||||
),
|
||||
"get_cli_adapter_instance": (
|
||||
"src.api.handlers.base.cli_adapter_base",
|
||||
"get_cli_adapter_instance",
|
||||
),
|
||||
"list_registered_cli_formats": (
|
||||
"src.api.handlers.base.cli_adapter_base",
|
||||
"list_registered_cli_formats",
|
||||
),
|
||||
# 请求构建器
|
||||
"RequestBuilder": ("src.api.handlers.base.request_builder", "RequestBuilder"),
|
||||
"PassthroughRequestBuilder": (
|
||||
"src.api.handlers.base.request_builder",
|
||||
"PassthroughRequestBuilder",
|
||||
),
|
||||
"build_passthrough_request": (
|
||||
"src.api.handlers.base.request_builder",
|
||||
"build_passthrough_request",
|
||||
),
|
||||
"SENSITIVE_HEADERS": ("src.api.handlers.base.request_builder", "SENSITIVE_HEADERS"),
|
||||
# 响应解析器
|
||||
"ResponseParser": ("src.api.handlers.base.response_parser", "ResponseParser"),
|
||||
"ParsedChunk": ("src.api.handlers.base.response_parser", "ParsedChunk"),
|
||||
"ParsedResponse": ("src.api.handlers.base.response_parser", "ParsedResponse"),
|
||||
"StreamStats": ("src.api.handlers.base.response_parser", "StreamStats"),
|
||||
}
|
||||
|
||||
|
||||
def __getattr__(name: str) -> Any:
|
||||
"""延迟导入以避免不必要的依赖加载。"""
|
||||
if name in _LAZY_IMPORTS:
|
||||
module_path, attr_name = _LAZY_IMPORTS[name]
|
||||
import importlib
|
||||
|
||||
module = importlib.import_module(module_path)
|
||||
return getattr(module, attr_name)
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
@@ -932,6 +932,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
endpoint,
|
||||
key,
|
||||
is_stream=upstream_is_stream,
|
||||
rules_original_body=working_request_body,
|
||||
extra_headers=prep.extra_headers if prep.extra_headers else None,
|
||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||
envelope=envelope,
|
||||
|
||||
@@ -452,13 +452,14 @@ class ChatSyncExecutor:
|
||||
provider_api_format = str(endpoint.api_format or api_format)
|
||||
client_api_format = api_format.value if hasattr(api_format, "value") else str(api_format)
|
||||
|
||||
attempt_body = request_state.build_attempt_body()
|
||||
# 构建 Provider 请求(模型映射、格式转换、envelope 包装)
|
||||
prep = await handler._prepare_provider_request(
|
||||
model=model,
|
||||
provider=provider,
|
||||
endpoint=endpoint,
|
||||
key=key,
|
||||
working_request_body=request_state.build_attempt_body(),
|
||||
working_request_body=attempt_body,
|
||||
original_headers=original_headers,
|
||||
client_api_format=client_api_format,
|
||||
provider_api_format=provider_api_format,
|
||||
@@ -487,6 +488,7 @@ class ChatSyncExecutor:
|
||||
endpoint,
|
||||
key,
|
||||
is_stream=upstream_is_stream,
|
||||
rules_original_body=attempt_body,
|
||||
extra_headers=prep.extra_headers if prep.extra_headers else None,
|
||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||
envelope=envelope,
|
||||
|
||||
@@ -188,6 +188,7 @@ class CliHandlerProtocol(Protocol):
|
||||
endpoint: Any,
|
||||
key: Any,
|
||||
request_body: dict[str, Any],
|
||||
rules_original_body: dict[str, Any] | None = ...,
|
||||
original_headers: dict[str, str],
|
||||
query_params: dict[str, str] | None,
|
||||
client_api_format: str,
|
||||
|
||||
@@ -199,6 +199,7 @@ class CliRequestMixin:
|
||||
endpoint: Any,
|
||||
key: Any,
|
||||
request_body: dict[str, Any],
|
||||
rules_original_body: dict[str, Any] | None = None,
|
||||
original_headers: dict[str, str],
|
||||
query_params: dict[str, str] | None,
|
||||
client_api_format: str,
|
||||
@@ -319,6 +320,7 @@ class CliRequestMixin:
|
||||
endpoint,
|
||||
key,
|
||||
is_stream=upstream_is_stream,
|
||||
rules_original_body=rules_original_body,
|
||||
extra_headers=extra_headers if extra_headers else None,
|
||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||
envelope=envelope,
|
||||
|
||||
@@ -318,6 +318,7 @@ class CliStreamMixin:
|
||||
endpoint=endpoint,
|
||||
key=key,
|
||||
request_body=request_body,
|
||||
rules_original_body=working_request_body,
|
||||
original_headers=original_headers,
|
||||
query_params=query_params,
|
||||
client_api_format=client_api_format,
|
||||
|
||||
@@ -120,7 +120,8 @@ class CliSyncMixin:
|
||||
provider_id=str(provider.id),
|
||||
)
|
||||
|
||||
request_body = request_state.build_attempt_body()
|
||||
attempt_body = request_state.build_attempt_body()
|
||||
request_body = attempt_body
|
||||
if mapped_model:
|
||||
mapped_model_result = mapped_model # 保存映射后的模型名,用于 Usage 记录
|
||||
request_body = self.apply_mapped_model(request_body, mapped_model)
|
||||
@@ -135,6 +136,7 @@ class CliSyncMixin:
|
||||
endpoint=endpoint,
|
||||
key=key,
|
||||
request_body=request_body,
|
||||
rules_original_body=attempt_body,
|
||||
original_headers=original_headers,
|
||||
query_params=query_params,
|
||||
client_api_format=client_api_format,
|
||||
|
||||
@@ -1140,6 +1140,7 @@ class RequestBuilder(ABC):
|
||||
envelope: ProviderEnvelope | None = None,
|
||||
body: dict[str, Any] | None = None,
|
||||
original_body: dict[str, Any] | None = None,
|
||||
rules_original_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""构建请求头"""
|
||||
pass
|
||||
@@ -1151,6 +1152,7 @@ class RequestBuilder(ABC):
|
||||
endpoint: Any,
|
||||
key: Any,
|
||||
*,
|
||||
rules_original_body: dict[str, Any] | None = None,
|
||||
mapped_model: str | None = None,
|
||||
is_stream: bool = False,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
@@ -1166,6 +1168,7 @@ class RequestBuilder(ABC):
|
||||
original_headers: 原始请求头
|
||||
endpoint: 端点配置
|
||||
key: Provider API Key
|
||||
rules_original_body: 规则条件评估用的“原始请求体”(如模型映射前的请求体);不传则使用 original_body
|
||||
mapped_model: 映射后的模型名
|
||||
is_stream: 是否为流式请求
|
||||
extra_headers: 额外请求头
|
||||
@@ -1175,6 +1178,9 @@ class RequestBuilder(ABC):
|
||||
Returns:
|
||||
Tuple[payload, headers]
|
||||
"""
|
||||
effective_rules_original_body = (
|
||||
rules_original_body if rules_original_body is not None else original_body
|
||||
)
|
||||
payload = self.build_payload(
|
||||
original_body,
|
||||
mapped_model=mapped_model,
|
||||
@@ -1187,7 +1193,7 @@ class RequestBuilder(ABC):
|
||||
payload = apply_body_rules(
|
||||
payload,
|
||||
body_rules,
|
||||
original_body=original_body,
|
||||
original_body=effective_rules_original_body,
|
||||
)
|
||||
|
||||
effective_provider_api_format = provider_api_format or getattr(endpoint, "api_format", None)
|
||||
@@ -1209,6 +1215,7 @@ class RequestBuilder(ABC):
|
||||
envelope=envelope,
|
||||
body=payload,
|
||||
original_body=original_body,
|
||||
rules_original_body=effective_rules_original_body,
|
||||
)
|
||||
return payload, headers
|
||||
|
||||
@@ -1314,6 +1321,7 @@ class PassthroughRequestBuilder(RequestBuilder):
|
||||
envelope: ProviderEnvelope | None = None,
|
||||
body: dict[str, Any] | None = None,
|
||||
original_body: dict[str, Any] | None = None,
|
||||
rules_original_body: dict[str, Any] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
|
||||
@@ -1325,6 +1333,8 @@ class PassthroughRequestBuilder(RequestBuilder):
|
||||
extra_headers: 额外请求头
|
||||
pre_computed_auth: 预先计算的认证信息 (auth_header, auth_value),
|
||||
用于 Service Account 等异步获取 token 的场景
|
||||
original_body: 原始请求体(对应 build() 的 original_body)
|
||||
rules_original_body: 规则条件评估用的“原始请求体”(如模型映射前的请求体);不传则使用 original_body
|
||||
"""
|
||||
raw_family = getattr(endpoint, "api_family", None)
|
||||
raw_kind = getattr(endpoint, "endpoint_kind", None)
|
||||
@@ -1366,7 +1376,9 @@ class PassthroughRequestBuilder(RequestBuilder):
|
||||
header_rules,
|
||||
protected_keys,
|
||||
body=body,
|
||||
original_body=original_body,
|
||||
original_body=rules_original_body
|
||||
if rules_original_body is not None
|
||||
else original_body,
|
||||
condition_evaluator=evaluate_condition,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user