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:
fawney19
2026-03-18 13:43:39 +08:00
parent 53ef35ec80
commit cbb66a5667
16 changed files with 374 additions and 85 deletions

View File

@@ -2179,6 +2179,7 @@ async def test_model_failover(
from src.services.scheduling.scheduling_config import SchedulingConfig
from src.services.task import TaskService
from src.services.task.core.protocol import AttemptKind, AttemptResult
from src.services.task.request_state import MutableRequestBodyState
provider = (
db.query(Provider)
@@ -2363,7 +2364,7 @@ async def test_model_failover(
user_api_key=None,
is_stream=False,
capability_requirements=None,
request_body_ref={"body": dict(request_payload)},
request_body_state=MutableRequestBodyState(dict(request_payload)),
request_headers=None,
request_body=dict(request_payload),
affinity_key=f"provider-test:{provider.id}",

View File

@@ -22,7 +22,6 @@ Chat Handler Base - Chat API 格式的通用基类
from __future__ import annotations
import asyncio
import copy
import json
from abc import ABC, abstractmethod
from collections.abc import AsyncGenerator, Awaitable, Callable
@@ -94,6 +93,7 @@ from src.services.provider.transport import (
)
from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.system.config import SystemConfigService
from src.services.task.request_state import MutableRequestBodyState
@dataclass
@@ -496,9 +496,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
)
api_format = self.allowed_api_formats[0]
# 可变请求体容器:允许 TaskService 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: dict[str, Any] = {"body": copy.deepcopy(original_request_body)}
request_state = MutableRequestBodyState(original_request_body)
# 创建类型安全的流式上下文
ctx = StreamContext(
@@ -542,7 +540,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider,
endpoint,
key,
request_body_ref["body"], # 使用容器中的请求体
request_state.build_attempt_body(),
original_headers,
query_params,
candidate,
@@ -577,7 +575,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
is_stream=True,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body_ref=request_body_ref,
request_body_state=request_state,
request_headers=original_headers,
request_body=original_request_body,
)
@@ -612,7 +610,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
if isinstance(scheduling_audit, dict):
ctx.scheduling_audit = scheduling_audit
# 同步整流状态(如果请求体被整流过)
ctx.rectified = request_body_ref.get("_rectified", False)
ctx.rectified = request_state.is_rectified()
# 创建遥测记录器
telemetry_recorder = StreamTelemetryRecorder(
@@ -692,7 +690,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
original_request_body: dict[str, Any],
working_request_body: dict[str, Any],
original_headers: dict[str, str],
client_api_format: str,
provider_api_format: str,
@@ -721,8 +719,8 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
api_format=provider_api_format,
)
# 应用模型映射到请求体
request_body = copy.deepcopy(original_request_body)
# `working_request_body` is already isolated per attempt.
request_body = working_request_body
if mapped_model:
request_body = self.apply_mapped_model(request_body, mapped_model)
@@ -859,7 +857,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider: Provider,
endpoint: ProviderEndpoint,
key: ProviderAPIKey,
original_request_body: dict[str, Any],
working_request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
candidate: ProviderCandidate | None = None,
@@ -893,7 +891,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
provider=provider,
endpoint=endpoint,
key=key,
original_request_body=original_request_body,
working_request_body=working_request_body,
original_headers=original_headers,
client_api_format=client_api_format,
provider_api_format=provider_api_format,

View File

@@ -11,7 +11,6 @@ ChatSyncExecutor - 非流式请求执行器
from __future__ import annotations
import copy
import json
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
@@ -47,6 +46,7 @@ from src.core.exceptions import (
UpstreamClientException,
)
from src.core.logger import logger
from src.services.task.request_state import MutableRequestBodyState
if TYPE_CHECKING:
from fastapi import Request
@@ -122,9 +122,7 @@ class ChatSyncExecutor:
request_body=original_request_body,
)
# 可变请求体容器:允许 TaskService 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: dict[str, Any] = {"body": copy.deepcopy(original_request_body)}
request_state = MutableRequestBodyState(original_request_body)
# 捕获的上下文变量
ctx = self._ctx
@@ -143,7 +141,7 @@ class ChatSyncExecutor:
model=model,
api_format=api_format,
original_headers=original_headers,
request_body_ref=request_body_ref,
request_state=request_state,
query_params=query_params,
client_content_encoding=effective_client_content_encoding,
)
@@ -175,7 +173,7 @@ class ChatSyncExecutor:
is_stream=False,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body_ref=request_body_ref,
request_body_state=request_state,
request_headers=original_headers,
request_body=original_request_body,
)
@@ -437,7 +435,7 @@ class ChatSyncExecutor:
model: str,
api_format: Any,
original_headers: dict[str, Any],
request_body_ref: dict[str, Any],
request_state: MutableRequestBodyState,
query_params: dict[str, str] | None = None,
client_content_encoding: str | None = None,
) -> dict[str, Any]:
@@ -458,7 +456,7 @@ class ChatSyncExecutor:
provider=provider,
endpoint=endpoint,
key=key,
original_request_body=request_body_ref["body"],
working_request_body=request_state.build_attempt_body(),
original_headers=original_headers,
client_api_format=client_api_format,
provider_api_format=provider_api_format,

