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

View File

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

View File

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

View File

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

View File

@@ -2,7 +2,6 @@
from __future__ import annotations from __future__ import annotations
import copy
import json import json
import time import time
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
@@ -33,6 +32,7 @@ from src.core.exceptions import (
) )
from src.core.logger import logger from src.core.logger import logger
from src.services.scheduling.aware_scheduler import ProviderCandidate from src.services.scheduling.aware_scheduler import ProviderCandidate
from src.services.task.request_state import MutableRequestBodyState
if TYPE_CHECKING: if TYPE_CHECKING:
from src.api.handlers.base.cli_protocol import CliHandlerProtocol from src.api.handlers.base.cli_protocol import CliHandlerProtocol
@@ -100,9 +100,7 @@ class CliSyncMixin:
needs_conversion = False # 是否需要格式转换(由 candidate 决定) needs_conversion = False # 是否需要格式转换(由 candidate 决定)
sync_proxy_info: dict[str, Any] | None = None # 代理信息 sync_proxy_info: dict[str, Any] | None = None # 代理信息
# 可变请求体容器:允许 TaskService 在遇到 Thinking 签名错误时整流请求体后重试 request_state = MutableRequestBodyState(original_request_body)
# 结构: {"body": 实际请求体, "_rectified": 是否已整流, "_rectified_this_turn": 本轮是否整流}
request_body_ref: dict[str, Any] = {"body": copy.deepcopy(original_request_body)}
async def sync_request_func( async def sync_request_func(
provider: "Provider", provider: "Provider",
@@ -122,8 +120,7 @@ class CliSyncMixin:
provider_id=str(provider.id), provider_id=str(provider.id),
) )
# 应用模型映射到请求体(子类可覆盖此方法处理不同格式) request_body = request_state.build_attempt_body()
request_body = copy.deepcopy(request_body_ref["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)
@@ -390,7 +387,7 @@ class CliSyncMixin:
is_stream=False, is_stream=False,
capability_requirements=capability_requirements or None, capability_requirements=capability_requirements or None,
preferred_key_ids=preferred_key_ids 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_headers=original_headers,
request_body=original_request_body, 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.core.provider_types import ProviderType
from src.services.request.candidate import RequestCandidateService from src.services.request.candidate import RequestCandidateService
from src.services.task.execute.pool import TaskPoolOperationsService from src.services.task.execute.pool import TaskPoolOperationsService
from src.services.task.request_state import RequestBodyState
class TaskErrorOperationsService: class TaskErrorOperationsService:
@@ -60,7 +61,7 @@ class TaskErrorOperationsService:
elapsed_ms: int, elapsed_ms: int,
captured_key_concurrent: int | None, captured_key_concurrent: int | None,
serializable_extra_data: dict[str, Any], serializable_extra_data: dict[str, Any],
request_body_ref: dict[str, Any] | None, request_body_state: RequestBodyState | None,
) -> str: ) -> str:
"""Try to rectify thinking signature errors and request a retry.""" """Try to rectify thinking signature errors and request a retry."""
from src.services.message.thinking_rectifier import ThinkingRectifier from src.services.message.thinking_rectifier import ThinkingRectifier
@@ -79,7 +80,7 @@ class TaskErrorOperationsService:
) )
raise converted_error raise converted_error
if request_body_ref is None: if request_body_state is None:
logger.warning(" [{}] Thinking 错误:无法获取请求体引用,终止重试", request_id) logger.warning(" [{}] Thinking 错误:无法获取请求体引用,终止重试", request_id)
self.mark_thinking_error_failed( self.mark_thinking_error_failed(
candidate_record_id, candidate_record_id,
@@ -93,12 +94,8 @@ class TaskErrorOperationsService:
provider_type_norm = str(provider_type or "").lower() provider_type_norm = str(provider_type or "").lower()
# Rectification may have multiple stages (Antigravity only). # Rectification may have multiple stages (Antigravity only).
stage_raw = request_body_ref.get("_rectify_stage", 0) stage = request_body_state.rectify_stage()
try: if stage <= 0 and request_body_state.is_rectified():
stage = int(stage_raw or 0)
except Exception:
stage = 0
if stage <= 0 and request_body_ref.get("_rectified", False):
stage = 1 stage = 1
if stage >= 2 or (stage >= 1 and provider_type_norm != ProviderType.ANTIGRAVITY): if stage >= 2 or (stage >= 1 and provider_type_norm != ProviderType.ANTIGRAVITY):
@@ -112,7 +109,7 @@ class TaskErrorOperationsService:
) )
raise converted_error raise converted_error
request_body = request_body_ref.get("body", {}) request_body = request_body_state.current_body
stage_label = "thinking_only" stage_label = "thinking_only"
next_stage = 1 next_stage = 1
@@ -129,10 +126,7 @@ class TaskErrorOperationsService:
next_stage = 2 next_stage = 2
if modified: if modified:
request_body_ref["body"] = rectified_body request_body_state.mark_rectified(rectified_body, stage=next_stage)
request_body_ref["_rectified"] = True
request_body_ref["_rectified_this_turn"] = True
request_body_ref["_rectify_stage"] = next_stage
if provider_type_norm == ProviderType.ANTIGRAVITY: if provider_type_norm == ProviderType.ANTIGRAVITY:
try: try:
@@ -188,7 +182,7 @@ class TaskErrorOperationsService:
attempt: int, attempt: int,
max_attempts: int, max_attempts: int,
error_classifier: Any, error_classifier: Any,
request_body_ref: dict[str, Any] | None = None, request_body_state: RequestBodyState | None = None,
) -> str: ) -> str:
""" """
Handle an execution error for a candidate. Handle an execution error for a candidate.
@@ -426,7 +420,7 @@ class TaskErrorOperationsService:
elapsed_ms=elapsed_ms, elapsed_ms=elapsed_ms,
captured_key_concurrent=captured_key_concurrent, captured_key_concurrent=captured_key_concurrent,
serializable_extra_data=serializable_extra_data, serializable_extra_data=serializable_extra_data,
request_body_ref=request_body_ref, request_body_state=request_body_state,
) )
if action == "continue": if action == "continue":
return "continue" return "continue"

