Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
51 changes: 38 additions & 13 deletions src/agents/sandbox/sandboxes/unix_local.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,7 @@
UnsafeTarMemberError,
safe_extract_tarfile,
should_skip_tar_member,
validate_tarfile,
)
from ..workspace_paths import _raise_if_filesystem_root
from . import _unix_local_file_ops
Expand Down Expand Up @@ -1148,19 +1149,17 @@ async def persist_workspace(self) -> io.IOBase:

def _archive_workspace() -> None:
with tarfile.open(fileobj=buf, mode="w") as tar:
tar.add(
root,
arcname=".",
filter=lambda ti: (
None
if should_skip_tar_member(
ti.name,
skip_rel_paths=skip,
root_name=None,
)
else ti
),
)

def filter_member(member: tarfile.TarInfo) -> tarfile.TarInfo | None:
# tarfile records inodes before filtering. Clear even excluded entries so
# every retained hardlink has its own payload. Unlike dereference=True,
# this preserves symlinks instead of reading their targets on the host.
getattr(tar, "inodes").clear() # noqa: B009 - Not exposed by typeshed.
if should_skip_tar_member(member.name, skip_rel_paths=skip, root_name=None):
return None
return member

tar.add(root, arcname=".", filter=filter_member)

try:
await run_blocking_workspace_io(_archive_workspace)
Expand All @@ -1170,6 +1169,32 @@ def _archive_workspace() -> None:
buf.seek(0)
return buf

async def _restore_snapshot_into_workspace_on_resume(self) -> None:
root = Path(self.state.manifest.root)
archive = await self.state.snapshot.restore(dependencies=self.dependencies)

def validate_archive() -> None:
try:
with tarfile.open(fileobj=archive, mode="r:*") as tar:
validate_tarfile(tar, allow_external_symlink_targets=False)
archive.seek(0)
except UnsafeTarMemberError as e:
raise WorkspaceArchiveWriteError(
path=root, context={"reason": e.reason, "member": e.member}, cause=e
) from e
except (tarfile.TarError, OSError) as e:
raise WorkspaceArchiveWriteError(path=root, cause=e) from e

try:
# Older snapshots may contain unsupported members. Reject them before discarding
# the live files; keep hydrate_workspace's own validation for direct callers too.
await run_blocking_workspace_io(validate_archive)
await self._clear_workspace_root_on_resume()
await self.hydrate_workspace(archive)
finally:
with suppress(Exception):
archive.close()

async def hydrate_workspace(self, data: io.IOBase) -> None:
root = Path(self.state.manifest.root)

Expand Down
156 changes: 153 additions & 3 deletions tests/sandbox/test_unix_local.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import asyncio
import contextlib
import io
import os
import shutil
import signal
import tarfile
Expand All @@ -14,8 +15,8 @@

import pytest

from agents.sandbox import SandboxPathGrant
from agents.sandbox.errors import PtySessionNotFoundError
from agents.sandbox import LocalSnapshotSpec, SandboxPathGrant
from agents.sandbox.errors import PtySessionNotFoundError, WorkspaceArchiveWriteError
from agents.sandbox.manifest import Environment, Manifest
from agents.sandbox.sandboxes import unix_local as unix_local_module
from agents.sandbox.sandboxes.unix_local import (
Expand All @@ -24,7 +25,7 @@
UnixLocalSandboxSessionState,
_UnixPtyProcessEntry,
)
from agents.sandbox.snapshot import NoopSnapshot
from agents.sandbox.snapshot import LocalSnapshot, NoopSnapshot
from agents.sandbox.types import ExecResult, User


