mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
refactor(task): 引入 MutableRequestBodyState 替代 request_body_ref 字典容器
将请求体可变状态从 {"body": dict} 字典容器重构为独立的
MutableRequestBodyState 类,统一管理 original_body / current_body /
build_attempt_body / rectify 等语义,消除各层通过 ref["body"] 间接
读写的隐式约定。
- 新增 src/services/task/request_state.py 定义 Protocol 与实现
- handler/executor/mixin 层改用 request_state 参数传递
- error_handler/state_transition 通过 request_state 判断整流状态
- 新增 request_state 单元测试与 chat/cli 请求体隔离测试
This commit is contained in:
200
tests/api/handlers/base/test_chat_request_body_isolation.py
Normal file
200
tests/api/handlers/base/test_chat_request_body_isolation.py
Normal file
@@ -0,0 +1,200 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
import src.api.handlers.base.chat_handler_base as chatmod
|
||||
from src.api.handlers.base.chat_handler_base import ChatHandlerBase
|
||||
from src.api.handlers.base.chat_sync_executor import ChatSyncExecutor
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.services.task.request_state import MutableRequestBodyState
|
||||
|
||||
|
||||
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 _DummyChatHandler(ChatHandlerBase):
|
||||
FORMAT_ID = "openai:chat"
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.request_id = "req-test"
|
||||
self.api_key = SimpleNamespace(id="user-key-1")
|
||||
self._request_builder = _CaptureBuilder()
|
||||
self.allowed_api_formats = ["openai:chat"]
|
||||
self.api_family = None
|
||||
self.endpoint_kind = None
|
||||
|
||||
async def _convert_request(self, request: Any) -> Any:
|
||||
return request
|
||||
|
||||
def _extract_usage(self, response: dict) -> dict[str, int]:
|
||||
return {}
|
||||
|
||||
async def _get_mapped_model(
|
||||
self,
|
||||
source_model: str,
|
||||
provider_id: str,
|
||||
api_format: str | None = None,
|
||||
) -> str | None:
|
||||
del source_model, provider_id, api_format
|
||||
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["messages"][0]["content"] = "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]:
|
||||
del mapped_model, provider_api_format
|
||||
request_body["messages"].append({"role": "assistant", "content": "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 "")
|
||||
|
||||
|
||||
def _patch_chat_upstream(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
async def _fake_get_provider_auth(endpoint: Any, key: Any) -> _DummyAuthInfo:
|
||||
return _DummyAuthInfo()
|
||||
|
||||
monkeypatch.setattr(chatmod, "get_provider_auth", _fake_get_provider_auth)
|
||||
monkeypatch.setattr(
|
||||
chatmod,
|
||||
"get_provider_behavior",
|
||||
lambda **kwargs: SimpleNamespace(
|
||||
envelope=None,
|
||||
same_format_variant=None,
|
||||
cross_format_variant=None,
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(chatmod, "get_upstream_stream_policy", lambda *args, **kwargs: None)
|
||||
monkeypatch.setattr(
|
||||
chatmod,
|
||||
"resolve_upstream_is_stream",
|
||||
lambda *, client_is_stream, policy: client_is_stream,
|
||||
)
|
||||
monkeypatch.setattr(chatmod, "enforce_stream_mode_for_upstream", lambda *args, **kwargs: None)
|
||||
monkeypatch.setattr(
|
||||
chatmod,
|
||||
"maybe_patch_request_with_prompt_cache_key",
|
||||
lambda request_body, **kwargs: request_body,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_execute_stream_request_does_not_mutate_original_request_body(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
_patch_chat_upstream(monkeypatch)
|
||||
|
||||
handler = _DummyChatHandler()
|
||||
ctx = StreamContext(model="gpt-test", api_format="openai:chat")
|
||||
ctx.client_api_format = "openai:chat"
|
||||
|
||||
provider = SimpleNamespace(name="provider", id="provider-1", provider_type="", proxy=None)
|
||||
endpoint = SimpleNamespace(id="endpoint-1", api_format="openai:chat", 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",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
}
|
||||
snapshot = copy.deepcopy(original_request_body)
|
||||
request_state = MutableRequestBodyState(original_request_body)
|
||||
|
||||
with pytest.raises(_StopBuild):
|
||||
await handler._execute_stream_request(
|
||||
ctx,
|
||||
object(),
|
||||
provider,
|
||||
endpoint,
|
||||
key,
|
||||
request_state.build_attempt_body(),
|
||||
{},
|
||||
candidate=candidate,
|
||||
)
|
||||
|
||||
assert original_request_body == snapshot
|
||||
assert handler._request_builder.request_body is not None
|
||||
assert handler._request_builder.request_body["messages"][0]["content"] == "prepared"
|
||||
assert handler._request_builder.request_body["messages"][-1]["content"] == "finalized"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_sync_request_func_does_not_mutate_original_request_body(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
_patch_chat_upstream(monkeypatch)
|
||||
|
||||
handler = _DummyChatHandler()
|
||||
executor = ChatSyncExecutor(handler)
|
||||
|
||||
provider = SimpleNamespace(name="provider", id="provider-1", provider_type="", proxy=None)
|
||||
endpoint = SimpleNamespace(id="endpoint-1", api_format="openai:chat", 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",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
}
|
||||
snapshot = copy.deepcopy(original_request_body)
|
||||
request_state = MutableRequestBodyState(original_request_body)
|
||||
|
||||
with pytest.raises(_StopBuild):
|
||||
await executor._sync_request_func(
|
||||
provider,
|
||||
endpoint,
|
||||
key,
|
||||
candidate,
|
||||
model="gpt-test",
|
||||
api_format="openai:chat",
|
||||
original_headers={},
|
||||
request_state=request_state,
|
||||
)
|
||||
|
||||
assert original_request_body == snapshot
|
||||
assert handler._request_builder.request_body is not None
|
||||
assert handler._request_builder.request_body["messages"][0]["content"] == "prepared"
|
||||
assert handler._request_builder.request_body["messages"][-1]["content"] == "finalized"
|
||||
@@ -10,6 +10,7 @@ import src.api.handlers.base.cli_request_mixin as request_mixmod
|
||||
from src.api.handlers.base.cli_request_mixin import CliRequestMixin
|
||||
from src.api.handlers.base.cli_stream_mixin import CliStreamMixin
|
||||
from src.api.handlers.base.stream_context import StreamContext
|
||||
from src.services.task.request_state import MutableRequestBodyState
|
||||
|
||||
|
||||
class _StopBuild(Exception):
|
||||
@@ -135,6 +136,7 @@ async def test_execute_stream_request_does_not_mutate_original_request_body(
|
||||
],
|
||||
}
|
||||
snapshot = copy.deepcopy(original_request_body)
|
||||
request_state = MutableRequestBodyState(original_request_body)
|
||||
|
||||
with pytest.raises(_StopBuild):
|
||||
await handler._execute_stream_request(
|
||||
@@ -142,7 +144,7 @@ async def test_execute_stream_request_does_not_mutate_original_request_body(
|
||||
provider,
|
||||
endpoint,
|
||||
key,
|
||||
original_request_body,
|
||||
request_state.build_attempt_body(),
|
||||
{},
|
||||
candidate=candidate,
|
||||
)
|
||||
|
||||
32
tests/api/handlers/base/test_request_state.py
Normal file
32
tests/api/handlers/base/test_request_state.py
Normal file
@@ -0,0 +1,32 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from src.services.task.request_state import MutableRequestBodyState
|
||||
|
||||
|
||||
def test_mutable_request_body_state_keeps_original_and_attempts_isolated() -> None:
|
||||
original = {
|
||||
"model": "gpt-5",
|
||||
"input": [{"role": "user", "content": [{"type": "input_text", "text": "hello"}]}],
|
||||
}
|
||||
|
||||
state = MutableRequestBodyState(original)
|
||||
|
||||
first_attempt = state.build_attempt_body()
|
||||
first_attempt["input"][0]["content"][0]["text"] = "attempt-1"
|
||||
|
||||
assert original["input"][0]["content"][0]["text"] == "hello"
|
||||
assert state.current_body["input"][0]["content"][0]["text"] == "hello"
|
||||
|
||||
rectified = state.build_attempt_body()
|
||||
rectified["input"][0]["content"][0]["text"] = "rectified"
|
||||
state.mark_rectified(rectified, stage=1)
|
||||
|
||||
second_attempt = state.build_attempt_body()
|
||||
second_attempt["input"][0]["content"][0]["text"] = "attempt-2"
|
||||
|
||||
assert state.is_rectified() is True
|
||||
assert state.rectify_stage() == 1
|
||||
assert state.current_body["input"][0]["content"][0]["text"] == "rectified"
|
||||
assert original["input"][0]["content"][0]["text"] == "hello"
|
||||
assert state.consume_rectified_this_turn() is True
|
||||
assert state.consume_rectified_this_turn() is False
|
||||
@@ -109,7 +109,7 @@ async def test_execute_sync_unified_temporarily_disables_expire_on_commit(
|
||||
is_stream=True,
|
||||
capability_requirements=None,
|
||||
preferred_key_ids=None,
|
||||
request_body_ref=None,
|
||||
request_body_state=None,
|
||||
request_headers=None,
|
||||
request_body=None,
|
||||
)
|
||||
|
||||
@@ -14,6 +14,7 @@ from src.services.task.execute.state_transition import (
|
||||
SyncExecutionState,
|
||||
resolve_execution_error_transition,
|
||||
)
|
||||
from src.services.task.request_state import MutableRequestBodyState
|
||||
|
||||
|
||||
def _make_candidate() -> SimpleNamespace:
|
||||
@@ -41,10 +42,11 @@ def test_classify_candidate_error_action(
|
||||
|
||||
|
||||
def test_resolve_execution_error_transition_retry_and_consume_rectify_flag() -> None:
|
||||
request_body_ref = {"_rectified_this_turn": True}
|
||||
request_body_state = MutableRequestBodyState({})
|
||||
request_body_state.mark_rectified({}, stage=1)
|
||||
state = SyncExecutionState(
|
||||
candidate_record_map={},
|
||||
request_body_ref=request_body_ref,
|
||||
request_body_state=request_body_state,
|
||||
)
|
||||
|
||||
transition = resolve_execution_error_transition(
|
||||
@@ -56,13 +58,15 @@ def test_resolve_execution_error_transition_retry_and_consume_rectify_flag() ->
|
||||
|
||||
assert transition.failover_action == FailoverAction.RETRY
|
||||
assert transition.max_retries == 3
|
||||
assert request_body_ref["_rectified_this_turn"] is False
|
||||
assert request_body_state.consume_rectified_this_turn() is False
|
||||
|
||||
|
||||
def test_resolve_execution_error_transition_next_candidate() -> None:
|
||||
request_body_state = MutableRequestBodyState({})
|
||||
request_body_state.mark_rectified({}, stage=1)
|
||||
state = SyncExecutionState(
|
||||
candidate_record_map={},
|
||||
request_body_ref={"_rectified_this_turn": True},
|
||||
request_body_state=request_body_state,
|
||||
)
|
||||
|
||||
transition = resolve_execution_error_transition(
|
||||
@@ -74,13 +78,14 @@ def test_resolve_execution_error_transition_next_candidate() -> None:
|
||||
|
||||
assert transition.failover_action == FailoverAction.CONTINUE
|
||||
assert transition.max_retries is None
|
||||
assert state.request_body_ref == {"_rectified_this_turn": True}
|
||||
assert state.request_body_state is request_body_state
|
||||
assert request_body_state.consume_rectified_this_turn() is True
|
||||
|
||||
|
||||
def test_sync_execution_state_resolve_candidate_record_id_fallback() -> None:
|
||||
state = SyncExecutionState(
|
||||
candidate_record_map={(2, 0): "r20"},
|
||||
request_body_ref=None,
|
||||
request_body_state=None,
|
||||
)
|
||||
|
||||
assert state.resolve_candidate_record_id(candidate_index=2, record_id=None) == "r20"
|
||||
@@ -92,7 +97,7 @@ def test_sync_execution_state_raise_classified_error_uses_last_error() -> None:
|
||||
candidate = _make_candidate()
|
||||
state = SyncExecutionState(
|
||||
candidate_record_map={},
|
||||
request_body_ref=None,
|
||||
request_body_state=None,
|
||||
last_error=err,
|
||||
last_candidate=candidate,
|
||||
)
|
||||
@@ -117,7 +122,7 @@ def test_sync_execution_state_raise_classified_error_uses_last_error() -> None:
|
||||
def test_sync_execution_state_raise_classified_error_fallback_error() -> None:
|
||||
state = SyncExecutionState(
|
||||
candidate_record_map={},
|
||||
request_body_ref=None,
|
||||
request_body_state=None,
|
||||
)
|
||||
failure_ops = MagicMock()
|
||||
|
||||
|
||||
@@ -10,8 +10,8 @@ from src.services.candidate.schema import CandidateKey
|
||||
from src.services.candidate.submit import SubmitOutcome
|
||||
from src.services.task.core.context import TaskMode
|
||||
from src.services.task.core.protocol import AttemptKind
|
||||
from src.services.task.service import pool_on_error
|
||||
from src.services.task.service import TaskService
|
||||
from src.services.task.request_state import MutableRequestBodyState
|
||||
from src.services.task.service import TaskService, pool_on_error
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -87,7 +87,7 @@ async def test_task_service_execute_sync_passes_request_headers_and_body() -> No
|
||||
|
||||
request_headers = {"authorization": "Bearer test", "x-trace-id": "abc123"}
|
||||
request_body = {"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]}
|
||||
request_body_ref = {"body": request_body}
|
||||
request_body_state = MutableRequestBodyState(request_body)
|
||||
|
||||
result = await svc.execute(
|
||||
task_type="chat",
|
||||
@@ -100,7 +100,7 @@ async def test_task_service_execute_sync_passes_request_headers_and_body() -> No
|
||||
is_stream=True,
|
||||
request_headers=request_headers,
|
||||
request_body=request_body,
|
||||
request_body_ref=request_body_ref,
|
||||
request_body_state=request_body_state,
|
||||
)
|
||||
|
||||
assert result is sentinel_result
|
||||
@@ -108,7 +108,7 @@ async def test_task_service_execute_sync_passes_request_headers_and_body() -> No
|
||||
kwargs = svc._sync_ops.execute_sync_unified.await_args.kwargs # type: ignore[attr-defined, union-attr]
|
||||
assert kwargs["request_headers"] == request_headers
|
||||
assert kwargs["request_body"] == request_body
|
||||
assert kwargs["request_body_ref"] == request_body_ref
|
||||
assert kwargs["request_body_state"] is request_body_state
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
Reference in New Issue
Block a user