from __future__ import annotations from dataclasses import dataclass, field from typing import cast import pytest from dify_agent.runtime_backend import ( BindingCreateError, ExecutionBindingCreateSpec, ExecutionBindingDestroySpec, HomeSnapshotCreateSpec, SharedWorkspaceUnsupportedError, WorkspacePreservationUnsupportedError, ) from dify_agent.runtime_backend.e2b import ( E2BExecutionBindingBackend, E2BHomeSnapshotBackend, E2BRuntimeLease, ) from dify_agent.runtime_backend.shellctl import ShellctlRuntimeLease @dataclass(slots=True) class _Files: paths: set[str] = field(default_factory=set) async def make_dir(self, path: str) -> bool: self.paths.add(path) return True async def exists(self, path: str) -> bool: return path in self.paths async def remove(self, path: str) -> None: self.paths.discard(path) @dataclass(slots=True) class _Snapshot: snapshot_id: str names: list[str] = field(default_factory=list) @dataclass(slots=True) class _Sandbox: sandbox_id: str files: _Files = field(default_factory=_Files) traffic_access_token: str | None = "traffic-token" pauses: list[bool] = field(default_factory=list) killed: int = 0 snapshots: int = 0 pause_error: Exception | None = None def get_host(self, port: int) -> str: return f"{self.sandbox_id}-{port}.example.test" async def pause(self, keep_memory: bool = True) -> bool: self.pauses.append(keep_memory) if self.pause_error is not None: raise self.pause_error return True async def kill(self) -> bool: self.killed += 1 return True async def create_snapshot(self, name: str | None = None) -> _Snapshot: del name self.snapshots += 1 return _Snapshot(snapshot_id=f"snapshot-{self.sandbox_id}-{self.snapshots}") @dataclass(slots=True) class _ControlPlane: created: list[tuple[str, str]] = field(default_factory=list) sandboxes: dict[str, _Sandbox] = field(default_factory=dict) killed: list[str] = field(default_factory=list) deleted_snapshots: list[str] = field(default_factory=list) pause_error: Exception | None = None async def create(self, template: str, *, timeout: int, metadata: dict[str, str], on_timeout: str) -> _Sandbox: del timeout sandbox_id = f"sandbox-{len(self.sandboxes) + 1}" sandbox = _Sandbox(sandbox_id=sandbox_id, pause_error=self.pause_error) self.sandboxes[sandbox_id] = sandbox self.created.append((template, on_timeout)) assert metadata["dify.resource"] == "runtime-sandbox" return sandbox async def connect(self, handle: str, *, timeout: int) -> _Sandbox: del timeout return self.sandboxes[handle] async def kill(self, handle: str) -> bool: self.killed.append(handle) return True async def delete_snapshot(self, snapshot_ref: str) -> bool: self.deleted_snapshots.append(snapshot_ref) return True @pytest.mark.anyio async def test_e2b_binding_uses_default_template_or_exact_snapshot_and_couples_refs() -> None: control = _ControlPlane() snapshots = E2BHomeSnapshotBackend( control_plane=control, # pyright: ignore[reportArgumentType] ) bindings = E2BExecutionBindingBackend( control_plane=control, # pyright: ignore[reportArgumentType] template="prepared-template", active_timeout_seconds=3600, ) default_allocation = await bindings.create_binding( ExecutionBindingCreateSpec( tenant_id="tenant-1", agent_id="agent-1", binding_id="binding-1", workspace_id="workspace-1", existing_workspace_ref=None, home_snapshot_ref=None, ) ) snapshot_allocation = await bindings.create_binding( ExecutionBindingCreateSpec( tenant_id="tenant-1", agent_id="agent-1", binding_id="binding-2", workspace_id="workspace-2", existing_workspace_ref=None, home_snapshot_ref="snapshot-1", ) ) assert control.created == [("prepared-template", "pause"), ("snapshot-1", "pause")] assert default_allocation.binding_ref == default_allocation.workspace_ref assert snapshot_allocation.binding_ref == snapshot_allocation.workspace_ref runtime = control.sandboxes[default_allocation.binding_ref] assert runtime.files.paths == {"/home/dify/workspace"} assert runtime.pauses == [True] for allocation in (default_allocation, snapshot_allocation): await bindings.destroy_binding( ExecutionBindingDestroySpec( binding_ref=allocation.binding_ref, workspace_ref=allocation.workspace_ref, destroy_workspace=True, ) ) await snapshots.delete("snapshot-1") assert control.killed == [default_allocation.binding_ref, snapshot_allocation.binding_ref] assert control.deleted_snapshots == ["snapshot-1"] @pytest.mark.anyio async def test_e2b_rejects_shared_workspace_and_binding_only_destroy() -> None: control = _ControlPlane() backend = E2BExecutionBindingBackend( control_plane=control, # pyright: ignore[reportArgumentType] template="prepared-template", active_timeout_seconds=3600, ) spec = ExecutionBindingCreateSpec( tenant_id="tenant-1", agent_id="agent-2", binding_id="binding-2", workspace_id="workspace-1", existing_workspace_ref="sandbox-1", home_snapshot_ref="snapshot-1", ) with pytest.raises(SharedWorkspaceUnsupportedError): await backend.create_binding(spec) assert control.created == [] assert control.sandboxes == {} with pytest.raises(WorkspacePreservationUnsupportedError): await backend.destroy_binding(ExecutionBindingDestroySpec(binding_ref="sandbox-1", destroy_workspace=False)) @pytest.mark.anyio async def test_e2b_binding_create_kills_sandbox_when_initialization_fails() -> None: control = _ControlPlane(pause_error=RuntimeError("pause failed")) backend = E2BExecutionBindingBackend( control_plane=control, # pyright: ignore[reportArgumentType] template="prepared-template", active_timeout_seconds=3600, ) with pytest.raises(BindingCreateError, match="pause failed"): await backend.create_binding( ExecutionBindingCreateSpec( tenant_id="tenant-1", agent_id="agent-1", binding_id="binding-1", workspace_id="workspace-1", existing_workspace_ref=None, home_snapshot_ref="snapshot-1", ) ) sandbox = next(iter(control.sandboxes.values())) assert sandbox.killed == 1 @pytest.mark.anyio async def test_e2b_missing_explicit_snapshot_does_not_fall_back_to_template() -> None: class _FailingControlPlane(_ControlPlane): async def create(self, template: str, *, timeout: int, metadata: dict[str, str], on_timeout: str) -> _Sandbox: del timeout, metadata self.created.append((template, on_timeout)) raise RuntimeError("snapshot unavailable") control = _FailingControlPlane() backend = E2BExecutionBindingBackend( control_plane=control, # pyright: ignore[reportArgumentType] template="prepared-template", active_timeout_seconds=3600, ) with pytest.raises(BindingCreateError, match="snapshot unavailable"): await backend.create_binding( ExecutionBindingCreateSpec( tenant_id="tenant-1", agent_id="agent-1", binding_id="binding-1", workspace_id="workspace-1", existing_workspace_ref=None, home_snapshot_ref="missing-snapshot", ) ) assert control.created == [("missing-snapshot", "pause")] @pytest.mark.anyio async def test_e2b_checkpoint_uses_exact_source_runtime() -> None: control = _ControlPlane() source_sandbox = _Sandbox(sandbox_id="source") source = E2BRuntimeLease( sandbox=source_sandbox, data_plane=cast(ShellctlRuntimeLease, object()), ) backend = E2BHomeSnapshotBackend( control_plane=control, # pyright: ignore[reportArgumentType] ) snapshot_ref = await backend.create_from_runtime( spec=HomeSnapshotCreateSpec(tenant_id="tenant-1", agent_id="agent-1", home_snapshot_id="home-2"), source=source, ) assert snapshot_ref == "snapshot-source-1" assert source_sandbox.snapshots == 1