From 29489c0f45fb957255f8ead03dfa512879b321e6 Mon Sep 17 00:00:00 2001 From: Huixin615 Date: Fri, 19 Jun 2026 16:10:56 +0800 Subject: [PATCH] fix: preserve sandbox reducer in middleware state (#3629) * fix: preserve sandbox reducer in middleware state * refactor: share sandbox state field annotation --- .../packages/harness/deerflow/agents/thread_state.py | 5 ++++- .../packages/harness/deerflow/sandbox/middleware.py | 4 ++-- backend/tests/test_sandbox_middleware.py | 12 +++++++++++- 3 files changed, 17 insertions(+), 4 deletions(-) diff --git a/backend/packages/harness/deerflow/agents/thread_state.py b/backend/packages/harness/deerflow/agents/thread_state.py index 8e5690779..b88b58129 100644 --- a/backend/packages/harness/deerflow/agents/thread_state.py +++ b/backend/packages/harness/deerflow/agents/thread_state.py @@ -39,6 +39,9 @@ def merge_sandbox(existing: SandboxState | None, new: SandboxState | None) -> Sa raise ValueError(f"Conflicting sandbox state updates: {existing_id!r} != {new_id!r}") +SandboxStateField = Annotated[NotRequired[SandboxState | None], merge_sandbox] + + def merge_artifacts(existing: list[str] | None, new: list[str] | None) -> list[str]: """Reducer for artifacts list - merges and deduplicates artifacts.""" if existing is None: @@ -106,7 +109,7 @@ def merge_promoted(existing: PromotedTools | None, new: PromotedTools | None) -> class ThreadState(AgentState): - sandbox: Annotated[NotRequired[SandboxState | None], merge_sandbox] + sandbox: SandboxStateField thread_data: NotRequired[ThreadDataState | None] title: NotRequired[str | None] artifacts: Annotated[list[str], merge_artifacts] diff --git a/backend/packages/harness/deerflow/sandbox/middleware.py b/backend/packages/harness/deerflow/sandbox/middleware.py index 5bdb5a700..0597cc557 100644 --- a/backend/packages/harness/deerflow/sandbox/middleware.py +++ b/backend/packages/harness/deerflow/sandbox/middleware.py @@ -11,7 +11,7 @@ from langgraph.prebuilt.tool_node import ToolCallRequest from langgraph.runtime import Runtime from langgraph.types import Command -from deerflow.agents.thread_state import SandboxState, ThreadDataState +from deerflow.agents.thread_state import SandboxStateField, ThreadDataState from deerflow.sandbox import get_sandbox_provider logger = logging.getLogger(__name__) @@ -20,7 +20,7 @@ logger = logging.getLogger(__name__) class SandboxMiddlewareState(AgentState): """Compatible with the `ThreadState` schema.""" - sandbox: NotRequired[SandboxState | None] + sandbox: SandboxStateField thread_data: NotRequired[ThreadDataState | None] diff --git a/backend/tests/test_sandbox_middleware.py b/backend/tests/test_sandbox_middleware.py index c584c759a..188ea242a 100644 --- a/backend/tests/test_sandbox_middleware.py +++ b/backend/tests/test_sandbox_middleware.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +from typing import get_type_hints import pytest from langchain.agents.middleware import AgentMiddleware @@ -10,7 +11,8 @@ from langgraph.prebuilt.tool_node import ToolCallRequest from langgraph.runtime import Runtime from langgraph.types import Command -from deerflow.sandbox.middleware import SandboxMiddleware +from deerflow.agents.thread_state import ThreadState +from deerflow.sandbox.middleware import SandboxMiddleware, SandboxMiddlewareState from deerflow.sandbox.sandbox import Sandbox from deerflow.sandbox.sandbox_provider import SandboxProvider, reset_sandbox_provider, set_sandbox_provider from deerflow.sandbox.search import GrepMatch @@ -90,6 +92,14 @@ class _AsyncOnlyProvider(SandboxProvider): return None +def test_sandbox_middleware_state_matches_thread_state_sandbox_field() -> None: + """Middleware-local schema must not drift from ThreadState.sandbox.""" + middleware_hints = get_type_hints(SandboxMiddlewareState, include_extras=True) + thread_hints = get_type_hints(ThreadState, include_extras=True) + + assert middleware_hints["sandbox"] == thread_hints["sandbox"] + + @pytest.mark.anyio async def test_provider_default_acquire_async_offloads_sync_acquire(monkeypatch: pytest.MonkeyPatch) -> None: provider = _SyncProvider()