mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
refactor(task): 引入 MutableRequestBodyState 替代 request_body_ref 字典容器
将请求体可变状态从 {"body": dict} 字典容器重构为独立的
MutableRequestBodyState 类,统一管理 original_body / current_body /
build_attempt_body / rectify 等语义,消除各层通过 ref["body"] 间接
读写的隐式约定。
- 新增 src/services/task/request_state.py 定义 Protocol 与实现
- handler/executor/mixin 层改用 request_state 参数传递
- error_handler/state_transition 通过 request_state 判断整流状态
- 新增 request_state 单元测试与 chat/cli 请求体隔离测试
This commit is contained in:
@@ -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}",
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
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.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)
|
||||||
|
|||||||
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_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,
|
||||||
)
|
)
|
||||||
|
|||||||
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,
|
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,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user