View File

@@ -4,7 +4,6 @@ from __future__ import annotations
import asyncio
import codecs
import copy
import json
import time
from collections.abc import AsyncGenerator
@@ -41,6 +40,7 @@ from src.core.logger import logger
from src.services.provider.behavior import get_provider_behavior
from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.system.config import SystemConfigService
from src.services.task.request_state import MutableRequestBodyState
from src.utils.sse_parser import SSEEventParser
from .cli_sse_helpers import _format_converted_events_to_sse
@@ -84,9 +84,7 @@ class CliStreamMixin:
client_content_encoding,
)
# 可变请求体容器:允许 TaskService 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: dict[str, Any] = {"body": copy.deepcopy(original_request_body)}
request_state = MutableRequestBodyState(original_request_body)
# 使用子类实现的方法提取 model不同 API 格式的 model 位置不同)
# 注意:使用 original_request_body因为整流只修改 messages不影响 model 字段
@@ -132,7 +130,7 @@ class CliStreamMixin:
provider,
endpoint,
key,
request_body_ref["body"], # 使用容器中的请求体
request_state.build_attempt_body(),
original_headers,
query_params,
candidate,
@@ -167,7 +165,7 @@ class CliStreamMixin:
is_stream=True,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body_ref=request_body_ref,
request_body_state=request_state,
request_headers=original_headers,
request_body=original_request_body,
)
@@ -206,7 +204,7 @@ class CliStreamMixin:
if isinstance(scheduling_audit, dict):
ctx.scheduling_audit = scheduling_audit
# 同步整流状态(如果请求体被整流过)
ctx.rectified = request_body_ref.get("_rectified", False)
ctx.rectified = request_state.is_rectified()
# 创建后台任务记录统计
background_tasks = BackgroundTasks()
@@ -253,7 +251,7 @@ class CliStreamMixin:
provider: "Provider",
endpoint: "ProviderEndpoint",
key: "ProviderAPIKey",
original_request_body: dict[str, Any],
working_request_body: dict[str, Any],
original_headers: dict[str, str],
query_params: dict[str, str] | None = None,
candidate: ProviderCandidate | None = None,
@@ -297,8 +295,8 @@ class CliStreamMixin:
provider_id=str(provider.id),
)
# 应用模型映射到请求体(子类可覆盖此方法处理不同格式)
request_body = copy.deepcopy(original_request_body)
# `working_request_body` is already isolated per attempt.
request_body = working_request_body
if mapped_model:
ctx.mapped_model = mapped_model # 保存映射后的模型名,用于 Usage 记录
request_body = self.apply_mapped_model(request_body, mapped_model)

View File

