fix: 覆写规则条件可使用映射前请求体

新增 rules_original_body 参数贯穿请求构建链路,确保 body_rules/header_rules
条件评估使用模型映射前的原始请求体;附带将 handlers __init__ 改为延迟导入。

Closes #255

Co-authored-by: hemo94931 <[email protected]>
This commit is contained in:
fawney19
2026-03-23 17:31:29 +08:00
co-authored by hemo94931
parent 46737d32f8
commit 4d5c591654
10 changed files with 151 additions and 42 deletions
+7 -6
View File
@@ -14,11 +14,9 @@ API Handlers - 请求处理器
注意:Handler 基类和具体 Handler 使用延迟导入以避免循环依赖。 注意:Handler 基类和具体 Handler 使用延迟导入以避免循环依赖。
""" """
# Adapter 基类(不会引起循环导入,可以直接导入) from __future__ import annotations
from src.api.handlers.base import (
ChatAdapterBase, from typing import Any
CliAdapterBase,
)
__all__ = [ __all__ = [
# Adapter 基类 # Adapter 基类
@@ -49,6 +47,9 @@ __all__ = [
# 延迟导入映射表 # 延迟导入映射表
_LAZY_IMPORTS = { _LAZY_IMPORTS = {
# Adapter 基类
"ChatAdapterBase": ("src.api.handlers.base.chat_adapter_base", "ChatAdapterBase"),
"CliAdapterBase": ("src.api.handlers.base.cli_adapter_base", "CliAdapterBase"),
# Handler 基类 # Handler 基类
"ChatHandlerBase": ("src.api.handlers.base.chat_handler_base", "ChatHandlerBase"), "ChatHandlerBase": ("src.api.handlers.base.chat_handler_base", "ChatHandlerBase"),
"CliMessageHandlerBase": ( "CliMessageHandlerBase": (
@@ -88,7 +89,7 @@ _LAZY_IMPORTS = {
} }
def __getattr__(name: str) -> None: def __getattr__(name: str) -> Any:
"""延迟导入以避免循环依赖""" """延迟导入以避免循环依赖"""
if name in _LAZY_IMPORTS: if name in _LAZY_IMPORTS:
module_path, attr_name = _LAZY_IMPORTS[name] module_path, attr_name = _LAZY_IMPORTS[name]
+59 -32
View File
@@ -8,39 +8,9 @@ Handler 基类模块
会形成循环导入。请直接从具体模块导入 Handler 基类。 会形成循环导入。请直接从具体模块导入 Handler 基类。
""" """
# Chat Adapter 基类(不会引起循环导入) from __future__ import annotations
from src.api.handlers.base.chat_adapter_base import (
ChatAdapterBase,
get_adapter_class,
get_adapter_instance,
list_registered_formats,
register_adapter,
)
# CLI Adapter 基类 from typing import Any
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,
)
__all__ = [ __all__ = [
# Chat Adapter # Chat Adapter
@@ -66,3 +36,60 @@ __all__ = [
"ParsedResponse", "ParsedResponse",
"StreamStats", "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, endpoint,
key, key,
is_stream=upstream_is_stream, is_stream=upstream_is_stream,
rules_original_body=working_request_body,
extra_headers=prep.extra_headers if prep.extra_headers else None, extra_headers=prep.extra_headers if prep.extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None, pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=envelope, envelope=envelope,
+3 -1
View File
@@ -452,13 +452,14 @@ class ChatSyncExecutor:
provider_api_format = str(endpoint.api_format or api_format) 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) client_api_format = api_format.value if hasattr(api_format, "value") else str(api_format)
attempt_body = request_state.build_attempt_body()
# 构建 Provider 请求(模型映射、格式转换、envelope 包装) # 构建 Provider 请求(模型映射、格式转换、envelope 包装)
prep = await handler._prepare_provider_request( prep = await handler._prepare_provider_request(
model=model, model=model,
provider=provider, provider=provider,
endpoint=endpoint, endpoint=endpoint,
key=key, key=key,
working_request_body=request_state.build_attempt_body(), working_request_body=attempt_body,
original_headers=original_headers, original_headers=original_headers,
client_api_format=client_api_format, client_api_format=client_api_format,
provider_api_format=provider_api_format, provider_api_format=provider_api_format,
@@ -487,6 +488,7 @@ class ChatSyncExecutor:
endpoint, endpoint,
key, key,
is_stream=upstream_is_stream, is_stream=upstream_is_stream,
rules_original_body=attempt_body,
extra_headers=prep.extra_headers if prep.extra_headers else None, extra_headers=prep.extra_headers if prep.extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None, pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=envelope, envelope=envelope,
+1
View File
@@ -188,6 +188,7 @@ class CliHandlerProtocol(Protocol):
endpoint: Any, endpoint: Any,
key: Any, key: Any,
request_body: dict[str, Any], request_body: dict[str, Any],
rules_original_body: dict[str, Any] | None = ...,
original_headers: dict[str, str], original_headers: dict[str, str],
query_params: dict[str, str] | None, query_params: dict[str, str] | None,
client_api_format: str, client_api_format: str,
@@ -199,6 +199,7 @@ class CliRequestMixin:
endpoint: Any, endpoint: Any,
key: Any, key: Any,
request_body: dict[str, Any], request_body: dict[str, Any],
rules_original_body: dict[str, Any] | None = None,
original_headers: dict[str, str], original_headers: dict[str, str],
query_params: dict[str, str] | None, query_params: dict[str, str] | None,
client_api_format: str, client_api_format: str,
@@ -319,6 +320,7 @@ class CliRequestMixin:
endpoint, endpoint,
key, key,
is_stream=upstream_is_stream, is_stream=upstream_is_stream,
rules_original_body=rules_original_body,
extra_headers=extra_headers if extra_headers else None, extra_headers=extra_headers if extra_headers else None,
pre_computed_auth=auth_info.as_tuple() if auth_info else None, pre_computed_auth=auth_info.as_tuple() if auth_info else None,
envelope=envelope, envelope=envelope,
@@ -318,6 +318,7 @@ class CliStreamMixin:
endpoint=endpoint, endpoint=endpoint,
key=key, key=key,
request_body=request_body, request_body=request_body,
rules_original_body=working_request_body,
original_headers=original_headers, original_headers=original_headers,
query_params=query_params, query_params=query_params,
client_api_format=client_api_format, client_api_format=client_api_format,
+3 -1
View File
@@ -120,7 +120,8 @@ class CliSyncMixin:
provider_id=str(provider.id), 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: if mapped_model:
mapped_model_result = mapped_model # 保存映射后的模型名,用于 Usage 记录 mapped_model_result = mapped_model # 保存映射后的模型名,用于 Usage 记录
request_body = self.apply_mapped_model(request_body, mapped_model) request_body = self.apply_mapped_model(request_body, mapped_model)
@@ -135,6 +136,7 @@ class CliSyncMixin:
endpoint=endpoint, endpoint=endpoint,
key=key, key=key,
request_body=request_body, request_body=request_body,
rules_original_body=attempt_body,
original_headers=original_headers, original_headers=original_headers,
query_params=query_params, query_params=query_params,
client_api_format=client_api_format, client_api_format=client_api_format,
+14 -2
View File
@@ -1140,6 +1140,7 @@ class RequestBuilder(ABC):
envelope: ProviderEnvelope | None = None, envelope: ProviderEnvelope | None = None,
body: dict[str, Any] | None = None, body: dict[str, Any] | None = None,
original_body: dict[str, Any] | None = None, original_body: dict[str, Any] | None = None,
rules_original_body: dict[str, Any] | None = None,
) -> dict[str, str]: ) -> dict[str, str]:
"""构建请求头""" """构建请求头"""
pass pass
@@ -1151,6 +1152,7 @@ class RequestBuilder(ABC):
endpoint: Any, endpoint: Any,
key: Any, key: Any,
*, *,
rules_original_body: dict[str, Any] | None = None,
mapped_model: str | None = None, mapped_model: str | None = None,
is_stream: bool = False, is_stream: bool = False,
extra_headers: dict[str, str] | None = None, extra_headers: dict[str, str] | None = None,
@@ -1166,6 +1168,7 @@ class RequestBuilder(ABC):
original_headers: 原始请求头 original_headers: 原始请求头
endpoint: 端点配置 endpoint: 端点配置
key: Provider API Key key: Provider API Key
rules_original_body: 规则条件评估用的“原始请求体”(如模型映射前的请求体);不传则使用 original_body
mapped_model: 映射后的模型名 mapped_model: 映射后的模型名
is_stream: 是否为流式请求 is_stream: 是否为流式请求
extra_headers: 额外请求头 extra_headers: 额外请求头
@@ -1175,6 +1178,9 @@ class RequestBuilder(ABC):
Returns: Returns:
Tuple[payload, headers] Tuple[payload, headers]
""" """
effective_rules_original_body = (
rules_original_body if rules_original_body is not None else original_body
)
payload = self.build_payload( payload = self.build_payload(
original_body, original_body,
mapped_model=mapped_model, mapped_model=mapped_model,
@@ -1187,7 +1193,7 @@ class RequestBuilder(ABC):
payload = apply_body_rules( payload = apply_body_rules(
payload, payload,
body_rules, 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) effective_provider_api_format = provider_api_format or getattr(endpoint, "api_format", None)
@@ -1209,6 +1215,7 @@ class RequestBuilder(ABC):
envelope=envelope, envelope=envelope,
body=payload, body=payload,
original_body=original_body, original_body=original_body,
rules_original_body=effective_rules_original_body,
) )
return payload, headers return payload, headers
@@ -1314,6 +1321,7 @@ class PassthroughRequestBuilder(RequestBuilder):
envelope: ProviderEnvelope | None = None, envelope: ProviderEnvelope | None = None,
body: dict[str, Any] | None = None, body: dict[str, Any] | None = None,
original_body: dict[str, Any] | None = None, original_body: dict[str, Any] | None = None,
rules_original_body: dict[str, Any] | None = None,
) -> dict[str, str]: ) -> dict[str, str]:
""" """
透传请求头 - 清理敏感头部(黑名单),透传其他所有头部 透传请求头 - 清理敏感头部(黑名单),透传其他所有头部
@@ -1325,6 +1333,8 @@ class PassthroughRequestBuilder(RequestBuilder):
extra_headers: 额外请求头 extra_headers: 额外请求头
pre_computed_auth: 预先计算的认证信息 (auth_header, auth_value), pre_computed_auth: 预先计算的认证信息 (auth_header, auth_value),
用于 Service Account 等异步获取 token 的场景 用于 Service Account 等异步获取 token 的场景
original_body: 原始请求体(对应 build() 的 original_body)
rules_original_body: 规则条件评估用的“原始请求体”(如模型映射前的请求体);不传则使用 original_body
""" """
raw_family = getattr(endpoint, "api_family", None) raw_family = getattr(endpoint, "api_family", None)
raw_kind = getattr(endpoint, "endpoint_kind", None) raw_kind = getattr(endpoint, "endpoint_kind", None)
@@ -1366,7 +1376,9 @@ class PassthroughRequestBuilder(RequestBuilder):
header_rules, header_rules,
protected_keys, protected_keys,
body=body, 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, condition_evaluator=evaluate_condition,
) )
@@ -0,0 +1,60 @@
from types import SimpleNamespace
from src.api.handlers.base.request_builder import PassthroughRequestBuilder
def test_passthrough_request_builder_rules_original_body_drives_conditions() -> None:
"""
When model mapping mutates the outbound body, body/header rules should still be
able to evaluate conditions against the pre-mapping payload via `source=original`.
"""
builder = PassthroughRequestBuilder()
endpoint = SimpleNamespace(
api_family="openai",
endpoint_kind="chat",
body_rules=[
{
"action": "set",
"path": "route",
"value": "by-original-model",
"condition": {
"source": "original",
"path": "model",
"op": "eq",
"value": "client-model",
},
}
],
header_rules=[
{
"action": "set",
"key": "X-Model-Source",
"value": "client",
"condition": {
"source": "original",
"path": "model",
"op": "eq",
"value": "client-model",
},
}
],
)
key = SimpleNamespace(api_key="unused")
mapped_body = {"model": "provider-model", "messages": []}
original_body = {"model": "client-model", "messages": []}
payload, headers = builder.build(
mapped_body,
{},
endpoint,
key,
rules_original_body=original_body,
pre_computed_auth=("Authorization", "Bearer upstream-token"),
)
assert payload["model"] == "provider-model"
assert payload["route"] == "by-original-model"
assert headers["X-Model-Source"] == "client"
assert headers["Authorization"] == "Bearer upstream-token"