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
24 changes: 21 additions & 3 deletions src/agent_env/providers/sandbox_providers/modal_vm_sandbox.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@
import re
import shlex
import time
from typing import Awaitable, Callable, ClassVar, Optional
from typing import TYPE_CHECKING, AsyncIterator, Awaitable, Callable, ClassVar, Optional

import modal

Expand All @@ -37,9 +37,12 @@
)
from agent_env.attribution import Attribution
from agent_env.config import get_config
from agent_env.providers.sandbox_providers.sandbox import NetworkPolicy, VmSandbox
from agent_env.providers.sandbox_providers.sandbox import NetworkPolicy, VmSandbox, push_object_over_stdin
from agent_env.providers.sandbox_providers.sandbox_provider import SANDBOX_MODE_VM, SandboxProvider

if TYPE_CHECKING:
from agent_env.store.object_store import ObjectStore

logger = logging.getLogger(__name__)

# Ubuntu 22.04 + Docker engine + the compose v2 plugin baked in — the Modal-native
Expand Down Expand Up @@ -136,6 +139,21 @@ async def exec(self, *command: str):
process = await self._sb.exec.aio(*cmd, text=False)
return _ModalProcessAdapter(process)

async def _exec_with_stdin(self, script: str, stdin: AsyncIterator[bytes]) -> tuple[int, str, str]:
process = await self._sb.exec.aio("bash", "-c", script, text=False)
async for piece in stdin:
process.stdin.write(piece)
await process.stdin.drain.aio()
process.stdin.write_eof()
await process.stdin.drain.aio()
exit_code = await process.wait.aio()
stdout, stderr = await asyncio.gather(process.stdout.read.aio(), process.stderr.read.aio())
return exit_code, stdout.decode(), stderr.decode()

async def _write_unsigned_object(self, object_store: ObjectStore, object_url: str, vm_path: str) -> None:
"""Over stdin, which on Modal carries an object several times as fast as exec arguments do."""
await push_object_over_stdin(self, object_store, object_url, vm_path)