@@ -2,7 +2,6 @@
from __future__ import annotations
import copy
import json
import time
from typing import TYPE_CHECKING, Any
@@ -33,6 +32,7 @@ from src.core.exceptions import (
)
from src.core.logger import logger
from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.task.request_state import MutableRequestBodyState
if TYPE_CHECKING:
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
@@ -100,9 +100,7 @@ class CliSyncMixin:
needs_conversion = False # 是否需要格式转换(由 candidate 决定)
sync_proxy_info: dict[str, Any] | None = None # 代理信息
# 可变请求体容器:允许 TaskService 在遇到 Thinking 签名错误时整流请求体后重试
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: dict[str, Any] = {"body": copy.deepcopy(original_request_body)}
request_state = MutableRequestBodyState(original_request_body)
async def sync_request_func(
provider: "Provider",
@@ -122,8 +120,7 @@ class CliSyncMixin:
provider_id=str(provider.id),
)
# 应用模型映射到请求体(子类可覆盖此方法处理不同格式)
request_body = copy.deepcopy(request_body_ref["body"])
request_body = request_state.build_attempt_body()
if mapped_model:
mapped_model_result = mapped_model # 保存映射后的模型名,用于 Usage 记录
request_body = self.apply_mapped_model(request_body, mapped_model)
@@ -390,7 +387,7 @@ class CliSyncMixin:
is_stream=False,
capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids or None,
request_body_ref=request_body_ref,
request_body_state=request_state,
request_headers=original_headers,
request_body=original_request_body,
)

View File

