fix: preserve sandbox reducer in middleware state (#3629)

* fix: preserve sandbox reducer in middleware state

* refactor: share sandbox state field annotation
This commit is contained in:
Huixin615 2026-06-19 16:10:56 +08:00 committed by GitHub
parent 3e5c76eb0a
commit 29489c0f45
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
3 changed files with 17 additions and 4 deletions

View File

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

View File

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

View File

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