async def setup_vm_for_gateway(self, exposed_ports: Optional[list[int]] = None) -> None:
"""Make the VM ready for the gateway deploy.

Expand Down Expand Up @@ -193,7 +211,7 @@ async def _ensure_dockerd_running(self) -> None:
)

async def _download_object_to_vm(self, object_url: str, vm_path: str) -> None:
"""Presigned objects go through aria2c (see _DL_*); a store that cannot presign keeps the base path."""
"""Presigned objects go through aria2c (see _DL_*); one a store cannot presign goes over stdin."""
store = get_config().get_object_store()
signed = await asyncio.to_thread(store.signed_get_url, object_url)
signed_at = time.monotonic()
Expand Down
215 changes: 212 additions & 3 deletions src/agent_env/providers/sandbox_providers/sandbox.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,18 +4,22 @@

import asyncio
import base64
import contextlib
import hashlib
import io
import logging
import os
import posixpath
import shlex
import stat
import tempfile
import time
import uuid
import weakref
from abc import ABC, abstractmethod
from dataclasses import dataclass, replace
from enum import Enum
from typing import TYPE_CHECKING, Any, Iterable, Optional
from typing import IO, TYPE_CHECKING, Any, AsyncIterator, Callable, Iterable, Optional

from agent_env.config import get_config
from agent_env.utils.paths import validate_relative_filename
Expand Down Expand Up @@ -55,6 +59,13 @@
# provider serves each exec a fraction of its bandwidth, so many ranges are in flight.
_READ_RANGE_BYTES = 4 * 1024 * 1024
_READS_IN_FLIGHT = 16
# An object pushed onto a VM host over exec arguments goes a chunk per exec, the chunk's base64 in the script; over stdin
# it goes as segments of at least _MIN_SEGMENT_BYTES, each in pieces. An object of one chunk or one segment goes in a
# single exec. Chunks and segments start at multiples of _PUSH_BLOCK, a multiple of the 4096-byte block dd seeks in and
# of 3, so a piece's base64 is unpadded.
_PUSH_BLOCK = 3 * 4096
_STDIN_PIECE_BYTES = 16 * _PUSH_BLOCK
_MIN_SEGMENT_BYTES = 1024 * 1024 // _PUSH_BLOCK * _PUSH_BLOCK


class NetworkPolicyUnsupportedError(NotImplementedError):
Expand Down Expand Up @@ -329,8 +340,14 @@ async def _download_object_to_vm(self, object_url: str, vm_path: str) -> None:
await self._write_unsigned_object(object_store, object_url, vm_path)

async def _write_unsigned_object(self, object_store: ObjectStore, object_url: str, vm_path: str) -> None:
"""Place an object the store cannot sign a URL for: its bytes, streamed over exec."""
await self._write_bytes_to_vm_path(await asyncio.to_thread(object_store.get, object_url), vm_path)
"""Place an object the store cannot sign a URL for: pushed over exec, a chunk per exec."""
await push_object_over_exec(self, object_store, object_url, vm_path)

async def _exec_with_stdin(self, script: str, stdin: AsyncIterator[bytes]) -> tuple[int, str, str]:
"""Run ``script`` with bash on the VM host, its standard input the bytes ``stdin`` yields, and return (exit_code,
stdout, stderr). A sandbox whose exec takes stdin implements it, and pushes objects with
``push_object_over_stdin``."""
raise NotImplementedError(f"{self.__class__.__name__} does not take stdin on exec")

async def _remove_vm_temp_file(self, *vm_paths: str) -> None:
try:
Expand Down Expand Up @@ -380,6 +397,9 @@ async def write_file_from_url(self, url: str, destination_path: str) -> None:
# One exec_script is a single `bash -c <script>` arg, capped by Linux MAX_ARG_STRLEN
# (128 KiB); 96 KiB leaves room for the printf wrapper.
_WFT_CHUNK_BYTES = 96 * 1024
# How many execs carry one object pushed onto the host at once, chunks or segments. A provider whose execs scale
# differently sets its own.
_PUSHES_IN_FLIGHT = 8

async def _write_bytes_to_vm_path(self, data: bytes, vm_path: str) -> None:
"""Stream bytes from agent-env onto the VM host at vm_path (base64 over exec)."""
Expand Down Expand Up @@ -488,3 +508,192 @@ async def upload_vm_file(sandbox: VmSandbox, vm_path: str, store: ObjectStore, o
local_path = os.path.join(directory, posixpath.basename(vm_path))
await read_vm_file(sandbox, vm_path, local_path)
await asyncio.to_thread(store.put_file_at, object_url, local_path)


async def push_object_over_exec(sandbox: VmSandbox, store: ObjectStore, object_url: str, vm_path: str) -> None:
"""Write the object at ``object_url`` to ``vm_path`` on the VM host a chunk per exec, the chunk's base64 in the
script and written at its offset, ``sandbox._PUSHES_IN_FLIGHT`` execs at once; an object of one chunk goes in a
single exec. The file is checked against the object's sha256."""
chunk = _exec_chunk_bytes(sandbox)
quoted = shlex.quote(vm_path)
digest = hashlib.sha256()
with contextlib.closing(await asyncio.to_thread(store.open, object_url)) as source:
data = await asyncio.to_thread(_read_exactly, source, chunk)
digest.update(data)
if len(data) < chunk:
written = await _write_in_one_exec(sandbox, quoted, data)
else:
await sandbox.exec_script(f": > {quoted}")
await _push_chunks(sandbox, quoted, source, chunk, data, digest)
written = (await sandbox.exec_script(f"sha256sum {quoted}")).split()[:1]
_check_written(vm_path, object_url, written, digest.hexdigest())