Expand All @@ -48,6 +49,155 @@ async def _exec_internal(
return ExecResult(stdout=b"", stderr=b"", exit_code=0)


@pytest.mark.asyncio
@pytest.mark.parametrize("exclude_first", [False, True])
async def test_unix_local_snapshot_round_trips_hardlinks(
tmp_path: Path, exclude_first: bool
) -> None:
workspace = tmp_path / "workspace"
client = UnixLocalSandboxClient(inherit_host_environment=False)
session = await client.create(
manifest=Manifest(root=str(workspace)),
snapshot=LocalSnapshotSpec(base_path=tmp_path / "snapshots"),
)
await session.start()
first = workspace / "a.py"
second = workspace / "b.py"
first.write_bytes(b"VALUE = 1\n")
first.chmod(0o755)
os.link(first, second)
(workspace / "link.py").symlink_to("b.py")
(workspace / "copy.py").write_bytes(b"independent\n")
if exclude_first:
session.register_persist_workspace_skip_path("a.py")
await session.stop()
archive = await session.state.snapshot.restore()
try:
with tarfile.open(fileobj=archive) as tar:
members = {member.name: member for member in tar.getmembers()}
assert members["./b.py"].isreg()
assert members["./link.py"].issym()
if exclude_first:
assert "./a.py" not in members
else:
assert members["./a.py"].isreg()
finally:
archive.close()

# Prove that resume actually restores the snapshot, not the surviving workspace.
second.write_bytes(b"changed after snapshot\n")
(workspace / "stale.txt").write_bytes(b"remove on resume")
resumed = await client.resume(session.state)
try:
await resumed.start()
assert second.read_bytes() == b"VALUE = 1\n"
assert second.stat().st_mode & 0o777 == 0o755
assert (workspace / "copy.py").read_bytes() == b"independent\n"
assert (workspace / "link.py").is_symlink()
assert (workspace / "link.py").read_bytes() == b"VALUE = 1\n"
assert not (workspace / "stale.txt").exists()
if exclude_first:
assert not first.exists()
else:
assert first.read_bytes() == b"VALUE = 1\n"
assert first.stat().st_ino != second.stat().st_ino
finally:
await resumed.shutdown()
await session.shutdown()


@pytest.mark.asyncio
@pytest.mark.parametrize("invalid_kind", ["hardlink", "external_symlink", "invalid_tar"])
async def test_unix_local_resume_rejects_invalid_snapshot_before_clearing_workspace(
tmp_path: Path, invalid_kind: str
) -> None:
workspace = tmp_path / "workspace"
client = UnixLocalSandboxClient(inherit_host_environment=False)
session = await client.create(
manifest=Manifest(root=str(workspace)),
snapshot=LocalSnapshotSpec(base_path=tmp_path / "snapshots"),
)
await session.start()
(workspace / "keep.txt").write_bytes(b"live workspace")
archive = io.BytesIO()
if invalid_kind == "invalid_tar":
archive.write(b"not a tar archive")
else:
with tarfile.open(fileobj=archive, mode="w") as tar:
member = tarfile.TarInfo("link")
member.type = tarfile.LNKTYPE if invalid_kind == "hardlink" else tarfile.SYMTYPE
member.linkname = "keep.txt" if invalid_kind == "hardlink" else "../outside"
tar.addfile(member)
archive.seek(0)
await session.state.snapshot.persist(archive)
archive.close()

resumed = await client.resume(session.state)
try:
with pytest.raises(WorkspaceArchiveWriteError):
await resumed.start()
assert (workspace / "keep.txt").read_bytes() == b"live workspace"
assert sorted(path.name for path in workspace.iterdir()) == ["keep.txt"]
assert not await resumed.running()
finally:
await resumed.shutdown()
await session.shutdown()


@pytest.mark.asyncio
async def test_unix_local_resume_cancellation_waits_for_archive_validation(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
workspace = tmp_path / "workspace"
client = UnixLocalSandboxClient(inherit_host_environment=False)
session = await client.create(
manifest=Manifest(root=str(workspace)),
snapshot=LocalSnapshotSpec(base_path=tmp_path / "snapshots"),
)
await session.start()
(workspace / "keep.txt").write_bytes(b"live workspace")
await session.stop()
archive = await session.state.snapshot.restore()
started = threading.Event()
release = threading.Event()
events: list[str] = []
validate = unix_local_module.validate_tarfile

async def restore(self: LocalSnapshot, **kwargs: object) -> io.IOBase:
return archive

def slow_validate(tar: tarfile.TarFile, **kwargs: object) -> None:
started.set()
assert release.wait(timeout=5)
validate(tar, allow_external_symlink_targets=False)
events.append("validated")

monkeypatch.setattr(LocalSnapshot, "restore", restore)
monkeypatch.setattr(unix_local_module, "validate_tarfile", slow_validate)
resumed = await client.resume(session.state)
task = asyncio.create_task(resumed.start())
try:
while not started.is_set():
if task.done():
await task
Comment thread
jbeckwith-oai marked this conversation as resolved.
pytest.fail("resume did not validate the archive")
await asyncio.sleep(0.005)
task.cancel()
await asyncio.sleep(0)
assert not task.done()
assert not archive.closed
release.set()
with pytest.raises(asyncio.CancelledError):
await task
Comment thread
jbeckwith-oai marked this conversation as resolved.
assert events == ["validated"]
assert archive.closed
assert (workspace / "keep.txt").read_bytes() == b"live workspace"
finally:
release.set()
await resumed.shutdown()
await session.shutdown()


@pytest.mark.asyncio
async def test_unix_local_inherits_host_environment_by_default(
tmp_path: Path,
Expand Down
Loading