View File

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

View File

@@ -31,6 +31,7 @@ from src.services.task.execute.state_transition import (
SyncExecutionState, SyncExecutionState,
resolve_execution_error_transition, resolve_execution_error_transition,
) )
from src.services.task.request_state import RequestBodyState
from src.services.usage.service import UsageService from src.services.usage.service import UsageService
@@ -65,7 +66,7 @@ class SyncTaskExecutionService:
is_stream: bool, is_stream: bool,
capability_requirements: dict[str, bool] | None, capability_requirements: dict[str, bool] | None,
preferred_key_ids: list[str] | 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_headers: dict[str, Any] | None,
request_body: dict[str, Any] | None, request_body: dict[str, Any] | None,
) -> ExecutionResult: ) -> ExecutionResult:
@@ -197,7 +198,7 @@ class SyncTaskExecutionService:
# Keep behavior consistent with previous behavior: last_candidate is updated even if skipped. # Keep behavior consistent with previous behavior: last_candidate is updated even if skipped.
execution_state = SyncExecutionState( execution_state = SyncExecutionState(
candidate_record_map=candidate_record_map, 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, last_candidate=all_candidates[-1] if all_candidates else None,
) )
@@ -334,7 +335,7 @@ class SyncTaskExecutionService:
request_id=request_id, request_id=request_id,
attempt=attempt_count, attempt=attempt_count,
max_attempts=int(max_attempts or 0), max_attempts=int(max_attempts or 0),
request_body_ref=request_body_ref, request_body_state=request_body_state,
error_classifier=error_classifier, error_classifier=error_classifier,
) )
action = classify_candidate_error_action(raw_action) 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.failure import TaskFailureOperationsService
from src.services.task.execute.pool import TaskPoolOperationsService from src.services.task.execute.pool import TaskPoolOperationsService
from src.services.task.execute.sync_execute import SyncTaskExecutionService 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.submit.submit_service import AsyncTaskSubmitService
from src.services.task.video.facade import TaskVideoFacadeService from src.services.task.video.facade import TaskVideoFacadeService
from src.services.task.video.operations import VideoTaskOperationsService from src.services.task.video.operations import VideoTaskOperationsService
@@ -117,7 +118,7 @@ class TaskService:
is_stream: bool = False, is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None, capability_requirements: dict[str, bool] | None = None,
preferred_key_ids: list[str] | 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_headers: dict[str, Any] | None = None,
request_body: dict[str, Any] | None = None, request_body: dict[str, Any] | None = None,
# ASYNC-only (video submit) # ASYNC-only (video submit)
@@ -138,7 +139,7 @@ class TaskService:
is_stream=is_stream, is_stream=is_stream,
capability_requirements=capability_requirements, capability_requirements=capability_requirements,
preferred_key_ids=preferred_key_ids, preferred_key_ids=preferred_key_ids,
request_body_ref=request_body_ref, request_body_state=request_body_state,
request_headers=request_headers, request_headers=request_headers,
request_body=request_body, request_body=request_body,
extract_external_task_id=extract_external_task_id, extract_external_task_id=extract_external_task_id,
@@ -160,7 +161,7 @@ class TaskService:
is_stream: bool = False, is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None, capability_requirements: dict[str, bool] | None = None,
preferred_key_ids: list[str] | 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_headers: dict[str, Any] | None = None,
request_body: dict[str, Any] | None = None, request_body: dict[str, Any] | None = None,
extract_external_task_id: Any | None = None, extract_external_task_id: Any | None = None,
@@ -246,7 +247,7 @@ class TaskService:
is_stream=is_stream, is_stream=is_stream,
capability_requirements=capability_requirements, capability_requirements=capability_requirements,
preferred_key_ids=preferred_key_ids, preferred_key_ids=preferred_key_ids,
request_body_ref=request_body_ref, request_body_state=request_body_state,
request_headers=request_headers, request_headers=request_headers,
request_body=request_body, request_body=request_body,
) )
@@ -263,7 +264,7 @@ class TaskService:
user_api_key: ApiKey | None = None, user_api_key: ApiKey | None = None,
is_stream: bool = False, is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None, 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_headers: dict[str, Any] | None = None,
request_body: dict[str, Any] | None = None, request_body: dict[str, Any] | None = None,
affinity_key: str | None = None, affinity_key: str | None = None,
@@ -411,7 +412,7 @@ class TaskService:
max_attempts = candidate_resolver.count_total_attempts(all_candidates) max_attempts = candidate_resolver.count_total_attempts(all_candidates)
execution_state = SyncExecutionState( execution_state = SyncExecutionState(
candidate_record_map=candidate_record_map, 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, last_candidate=all_candidates[-1] if all_candidates else None,
) )
@@ -540,7 +541,7 @@ class TaskService:
request_id=request_id, request_id=request_id,
attempt=attempt_count, attempt=attempt_count,
max_attempts=int(max_attempts or 0), max_attempts=int(max_attempts or 0),
request_body_ref=request_body_ref, request_body_state=request_body_state,
error_classifier=error_classifier, error_classifier=error_classifier,
) )
action = classify_candidate_error_action(raw_action) 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_request_mixin import CliRequestMixin
from src.api.handlers.base.cli_stream_mixin import CliStreamMixin from src.api.handlers.base.cli_stream_mixin import CliStreamMixin
from src.api.handlers.base.stream_context import StreamContext from src.api.handlers.base.stream_context import StreamContext
from src.services.task.request_state import MutableRequestBodyState
class _StopBuild(Exception): 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) snapshot = copy.deepcopy(original_request_body)
request_state = MutableRequestBodyState(original_request_body)
with pytest.raises(_StopBuild): with pytest.raises(_StopBuild):
await handler._execute_stream_request( await handler._execute_stream_request(
@@ -142,7 +144,7 @@ async def test_execute_stream_request_does_not_mutate_original_request_body(
provider, provider,
endpoint, endpoint,
key, key,
original_request_body, request_state.build_attempt_body(),
{}, {},
candidate=candidate, 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, is_stream=True,
capability_requirements=None, capability_requirements=None,
preferred_key_ids=None, preferred_key_ids=None,
request_body_ref=None, request_body_state=None,
request_headers=None, request_headers=None,
request_body=None, request_body=None,
) )

View File

@@ -14,6 +14,7 @@ from src.services.task.execute.state_transition import (
SyncExecutionState, SyncExecutionState,
resolve_execution_error_transition, resolve_execution_error_transition,
) )
from src.services.task.request_state import MutableRequestBodyState
def _make_candidate() -> SimpleNamespace: 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: 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( state = SyncExecutionState(
candidate_record_map={}, candidate_record_map={},
request_body_ref=request_body_ref, request_body_state=request_body_state,
) )
transition = resolve_execution_error_transition( 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.failover_action == FailoverAction.RETRY
assert transition.max_retries == 3 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: def test_resolve_execution_error_transition_next_candidate() -> None:
request_body_state = MutableRequestBodyState({})
request_body_state.mark_rectified({}, stage=1)
state = SyncExecutionState( state = SyncExecutionState(
candidate_record_map={}, candidate_record_map={},
request_body_ref={"_rectified_this_turn": True}, request_body_state=request_body_state,
) )
transition = resolve_execution_error_transition( 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.failover_action == FailoverAction.CONTINUE
assert transition.max_retries is None 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: def test_sync_execution_state_resolve_candidate_record_id_fallback() -> None:
state = SyncExecutionState( state = SyncExecutionState(
candidate_record_map={(2, 0): "r20"}, 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" 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() candidate = _make_candidate()
state = SyncExecutionState( state = SyncExecutionState(
candidate_record_map={}, candidate_record_map={},
request_body_ref=None, request_body_state=None,
last_error=err, last_error=err,
last_candidate=candidate, 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: def test_sync_execution_state_raise_classified_error_fallback_error() -> None:
state = SyncExecutionState( state = SyncExecutionState(
candidate_record_map={}, candidate_record_map={},
request_body_ref=None, request_body_state=None,
) )
failure_ops = MagicMock() 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.candidate.submit import SubmitOutcome
from src.services.task.core.context import TaskMode from src.services.task.core.context import TaskMode
from src.services.task.core.protocol import AttemptKind from src.services.task.core.protocol import AttemptKind
from src.services.task.service import pool_on_error from src.services.task.request_state import MutableRequestBodyState
from src.services.task.service import TaskService from src.services.task.service import TaskService, pool_on_error
@pytest.mark.asyncio @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_headers = {"authorization": "Bearer test", "x-trace-id": "abc123"}
request_body = {"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]} 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( result = await svc.execute(
task_type="chat", task_type="chat",
@@ -100,7 +100,7 @@ async def test_task_service_execute_sync_passes_request_headers_and_body() -> No
is_stream=True, is_stream=True,
request_headers=request_headers, request_headers=request_headers,
request_body=request_body, request_body=request_body,
request_body_ref=request_body_ref, request_body_state=request_body_state,
) )
assert result is sentinel_result 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] 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_headers"] == request_headers
assert kwargs["request_body"] == request_body 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 @pytest.mark.asyncio