deer-flow/backend/tests/test_ensure_sandbox_initialized.py
Aari 9e0fbd60fa
fix(sandbox): isolate concurrent subagent shell sessions (#5134)
* fix(sandbox): isolate concurrent subagent shell sessions

* fix(sandbox): make execution acquire idempotent

* fix(sandbox): close execution lifecycle gaps

* fix(sandbox): serialize retained client lifecycle

* fix(sandbox): close remaining client lifecycle gaps

* fix(sandbox): unwind failed client lookup

* fix(sandbox): protect internal lease identities

* fix(sandbox): make cancellation reconciliation durable

* fix(sandbox): fence cancelled workers and IM uploads
2026-09-02 21:05:23 +08:00

556 lines
19 KiB
Python

"""Tests for ensure_sandbox_initialized with fork-restored channel values."""
from __future__ import annotations
import asyncio
import threading
import pytest
from langchain.tools import ToolRuntime
from langgraph.types import Overwrite
from deerflow.sandbox.exceptions import SandboxNotFoundError
from deerflow.sandbox.lease import SANDBOX_LEASE_OWNER_CONTEXT_KEY, get_sandbox_lease_manager
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
from deerflow.sandbox.tools import (
_run_sync_tool_after_async_sandbox_init,
ensure_sandbox_initialized,
ensure_sandbox_initialized_async,
)
class _StubSandbox(Sandbox):
def __init__(self, sandbox_id: str) -> None:
super().__init__(sandbox_id)
self.released_scopes: list[str] = []
def execute_command(self, command: str, env: dict | None = None, timeout: float | None = None) -> str:
del env, timeout
return "OK"
def read_file(self, path: str) -> str:
return "content"
def download_file(self, path: str) -> bytes:
return b"content"
def list_dir(self, path: str, max_depth: int = 2) -> list[str]:
return ["/mnt/user-data/workspace/file.txt"]
def write_file(self, path: str, content: str, append: bool = False) -> None:
return None
def glob(self, path: str, pattern: str, *, include_dirs: bool = False, max_results: int = 200) -> tuple[list[str], bool]:
return [], False
def grep(
self,
path: str,
pattern: str,
*,
glob: str | None = None,
literal: bool = False,
case_sensitive: bool = False,
max_results: int = 100,
) -> tuple[list[GrepMatch], bool]:
return [], False
def update_file(self, path: str, content: bytes) -> None:
return None
def release_command_scope(self, scope_id: str) -> None:
self.released_scopes.append(scope_id)
class _RecordingProvider(SandboxProvider):
def __init__(self) -> None:
self.sandbox = _StubSandbox("stub")
self.released: list[str] = []
def acquire(self, thread_id: str | None = None, *, user_id: str | None = None) -> str:
raise AssertionError("state already carries a sandbox; acquire must not run")
async def acquire_async(self, thread_id: str | None = None, *, user_id: str | None = None) -> str:
raise AssertionError("state already carries a sandbox; acquire must not run")
def get(self, sandbox_id: str) -> Sandbox | None:
if sandbox_id == "parent-sandbox":
return self.sandbox
return None
def release(self, sandbox_id: str) -> None:
self.released.append(sandbox_id)
class _FallthroughProvider(SandboxProvider):
"""Provider whose parent id has expired, forcing a fresh acquire."""
def __init__(self) -> None:
self.sandbox = _StubSandbox("fresh")
self.acquired: list[str | None] = []
self.released: list[str] = []
def acquire(self, thread_id: str | None = None, *, user_id: str | None = None) -> str:
self.acquired.append(thread_id)
return "fresh-sandbox"
async def acquire_async(self, thread_id: str | None = None, *, user_id: str | None = None) -> str:
self.acquired.append(thread_id)
return "fresh-sandbox"
def get(self, sandbox_id: str) -> Sandbox | None:
if sandbox_id == "fresh-sandbox":
return self.sandbox
return None
def release(self, sandbox_id: str) -> None:
self.released.append(sandbox_id)
class _PostAcquireLookupFailureProvider(SandboxProvider):
"""Provider that binds an id but cannot return its active client."""
def __init__(self) -> None:
self.released: list[str] = []
def acquire(self, thread_id: str | None = None, *, user_id: str | None = None) -> str:
return "lost-after-acquire"
async def acquire_async(self, thread_id: str | None = None, *, user_id: str | None = None) -> str:
return "lost-after-acquire"
def get(self, sandbox_id: str) -> Sandbox | None:
return None
def release(self, sandbox_id: str) -> None:
self.released.append(sandbox_id)
def _make_runtime(state: dict) -> ToolRuntime:
return ToolRuntime(
state=state,
context={},
config={"configurable": {}},
stream_writer=lambda _: None,
tools=[],
tool_call_id="call-1",
store=None,
)
def test_post_acquire_lookup_failure_unwinds_sync_execution_lease() -> None:
provider = _PostAcquireLookupFailureProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({})
runtime.context.update(
{
SANDBOX_LEASE_OWNER_CONTEXT_KEY: "sync-owner",
"thread_id": "thread-1",
"user_id": "user-1",
}
)
manager = get_sandbox_lease_manager(provider)
with pytest.raises(SandboxNotFoundError, match="Sandbox not found after acquisition"):
ensure_sandbox_initialized(runtime)
assert manager.binding_for("sync-owner") is None
assert provider.released == ["lost-after-acquire"]
finally:
reset_sandbox_provider()
@pytest.mark.anyio
async def test_post_acquire_lookup_failure_unwinds_async_execution_lease() -> None:
provider = _PostAcquireLookupFailureProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({})
runtime.context.update(
{
SANDBOX_LEASE_OWNER_CONTEXT_KEY: "async-owner",
"thread_id": "thread-1",
"user_id": "user-1",
}
)
manager = get_sandbox_lease_manager(provider)
with pytest.raises(SandboxNotFoundError, match="Sandbox not found after acquisition"):
await ensure_sandbox_initialized_async(runtime)
assert manager.binding_for("async-owner") is None
assert provider.released == ["lost-after-acquire"]
finally:
reset_sandbox_provider()
def test_ensure_sandbox_initialized_unwraps_overwrite_state() -> None:
"""Fork-restored state must not crash on the Overwrite wrapper."""
provider = _RecordingProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": Overwrite({"sandbox_id": "parent-sandbox"})})
sandbox = ensure_sandbox_initialized(runtime)
finally:
reset_sandbox_provider()
assert sandbox is provider.sandbox
assert runtime.context["sandbox_id"] == "parent-sandbox"
# The reuse path must not take ownership: the wrapped state is left
# untouched, so after_agent still sees fork_restored and skips release.
assert isinstance(runtime.state["sandbox"], Overwrite)
@pytest.mark.anyio
async def test_ensure_sandbox_initialized_async_unwraps_overwrite_state() -> None:
provider = _RecordingProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": Overwrite({"sandbox_id": "parent-sandbox"})})
sandbox = await ensure_sandbox_initialized_async(runtime)
finally:
reset_sandbox_provider()
assert sandbox is provider.sandbox
assert runtime.context["sandbox_id"] == "parent-sandbox"
def test_fork_restored_owner_holds_non_releasing_scope_lease() -> None:
provider = _RecordingProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": Overwrite({"sandbox_id": "parent-sandbox"})})
runtime.context.update(
{
SANDBOX_LEASE_OWNER_CONTEXT_KEY: "fork-child",
"thread_id": "thread-1",
"user_id": "user-1",
}
)
sandbox = ensure_sandbox_initialized(runtime)
manager = get_sandbox_lease_manager(provider)
assert sandbox is provider.sandbox
assert manager.binding_for("fork-child") == "parent-sandbox"
assert isinstance(runtime.state["sandbox"], Overwrite)
manager.release("fork-child")
assert provider.sandbox.released_scopes == ["fork-child"]
assert provider.released == []
finally:
reset_sandbox_provider()
@pytest.mark.anyio
async def test_async_fork_restored_owner_holds_non_releasing_scope_lease() -> None:
provider = _RecordingProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": Overwrite({"sandbox_id": "parent-sandbox"})})
runtime.context.update(
{
SANDBOX_LEASE_OWNER_CONTEXT_KEY: "fork-child",
"thread_id": "thread-1",
"user_id": "user-1",
}
)
sandbox = await ensure_sandbox_initialized_async(runtime)
manager = get_sandbox_lease_manager(provider)
assert sandbox is provider.sandbox
assert manager.binding_for("fork-child") == "parent-sandbox"
assert isinstance(runtime.state["sandbox"], Overwrite)
await manager.release_async("fork-child")
assert provider.sandbox.released_scopes == ["fork-child"]
assert provider.released == []
finally:
reset_sandbox_provider()
def test_ensure_sandbox_initialized_plain_state_unchanged() -> None:
provider = _RecordingProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": {"sandbox_id": "parent-sandbox"}})
sandbox = ensure_sandbox_initialized(runtime)
finally:
reset_sandbox_provider()
assert sandbox is provider.sandbox
assert runtime.context["sandbox_id"] == "parent-sandbox"
def test_ensure_sandbox_initialized_acquires_fresh_when_parent_missing() -> None:
"""Acquire fall-through: the fork-restored id is gone from the provider,
so a fresh sandbox is acquired and the stale wrapped state is replaced
by the freshly acquired plain dict."""
provider = _FallthroughProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": Overwrite({"sandbox_id": "parent-sandbox"})})
runtime.context["thread_id"] = "t-1"
sandbox = ensure_sandbox_initialized(runtime)
finally:
reset_sandbox_provider()
assert provider.acquired == ["t-1"]
assert sandbox is provider.sandbox
assert runtime.state["sandbox"] == {"sandbox_id": "fresh-sandbox"}
assert runtime.context["sandbox_id"] == "fresh-sandbox"
def test_fork_restored_owner_normally_releases_fresh_replacement() -> None:
provider = _FallthroughProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": Overwrite({"sandbox_id": "parent-sandbox"})})
runtime.context.update(
{
SANDBOX_LEASE_OWNER_CONTEXT_KEY: "fork-child",
"thread_id": "thread-1",
"user_id": "user-1",
}
)
sandbox = ensure_sandbox_initialized(runtime)
manager = get_sandbox_lease_manager(provider)
assert sandbox is provider.sandbox
assert manager.binding_for("fork-child") == "fresh-sandbox"
manager.release("fork-child")
assert provider.sandbox.released_scopes == ["fork-child"]
assert provider.released == ["fresh-sandbox"]
finally:
reset_sandbox_provider()
@pytest.mark.anyio
async def test_async_fork_restored_owner_normally_releases_fresh_replacement() -> None:
provider = _FallthroughProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": Overwrite({"sandbox_id": "parent-sandbox"})})
runtime.context.update(
{
SANDBOX_LEASE_OWNER_CONTEXT_KEY: "fork-child",
"thread_id": "thread-1",
"user_id": "user-1",
}
)
sandbox = await ensure_sandbox_initialized_async(runtime)
manager = get_sandbox_lease_manager(provider)
assert sandbox is provider.sandbox
assert manager.binding_for("fork-child") == "fresh-sandbox"
await manager.release_async("fork-child")
assert provider.sandbox.released_scopes == ["fork-child"]
assert provider.released == ["fresh-sandbox"]
finally:
reset_sandbox_provider()
@pytest.mark.anyio
async def test_ensure_sandbox_initialized_async_plain_state_unchanged() -> None:
provider = _RecordingProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": {"sandbox_id": "parent-sandbox"}})
sandbox = await ensure_sandbox_initialized_async(runtime)
finally:
reset_sandbox_provider()
assert sandbox is provider.sandbox
assert runtime.context["sandbox_id"] == "parent-sandbox"
def test_reuse_with_config_only_thread_id_binds_execution_owner() -> None:
provider = _RecordingProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": {"sandbox_id": "parent-sandbox"}})
runtime.context[SANDBOX_LEASE_OWNER_CONTEXT_KEY] = "config-owner"
runtime.config["configurable"]["thread_id"] = "thread-from-config"
sandbox = ensure_sandbox_initialized(runtime)
assert sandbox is provider.sandbox
assert get_sandbox_lease_manager(provider).binding_for("config-owner") == "parent-sandbox"
finally:
reset_sandbox_provider()
@pytest.mark.anyio
async def test_async_reuse_with_config_only_thread_id_binds_execution_owner() -> None:
provider = _RecordingProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": {"sandbox_id": "parent-sandbox"}})
runtime.context[SANDBOX_LEASE_OWNER_CONTEXT_KEY] = "config-owner"
runtime.config["configurable"]["thread_id"] = "thread-from-config"
sandbox = await ensure_sandbox_initialized_async(runtime)
assert sandbox is provider.sandbox
assert get_sandbox_lease_manager(provider).binding_for("config-owner") == "parent-sandbox"
finally:
reset_sandbox_provider()
def test_reuse_replaces_stale_checkpoint_and_owner_binding() -> None:
provider = _FallthroughProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": {"sandbox_id": "parent-sandbox"}})
runtime.context[SANDBOX_LEASE_OWNER_CONTEXT_KEY] = "stale-owner"
runtime.context["thread_id"] = "thread-1"
runtime.context["user_id"] = "user-1"
manager = get_sandbox_lease_manager(provider)
manager.retain(
"stale-owner",
"parent-sandbox",
thread_id="thread-1",
user_id="user-1",
)
sandbox = ensure_sandbox_initialized(runtime)
assert sandbox is provider.sandbox
assert runtime.state["sandbox"] == {"sandbox_id": "fresh-sandbox"}
assert runtime.context["sandbox_id"] == "fresh-sandbox"
assert manager.binding_for("stale-owner") == "fresh-sandbox"
assert provider.acquired == ["thread-1"]
finally:
reset_sandbox_provider()
@pytest.mark.anyio
async def test_async_reuse_replaces_stale_checkpoint_and_owner_binding() -> None:
provider = _FallthroughProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": {"sandbox_id": "parent-sandbox"}})
runtime.context[SANDBOX_LEASE_OWNER_CONTEXT_KEY] = "stale-owner"
runtime.context["thread_id"] = "thread-1"
runtime.context["user_id"] = "user-1"
manager = get_sandbox_lease_manager(provider)
manager.retain(
"stale-owner",
"parent-sandbox",
thread_id="thread-1",
user_id="user-1",
)
sandbox = await ensure_sandbox_initialized_async(runtime)
assert sandbox is provider.sandbox
assert runtime.state["sandbox"] == {"sandbox_id": "fresh-sandbox"}
assert runtime.context["sandbox_id"] == "fresh-sandbox"
assert manager.binding_for("stale-owner") == "fresh-sandbox"
assert provider.acquired == ["thread-1"]
finally:
reset_sandbox_provider()
@pytest.mark.anyio
async def test_ensure_sandbox_initialized_async_acquires_fresh_when_parent_missing() -> None:
"""Same fall-through as the sync path: the fork-restored id is gone from
the provider, so a fresh sandbox is acquired and the stale wrapped state
is replaced by the freshly acquired plain dict."""
provider = _FallthroughProvider()
set_sandbox_provider(provider)
try:
runtime = _make_runtime({"sandbox": Overwrite({"sandbox_id": "parent-sandbox"})})
runtime.context["thread_id"] = "t-1"
sandbox = await ensure_sandbox_initialized_async(runtime)
finally:
reset_sandbox_provider()
assert provider.acquired == ["t-1"]
assert sandbox is provider.sandbox
assert runtime.state["sandbox"] == {"sandbox_id": "fresh-sandbox"}
assert runtime.context["sandbox_id"] == "fresh-sandbox"
@pytest.mark.anyio
async def test_cancelled_async_tool_drains_worker_before_execution_lease_cleanup(monkeypatch) -> None:
"""Cancellation must not let a late sync body re-admit a released owner."""
provider = _FallthroughProvider()
set_sandbox_provider(provider)
worker_started = threading.Event()
allow_worker = threading.Event()
worker_finished = threading.Event()
try:
runtime = _make_runtime({"sandbox": {"sandbox_id": "fresh-sandbox"}})
runtime.context.update(
{
SANDBOX_LEASE_OWNER_CONTEXT_KEY: "cancelled-child",
"thread_id": "thread-1",
"user_id": "user-1",
}
)
manager = get_sandbox_lease_manager(provider)
await manager.acquire_async("cancelled-child", "thread-1", user_id="user-1")
await manager.acquire_async("parallel-sibling", "thread-1", user_id="user-1")
async def _allow_sandbox(*, context, app_config):
del context, app_config
async def _safe_config():
return None
monkeypatch.setattr("deerflow.sandbox.tools.authorize_sandbox_execution_async", _allow_sandbox)
monkeypatch.setattr("deerflow.sandbox.tools.safe_app_config_async", _safe_config)
def _blocking_tool(inner_runtime: ToolRuntime) -> str:
worker_started.set()
assert allow_worker.wait(timeout=2)
try:
sandbox = ensure_sandbox_initialized(inner_runtime)
return sandbox.execute_command("late command")
finally:
worker_finished.set()
async def _run_then_cleanup() -> str:
try:
return await _run_sync_tool_after_async_sandbox_init(_blocking_tool, runtime)
finally:
await manager.release_async("cancelled-child")
execution = asyncio.create_task(_run_then_cleanup())
assert await asyncio.to_thread(worker_started.wait, 1)
for _ in range(3):
execution.cancel()
await asyncio.sleep(0)
assert not execution.done()
assert manager.binding_for("cancelled-child") == "fresh-sandbox"
assert provider.released == []
allow_worker.set()
with pytest.raises(asyncio.CancelledError):
await execution
assert worker_finished.is_set()
assert manager.binding_for("cancelled-child") is None
assert manager.binding_for("parallel-sibling") == "fresh-sandbox"
assert provider.released == []
await manager.release_async("parallel-sibling")
assert provider.released == ["fresh-sandbox"]
finally:
allow_worker.set()
await asyncio.to_thread(worker_finished.wait, 2)
reset_sandbox_provider()