mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-08 21:19:59 +00:00
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:
parent
3e5c76eb0a
commit
29489c0f45
@ -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}")
|
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]:
|
def merge_artifacts(existing: list[str] | None, new: list[str] | None) -> list[str]:
|
||||||
"""Reducer for artifacts list - merges and deduplicates artifacts."""
|
"""Reducer for artifacts list - merges and deduplicates artifacts."""
|
||||||
if existing is None:
|
if existing is None:
|
||||||
@ -106,7 +109,7 @@ def merge_promoted(existing: PromotedTools | None, new: PromotedTools | None) ->
|
|||||||
|
|
||||||
|
|
||||||
class ThreadState(AgentState):
|
class ThreadState(AgentState):
|
||||||
sandbox: Annotated[NotRequired[SandboxState | None], merge_sandbox]
|
sandbox: SandboxStateField
|
||||||
thread_data: NotRequired[ThreadDataState | None]
|
thread_data: NotRequired[ThreadDataState | None]
|
||||||
title: NotRequired[str | None]
|
title: NotRequired[str | None]
|
||||||
artifacts: Annotated[list[str], merge_artifacts]
|
artifacts: Annotated[list[str], merge_artifacts]
|
||||||
|
|||||||
@ -11,7 +11,7 @@ from langgraph.prebuilt.tool_node import ToolCallRequest
|
|||||||
from langgraph.runtime import Runtime
|
from langgraph.runtime import Runtime
|
||||||
from langgraph.types import Command
|
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
|
from deerflow.sandbox import get_sandbox_provider
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@ -20,7 +20,7 @@ logger = logging.getLogger(__name__)
|
|||||||
class SandboxMiddlewareState(AgentState):
|
class SandboxMiddlewareState(AgentState):
|
||||||
"""Compatible with the `ThreadState` schema."""
|
"""Compatible with the `ThreadState` schema."""
|
||||||
|
|
||||||
sandbox: NotRequired[SandboxState | None]
|
sandbox: SandboxStateField
|
||||||
thread_data: NotRequired[ThreadDataState | None]
|
thread_data: NotRequired[ThreadDataState | None]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -1,6 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
from typing import get_type_hints
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from langchain.agents.middleware import AgentMiddleware
|
from langchain.agents.middleware import AgentMiddleware
|
||||||
@ -10,7 +11,8 @@ from langgraph.prebuilt.tool_node import ToolCallRequest
|
|||||||
from langgraph.runtime import Runtime
|
from langgraph.runtime import Runtime
|
||||||
from langgraph.types import Command
|
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 import Sandbox
|
||||||
from deerflow.sandbox.sandbox_provider import SandboxProvider, reset_sandbox_provider, set_sandbox_provider
|
from deerflow.sandbox.sandbox_provider import SandboxProvider, reset_sandbox_provider, set_sandbox_provider
|
||||||
from deerflow.sandbox.search import GrepMatch
|
from deerflow.sandbox.search import GrepMatch
|
||||||
@ -90,6 +92,14 @@ class _AsyncOnlyProvider(SandboxProvider):
|
|||||||
return None
|
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
|
@pytest.mark.anyio
|
||||||
async def test_provider_default_acquire_async_offloads_sync_acquire(monkeypatch: pytest.MonkeyPatch) -> None:
|
async def test_provider_default_acquire_async_offloads_sync_acquire(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
provider = _SyncProvider()
|
provider = _SyncProvider()
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user