mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +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:
@@ -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}",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
64
src/services/task/request_state.py
Normal file
64
src/services/task/request_state.py
Normal 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"]
|
||||
@@ -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)
|
||||
|
||||
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