mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
fix(request-body): 使用 deepcopy 防止请求体在处理流程中被意外修改
handler 基类和格式转换 registry 中,原始请求体通过浅拷贝或直接引用传递, 导致下游处理(模型映射、格式转换、重试整流)可能修改原始数据, 影响后续重试或并发请求的正确性。统一改用 copy.deepcopy 隔离副本。
This commit is contained in:
146
tests/api/handlers/base/test_cli_request_body_isolation.py
Normal file
146
tests/api/handlers/base/test_cli_request_body_isolation.py
Normal file
@@ -0,0 +1,146 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
import src.api.handlers.base.cli_stream_mixin as mixmod
|
||||
from src.api.handlers.base.cli_stream_mixin import CliStreamMixin
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
|
||||
|
||||
class _StopBuild(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class _DummyAuthInfo:
|
||||
auth_header = "authorization"
|
||||
auth_value = "Bearer test"
|
||||
decrypted_auth_config = None
|
||||
|
||||
def as_tuple(self) -> tuple[str, str]:
|
||||
return self.auth_header, self.auth_value
|
||||
|
||||
|
||||
class _CaptureBuilder:
|
||||
def __init__(self) -> None:
|
||||
self.request_body: dict[str, Any] | None = None
|
||||
|
||||
def build(self, request_body: dict[str, Any], *args: Any, **kwargs: Any) -> Any:
|
||||
self.request_body = request_body
|
||||
raise _StopBuild()
|
||||
|
||||
|
||||
class _DummyCliStreamHandler(CliStreamMixin):
|
||||
FORMAT_ID = "openai:cli"
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.primary_api_format = "openai:cli"
|
||||
self.request_id = "req-test"
|
||||
self.api_key = SimpleNamespace(id="user-key-1")
|
||||
self._request_builder = _CaptureBuilder()
|
||||
|
||||
async def _get_mapped_model(self, source_model: str, provider_id: str) -> str | None:
|
||||
return None
|
||||
|
||||
def apply_mapped_model(self, request_body: dict[str, Any], mapped_model: str) -> dict[str, Any]:
|
||||
out = dict(request_body)
|
||||
out["model"] = mapped_model
|
||||
return out
|
||||
|
||||
def prepare_provider_request_body(self, request_body: dict[str, Any]) -> dict[str, Any]:
|
||||
request_body.pop("_aether_compact", None)
|
||||
request_body["input"][0]["content"][0]["text"] = "prepared"
|
||||
return request_body
|
||||
|
||||
def finalize_provider_request(
|
||||
self,
|
||||
request_body: dict[str, Any],
|
||||
*,
|
||||
mapped_model: str | None,
|
||||
provider_api_format: str | None,
|
||||
) -> dict[str, Any]:
|
||||
request_body["input"][0]["content"].append({"type": "input_text", "text": "finalized"})
|
||||
return request_body
|
||||
|
||||
def get_model_for_url(
|
||||
self,
|
||||
request_body: dict[str, Any],
|
||||
mapped_model: str | None,
|
||||
) -> str | None:
|
||||
return mapped_model or str(request_body.get("model") or "")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_stream_request_does_not_mutate_original_request_body(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
async def _fake_get_provider_auth(endpoint: Any, key: Any) -> _DummyAuthInfo:
|
||||
return _DummyAuthInfo()
|
||||
|
||||
monkeypatch.setattr(mixmod, "get_provider_auth", _fake_get_provider_auth)
|
||||
monkeypatch.setattr(
|
||||
mixmod,
|
||||
"get_provider_behavior",
|
||||
lambda **kwargs: SimpleNamespace(
|
||||
envelope=None,
|
||||
same_format_variant=None,
|
||||
cross_format_variant=None,
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(mixmod, "get_upstream_stream_policy", lambda *args, **kwargs: None)
|
||||
monkeypatch.setattr(
|
||||
mixmod,
|
||||
"resolve_upstream_is_stream",
|
||||
lambda *, client_is_stream, policy: client_is_stream,
|
||||
)
|
||||
monkeypatch.setattr(mixmod, "enforce_stream_mode_for_upstream", lambda *args, **kwargs: None)
|
||||
monkeypatch.setattr(
|
||||
mixmod,
|
||||
"maybe_patch_request_with_prompt_cache_key",
|
||||
lambda request_body, **kwargs: request_body,
|
||||
)
|
||||
|
||||
handler = _DummyCliStreamHandler()
|
||||
ctx = StreamContext(model="gpt-test", api_format="openai:cli")
|
||||
ctx.client_api_format = "openai:cli"
|
||||
|
||||
provider = SimpleNamespace(name="provider", id="provider-1", provider_type="", proxy=None)
|
||||
endpoint = SimpleNamespace(id="endpoint-1", api_format="openai:cli", base_url="https://x")
|
||||
key = SimpleNamespace(id="key-1", proxy=None)
|
||||
candidate = SimpleNamespace(
|
||||
mapping_matched_model=None, needs_conversion=False, output_limit=None
|
||||
)
|
||||
|
||||
original_request_body = {
|
||||
"model": "gpt-test",
|
||||
"_aether_compact": True,
|
||||
"input": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "input_text", "text": "hello"},
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
snapshot = copy.deepcopy(original_request_body)
|
||||
|
||||
with pytest.raises(_StopBuild):
|
||||
await handler._execute_stream_request(
|
||||
ctx,
|
||||
provider,
|
||||
endpoint,
|
||||
key,
|
||||
original_request_body,
|
||||
{},
|
||||
candidate=candidate,
|
||||
)
|
||||
|
||||
assert original_request_body == snapshot
|
||||
assert handler._request_builder.request_body is not None
|
||||
assert "_aether_compact" not in handler._request_builder.request_body
|
||||
assert handler._request_builder.request_body["input"][0]["content"][0]["text"] == "prepared"
|
||||
assert handler._request_builder.request_body["input"][0]["content"][-1]["text"] == "finalized"
|
||||
132
tests/core/api_format/conversion/test_registry_non_mutation.py
Normal file
132
tests/core/api_format/conversion/test_registry_non_mutation.py
Normal file
@@ -0,0 +1,132 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
from typing import Any
|
||||
|
||||
from src.core.api_format.conversion.internal import (
|
||||
FormatCapabilities,
|
||||
InternalMessage,
|
||||
InternalRequest,
|
||||
InternalResponse,
|
||||
Role,
|
||||
TextBlock,
|
||||
)
|
||||
from src.core.api_format.conversion.normalizer import FormatNormalizer
|
||||
from src.core.api_format.conversion.registry import FormatConversionRegistry
|
||||
|
||||
|
||||
class _BaseTestNormalizer(FormatNormalizer):
|
||||
capabilities = FormatCapabilities()
|
||||
|
||||
def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
|
||||
return InternalResponse(id=str(response.get("id") or ""), model="", content=[])
|
||||
|
||||
def response_from_internal(
|
||||
self,
|
||||
internal: InternalResponse,
|
||||
*,
|
||||
requested_model: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"id": internal.id,
|
||||
"model": requested_model or internal.model,
|
||||
}
|
||||
|
||||
|
||||
class _SameFormatNormalizer(_BaseTestNormalizer):
|
||||
FORMAT_ID = "TEST:SAME"
|
||||
|
||||
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
|
||||
return InternalRequest(model=str(request.get("model") or ""), messages=[])
|
||||
|
||||
def request_from_internal(
|
||||
self,
|
||||
internal: InternalRequest,
|
||||
*,
|
||||
target_variant: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return {"model": internal.model, "variant": target_variant}
|
||||
|
||||
|
||||
class _MutatingSourceNormalizer(_BaseTestNormalizer):
|
||||
FORMAT_ID = "TEST:MUTSRC"
|
||||
|
||||
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
|
||||
request.pop("ephemeral", None)
|
||||
request["messages"][0]["content"][0]["text"] = "mutated"
|
||||
return InternalRequest(
|
||||
model=str(request.get("model") or ""),
|
||||
messages=[
|
||||
InternalMessage(
|
||||
role=Role.USER,
|
||||
content=[TextBlock(text=str(request["messages"][0]["content"][0]["text"]))],
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
def request_from_internal(
|
||||
self,
|
||||
internal: InternalRequest,
|
||||
*,
|
||||
target_variant: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return {"model": internal.model, "variant": target_variant}
|
||||
|
||||
|
||||
class _TargetNormalizer(_BaseTestNormalizer):
|
||||
FORMAT_ID = "TEST:MUTTGT"
|
||||
|
||||
def request_to_internal(self, request: dict[str, Any]) -> InternalRequest:
|
||||
return InternalRequest(model=str(request.get("model") or ""), messages=[])
|
||||
|
||||
def request_from_internal(
|
||||
self,
|
||||
internal: InternalRequest,
|
||||
*,
|
||||
target_variant: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
text = ""
|
||||
if internal.messages and internal.messages[0].content:
|
||||
first = internal.messages[0].content[0]
|
||||
if isinstance(first, TextBlock):
|
||||
text = first.text
|
||||
return {"model": internal.model, "text": text, "variant": target_variant}
|
||||
|
||||
|
||||
def test_convert_request_same_format_returns_detached_copy() -> None:
|
||||
registry = FormatConversionRegistry()
|
||||
registry.register(_SameFormatNormalizer())
|
||||
|
||||
original = {
|
||||
"model": "gpt-test",
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": "hello"}]}],
|
||||
}
|
||||
|
||||
out = registry.convert_request(original, "test:same", "test:same")
|
||||
|
||||
assert out == original
|
||||
assert out is not original
|
||||
assert out["messages"] is not original["messages"]
|
||||
|
||||
out["messages"][0]["content"][0]["text"] = "changed"
|
||||
|
||||
assert original["messages"][0]["content"][0]["text"] == "hello"
|
||||
|
||||
|
||||
def test_convert_request_cross_format_does_not_mutate_original_input() -> None:
|
||||
registry = FormatConversionRegistry()
|
||||
registry.register(_MutatingSourceNormalizer())
|
||||
registry.register(_TargetNormalizer())
|
||||
|
||||
original = {
|
||||
"model": "gpt-test",
|
||||
"ephemeral": "keep-me",
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": "hello"}]}],
|
||||
}
|
||||
snapshot = copy.deepcopy(original)
|
||||
|
||||
out = registry.convert_request(original, "test:mutsrc", "test:muttgt")
|
||||
|
||||
assert out["model"] == "gpt-test"
|
||||
assert out["text"] == "mutated"
|
||||
assert original == snapshot
|
||||
Reference in New Issue
Block a user