def _exec_chunk_bytes(sandbox: VmSandbox) -> int:
"""How much of an object one exec's script carries, as base64, within the sandbox's command limit."""
return max(_PUSH_BLOCK, sandbox._WFT_CHUNK_BYTES // 4 * 3 // _PUSH_BLOCK * _PUSH_BLOCK)


async def _write_in_one_exec(sandbox: VmSandbox, quoted: str, data: bytes) -> list[str]:
"""Write ``data`` as the whole VM file in one exec whose script carries it, and return the sha256 it reports."""
encoded = shlex.quote(base64.b64encode(data).decode())
reported = await sandbox.exec_script(f"printf %s {encoded} | base64 -d > {quoted} && sha256sum {quoted}",
max_retries=2)
return reported.split()[:1]


def _check_written(vm_path: str, object_url: str, written: list[str], expected: str) -> None:
if written != [expected]:
raise RuntimeError(f"{vm_path} on the VM doesn't match {object_url}: its sha256 is {written}, the object's "
f"{expected}")


async def _push_chunks(
sandbox: VmSandbox, quoted: str, source: IO[bytes], chunk: int, data: bytes, digest: Any,
) -> None:
"""Write ``data``, the object's first chunk, and every chunk ``source`` has after it, each at its offset."""
gate = asyncio.Semaphore(sandbox._PUSHES_IN_FLIGHT)
failed = False

async def write(offset: int, data: bytes) -> None:
nonlocal failed
try:
encoded = shlex.quote(base64.b64encode(data).decode())
await sandbox.exec_script(
f"printf %s {encoded} | base64 -d | dd of={quoted} bs=4096 seek={offset // 4096} conv=notrunc",
max_retries=2,
)
except BaseException:
failed = True
raise
finally:
gate.release()

writes: list[asyncio.Future] = []
try:
offset = 0
while not failed and data:
await gate.acquire()
if failed:
gate.release()
break
writes.append(asyncio.ensure_future(write(offset, data)))
offset += len(data)
data = await asyncio.to_thread(_read_exactly, source, chunk)
digest.update(data)
await asyncio.gather(*writes)
finally:
for pending in writes:
pending.cancel()
await asyncio.gather(*writes, return_exceptions=True)


async def push_object_over_stdin(sandbox: VmSandbox, store: ObjectStore, object_url: str, vm_path: str) -> None:
"""Write the object at ``object_url`` to ``vm_path`` on the VM host over exec stdin. When the store opens it as a
file, it goes as up to ``sandbox._PUSHES_IN_FLIGHT`` segments at once, all read from that one open file, so one
version of it; one that fits an exec's script goes in it, since a stdin exec takes more round trips. An object of
one segment, or any other reader, goes in a single stdin exec. Each segment is written at its offset and checked
against its sha256 by the exec that writes it."""
quoted = shlex.quote(vm_path)
with contextlib.closing(await asyncio.to_thread(store.open, object_url)) as source:
descriptor = _file_descriptor(source)
if descriptor is None:
await _push_segment(sandbox, quoted, lambda size: _read_exactly(source, size), 0, None, whole_file=True)
return
size = os.fstat(descriptor).st_size
if size < _exec_chunk_bytes(sandbox):
data = await asyncio.to_thread(_reader_at(descriptor, 0), size)
if len(data) != size:
raise RuntimeError(f"{object_url} ended {size - len(data)} bytes short; it changed during the push")
written = await _write_in_one_exec(sandbox, quoted, data)
_check_written(vm_path, object_url, written, hashlib.sha256(data).hexdigest())
return
segment = max(_MIN_SEGMENT_BYTES, -(-size // (sandbox._PUSHES_IN_FLIGHT * _PUSH_BLOCK)) * _PUSH_BLOCK)
if size <= segment:
await _push_segment(sandbox, quoted, _reader_at(descriptor, 0), 0, size, whole_file=True)
return
await sandbox.exec_script(f": > {quoted}")
pushes = [
asyncio.ensure_future(_push_segment(
sandbox, quoted, _reader_at(descriptor, offset), offset, min(segment, size - offset), whole_file=False))
for offset in range(0, size, segment)
]
try:
await asyncio.gather(*pushes)
finally:
for pending in pushes:
pending.cancel()
await asyncio.gather(*pushes, return_exceptions=True)


async def _push_segment(
sandbox: VmSandbox, quoted: str, read: Callable[[int], bytes], offset: int, length: int | None, *, whole_file: bool,
) -> None:
"""Stream ``length`` bytes from ``read`` (all it gives when None) into the VM file at ``offset``, and check them.
A ``whole_file`` segment creates the file, which holds nothing else."""
digest = hashlib.sha256()

async def pieces() -> AsyncIterator[bytes]:
left = length
while left is None or left > 0:
data = await asyncio.to_thread(read, _STDIN_PIECE_BYTES if left is None else min(_STDIN_PIECE_BYTES, left))
if not data:
if left is not None:
raise RuntimeError(f"The object ended {left} bytes short of the segment at offset {offset} of "

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Failed push leaves writer waiting

When a local object is cut short during a Modal push, pieces() raises while _exec_with_stdin is feeding a running writer. The error skips write_eof() and wait(), leaving that VM process waiting for input after the push has failed. Close the input and clean up the process when feeding fails.

Prompt To Fix With AI
This is a comment left during a code review.
Path: src/agent_env/providers/sandbox_providers/sandbox.py
Line: 595

Comment:
**Failed push leaves writer waiting**

When a local object is cut short during a Modal push, `pieces()` raises while `_exec_with_stdin` is feeding a running writer. The error skips `write_eof()` and `wait()`, leaving that VM process waiting for input after the push has failed. Close the input and clean up the process when feeding fails.

---

For each issue above, determine whether it is valid and should be fixed. If so, fix it directly.

Fix in Cursor Fix in Claude Code Fix in Codex

f"{quoted}; it changed during the push")
return
Comment thread
greptile-apps[bot] marked this conversation as resolved.
digest.update(data)
if left is not None:
left -= len(data)
yield base64.b64encode(data)

if whole_file:
script = f"base64 -d > {quoted} && sha256sum {quoted}"
else:
script = (f"base64 -d | dd of={quoted} bs=4096 seek={offset // 4096} conv=notrunc && "
f"tail -c +{offset + 1} {quoted} | head -c {length} | sha256sum")
code, stdout, stderr = await sandbox._exec_with_stdin(script, pieces())
if code != 0:
raise RuntimeError(f"Writing {quoted} at offset {offset} on the VM failed (exit {code}): {stderr[-1500:]}")
if stdout.split()[:1] != [digest.hexdigest()]:
raise RuntimeError(f"{quoted} on the VM doesn't match at offset {offset}: its sha256 is {stdout.split()[:1]}, "
f"the object's {digest.hexdigest()}")


def _file_descriptor(reader: IO[bytes]) -> int | None:
"""The descriptor of the regular file ``reader`` reads as it is, or None when it reads anything else. A reader that
transforms what it reads, such as a decompressing one, can still name its file's descriptor, so only a plain file
reader counts."""
if not isinstance(reader, (io.BufferedReader, io.FileIO)):
return None
try:
descriptor = reader.fileno()
except OSError:
return None
return descriptor if stat.S_ISREG(os.fstat(descriptor).st_mode) else None


def _reader_at(descriptor: int, offset: int) -> Callable[[int], bytes]:
"""Reads the file ``descriptor`` names from ``offset`` on, by position, so readers of one file don't share a
file offset."""
position = offset

def read(size: int) -> bytes:
nonlocal position
parts, got = [], 0
while got < size and (part := os.pread(descriptor, size - got, position + got)):
parts.append(part)
got += len(part)
position += got
return b"".join(parts)

return read


def _read_exactly(reader: IO[bytes], size: int) -> bytes:
"""``size`` bytes of ``reader``, fewer only at its end: a stream may return less than asked for."""
parts, got = [], 0
while got < size and (part := reader.read(size - got)):
parts.append(part)
got += len(part)
return b"".join(parts)
37 changes: 32 additions & 5 deletions tst/unit/providers/sandbox_providers/modal_vm_sandbox_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -304,21 +304,48 @@ async def test_download_attempts_exhausted_cleans_up_and_raises(store, monkeypat


@pytest.mark.asyncio
async def test_unsigned_store_keeps_the_base_class_path():
async def test_an_object_the_store_cannot_sign_goes_over_stdin(monkeypatch):
class _Local:
def signed_get_url(self, object_url, expires_in=3600):
return None

def get(self, object_url):
return b"BYTES"
pushed = []

set_object_store(_Local())
async def push(sandbox, store, object_url, vm_path):
pushed.append((sandbox, store, object_url, vm_path))

monkeypatch.setattr(mvs, "push_object_over_stdin", push)
store = _Local()
set_object_store(store)
try:
vm = _ScriptedVm([])
await vm.load_s3_file("file:///store/x", "/app/x")
finally:
reset_config()
assert not vm.downloads() and any("base64 -d" in s for s in vm.scripts)
assert pushed == [(vm, store, "file:///store/x", "/app/x")]
assert not vm.downloads()


@pytest.mark.asyncio
async def test_a_stdin_exec_feeds_the_script_and_returns_what_it_printed():
process = MagicMock()
process.stdin.drain.aio = AsyncMock()
process.wait.aio = AsyncMock(return_value=0)
process.stdout.read.aio = AsyncMock(return_value=b"abc123 -\n")
process.stderr.read.aio = AsyncMock(return_value=b"2+0 records in\n")
sb = MagicMock()
sb.exec.aio = AsyncMock(return_value=process)

async def pieces():
yield b"QUJD"
yield b"REVG"

result = await _sandbox(sb)._exec_with_stdin("base64 -d | dd of=/tmp/x && sha256sum /tmp/x", pieces())

assert result == (0, "abc123 -\n", "2+0 records in\n")
sb.exec.aio.assert_awaited_once_with("bash", "-c", "base64 -d | dd of=/tmp/x && sha256sum /tmp/x", text=False)
assert [c.args for c in process.stdin.write.call_args_list] == [(b"QUJD",), (b"REVG",)]
process.stdin.write_eof.assert_called_once()


@pytest.mark.asyncio
Expand Down
Loading
Loading