@@ -18,6 +18,7 @@ from src.core.logger import logger
from src.core.provider_types import ProviderType
from src.services.request.candidate import RequestCandidateService
from src.services.task.execute.pool import TaskPoolOperationsService
from src.services.task.request_state import RequestBodyState
class TaskErrorOperationsService:
@@ -60,7 +61,7 @@ class TaskErrorOperationsService:
elapsed_ms: int,
captured_key_concurrent: int | None,
serializable_extra_data: dict[str, Any],
request_body_ref: dict[str, Any] | None,
request_body_state: RequestBodyState | None,
) -> str:
"""Try to rectify thinking signature errors and request a retry."""
from src.services.message.thinking_rectifier import ThinkingRectifier
@@ -79,7 +80,7 @@ class TaskErrorOperationsService:
)
raise converted_error
if request_body_ref is None:
if request_body_state is None:
logger.warning(" [{}] Thinking 错误:无法获取请求体引用,终止重试", request_id)
self.mark_thinking_error_failed(
candidate_record_id,
@@ -93,12 +94,8 @@ class TaskErrorOperationsService:
provider_type_norm = str(provider_type or "").lower()
# Rectification may have multiple stages (Antigravity only).
stage_raw = request_body_ref.get("_rectify_stage", 0)
try:
stage = int(stage_raw or 0)
except Exception:
stage = 0
if stage <= 0 and request_body_ref.get("_rectified", False):
stage = request_body_state.rectify_stage()
if stage <= 0 and request_body_state.is_rectified():
stage = 1
if stage >= 2 or (stage >= 1 and provider_type_norm != ProviderType.ANTIGRAVITY):
@@ -112,7 +109,7 @@ class TaskErrorOperationsService:
)
raise converted_error
request_body = request_body_ref.get("body", {})
request_body = request_body_state.current_body
stage_label = "thinking_only"
next_stage = 1
@@ -129,10 +126,7 @@ class TaskErrorOperationsService:
next_stage = 2
if modified:
request_body_ref["body"] = rectified_body
request_body_ref["_rectified"] = True
request_body_ref["_rectified_this_turn"] = True
request_body_ref["_rectify_stage"] = next_stage
request_body_state.mark_rectified(rectified_body, stage=next_stage)
if provider_type_norm == ProviderType.ANTIGRAVITY:
try:
@@ -188,7 +182,7 @@ class TaskErrorOperationsService:
attempt: int,
max_attempts: int,
error_classifier: Any,
request_body_ref: dict[str, Any] | None = None,
request_body_state: RequestBodyState | None = None,
) -> str:
"""
Handle an execution error for a candidate.
@@ -426,7 +420,7 @@ class TaskErrorOperationsService:
elapsed_ms=elapsed_ms,
captured_key_concurrent=captured_key_concurrent,
serializable_extra_data=serializable_extra_data,
request_body_ref=request_body_ref,
request_body_state=request_body_state,
)
if action == "continue":
return "continue"

View File

@@ -1,10 +1,11 @@
from __future__ import annotations
from dataclasses import dataclass, field
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from src.services.candidate.policy import FailoverAction
from src.services.task.execute.exception_classification import CandidateErrorAction
from src.services.task.request_state import RequestBodyState
if TYPE_CHECKING:
from src.services.task.execute.failure import TaskFailureOperationsService
@@ -26,10 +27,9 @@ class SyncExecutionState:
"""同步执行阶段状态容器(候选上下文 + 异常上下文)。"""
candidate_record_map: dict[tuple[int, int], str]
request_body_ref: dict[str, Any] | None
request_body_state: RequestBodyState | None
last_error: Exception | None = None
last_candidate: Any | None = None
_rectify_flag_key: str = field(default="_rectified_this_turn", repr=False)
def touch_candidate(self, candidate: Any) -> None:
self.last_candidate = candidate
@@ -47,12 +47,10 @@ class SyncExecutionState:
def consume_rectify_retry_extension(
self, *, max_retries_for_candidate: int, retry_index: int
) -> int | None:
if not self.request_body_ref:
if not self.request_body_state:
return None
if not self.request_body_ref.get(self._rectify_flag_key, False):
if not self.request_body_state.consume_rectified_this_turn():
return None
self.request_body_ref[self._rectify_flag_key] = False
return max(max_retries_for_candidate, retry_index + 2)
def raise_classified_error(

View File

@@ -31,6 +31,7 @@ from src.services.task.execute.state_transition import (
SyncExecutionState,
resolve_execution_error_transition,
)
from src.services.task.request_state import RequestBodyState
from src.services.usage.service import UsageService
@@ -65,7 +66,7 @@ class SyncTaskExecutionService:
is_stream: bool,
capability_requirements: dict[str, bool] | None,
preferred_key_ids: list[str] | None,
request_body_ref: dict[str, Any] | None,
request_body_state: RequestBodyState | None,
request_headers: dict[str, Any] | None,
request_body: dict[str, Any] | None,
) -> ExecutionResult:
@@ -197,7 +198,7 @@ class SyncTaskExecutionService:
# Keep behavior consistent with previous behavior: last_candidate is updated even if skipped.
execution_state = SyncExecutionState(
candidate_record_map=candidate_record_map,
request_body_ref=request_body_ref,
request_body_state=request_body_state,
last_candidate=all_candidates[-1] if all_candidates else None,
)
@@ -334,7 +335,7 @@ class SyncTaskExecutionService:
request_id=request_id,
attempt=attempt_count,
max_attempts=int(max_attempts or 0),
request_body_ref=request_body_ref,
request_body_state=request_body_state,
error_classifier=error_classifier,
)
action = classify_candidate_error_action(raw_action)

View File

@@ -0,0 +1,64 @@
from __future__ import annotations
import copy
from dataclasses import dataclass, field
from typing import Any, Protocol
class RequestBodyState(Protocol):
@property
def current_body(self) -> dict[str, Any]: ...
def build_attempt_body(self) -> dict[str, Any]: ...
def is_rectified(self) -> bool: ...
def rectify_stage(self) -> int: ...
def mark_rectified(self, body: dict[str, Any], *, stage: int) -> None: ...
def consume_rectified_this_turn(self) -> bool: ...
@dataclass(slots=True)
class MutableRequestBodyState:
"""Owns the mutable working request body used across retries."""
original_body: dict[str, Any]
_current_body: dict[str, Any] = field(init=False, repr=False)
_rectified: bool = field(default=False, init=False, repr=False)
_rectified_this_turn: bool = field(default=False, init=False, repr=False)
_rectify_stage: int = field(default=0, init=False, repr=False)
def __post_init__(self) -> None:
# Attempts already deep-copy per dispatch, and rectification paths clone before
# rewriting. Keep the initial working body as a direct view to avoid an eager copy.
self._current_body = self.original_body
@property
def current_body(self) -> dict[str, Any]:
return self._current_body
def build_attempt_body(self) -> dict[str, Any]:
return copy.deepcopy(self._current_body)
def is_rectified(self) -> bool:
return self._rectified
def rectify_stage(self) -> int:
return self._rectify_stage
def mark_rectified(self, body: dict[str, Any], *, stage: int) -> None:
self._current_body = body
self._rectified = True
self._rectified_this_turn = True
self._rectify_stage = stage
def consume_rectified_this_turn(self) -> bool:
if not self._rectified_this_turn:
return False
self._rectified_this_turn = False
return True
__all__ = ["MutableRequestBodyState", "RequestBodyState"]

View File

@@ -15,6 +15,7 @@ from src.services.task.execute.error_handler import TaskErrorOperationsService
from src.services.task.execute.failure import TaskFailureOperationsService
from src.services.task.execute.pool import TaskPoolOperationsService
from src.services.task.execute.sync_execute import SyncTaskExecutionService
from src.services.task.request_state import RequestBodyState
from src.services.task.submit.submit_service import AsyncTaskSubmitService
from src.services.task.video.facade import TaskVideoFacadeService
from src.services.task.video.operations import VideoTaskOperationsService
@@ -117,7 +118,7 @@ class TaskService:
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
preferred_key_ids: list[str] | None = None,
request_body_ref: dict[str, Any] | None = None,
request_body_state: RequestBodyState | None = None,
request_headers: dict[str, Any] | None = None,
request_body: dict[str, Any] | None = None,
# ASYNC-only (video submit)
@@ -138,7 +139,7 @@ class TaskService:
is_stream=is_stream,
capability_requirements=capability_requirements,
preferred_key_ids=preferred_key_ids,
request_body_ref=request_body_ref,
request_body_state=request_body_state,
request_headers=request_headers,
request_body=request_body,
extract_external_task_id=extract_external_task_id,
@@ -160,7 +161,7 @@ class TaskService:
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
preferred_key_ids: list[str] | None = None,
request_body_ref: dict[str, Any] | None = None,
request_body_state: RequestBodyState | None = None,
request_headers: dict[str, Any] | None = None,
request_body: dict[str, Any] | None = None,
extract_external_task_id: Any | None = None,
@@ -246,7 +247,7 @@ class TaskService:
is_stream=is_stream,
capability_requirements=capability_requirements,
preferred_key_ids=preferred_key_ids,
request_body_ref=request_body_ref,
request_body_state=request_body_state,
request_headers=request_headers,
request_body=request_body,
)
@@ -263,7 +264,7 @@ class TaskService:
user_api_key: ApiKey | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
request_body_ref: dict[str, Any] | None = None,
request_body_state: RequestBodyState | None = None,
request_headers: dict[str, Any] | None = None,
request_body: dict[str, Any] | None = None,
affinity_key: str | None = None,
@@ -411,7 +412,7 @@ class TaskService:
max_attempts = candidate_resolver.count_total_attempts(all_candidates)
execution_state = SyncExecutionState(
candidate_record_map=candidate_record_map,
request_body_ref=request_body_ref,
request_body_state=request_body_state,
last_candidate=all_candidates[-1] if all_candidates else None,
)
@@ -540,7 +541,7 @@ class TaskService:
request_id=request_id,
attempt=attempt_count,
max_attempts=int(max_attempts or 0),
request_body_ref=request_body_ref,
request_body_state=request_body_state,
error_classifier=error_classifier,
)
action = classify_candidate_error_action(raw_action)

View 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"

View File

@@ -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,
)

View 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

View File

@@ -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,
)

View File

@@ -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()

View File

@@ -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