mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-07-23 06:28:34 +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}")
|
||||
|
||||
|
||||
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]
|
||||
|
||||
@ -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]
|
||||
|
||||
|
||||
|
||||
@ -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()
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user