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
26 changes: 25 additions & 1 deletion src/agent_env/task_step/task_steps/load_artifact.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from __future__ import annotations

import asyncio
import contextlib
import logging
import os
import posixpath
Expand Down Expand Up @@ -356,6 +357,21 @@ async def _load_url_onto_vm(self, sandbox, url: str, destination_path: str) -> N
f"curl -fsSL {CURL_RETRY_FLAGS} {shlex.quote(url)} -o {shlex.quote(destination_path)}"
)

async def _load_url_into_container(self, sandbox, container: str, url: str, destination_path: str) -> None:
"""Download ``url`` onto the VM host, then copy it into ``container`` at ``destination_path``:
``write_file_from_url`` reaches only the sandbox's own agent container."""
vm_temp = f"/tmp/_load_url_{uuid.uuid4().hex[:8]}"
Comment thread
greptile-apps[bot] marked this conversation as resolved.
try:
await self._load_url_onto_vm(sandbox, url, vm_temp)
await sandbox.exec_script(
f"docker exec -u 0 {shlex.quote(container)} mkdir -p {shlex.quote(posixpath.dirname(destination_path))}"
)
await sandbox.docker_cp(vm_temp, f"{container}:{destination_path}", remove_source=True)
except BaseException: # a cancelled load leaves nothing on the host either
with contextlib.suppress(Exception): # best effort: the load's own error is the one to report
await sandbox.exec_script(f"rm -f {shlex.quote(vm_temp)}")
raise

async def execute(self, context: TaskStepContext) -> TaskStepContext:
from agent_env.a2a_agent import A2AAgent
from agent_env.a2a_agent.store import get_a2a_agent_instance_store
Expand Down Expand Up @@ -599,14 +615,22 @@ async def execute(self, context: TaskStepContext) -> TaskStepContext:
destination = (destination_path or "/tmp/file_artifacts").rstrip("/") or "/"
semaphore = asyncio.Semaphore(8)

if onto_vm_host:
target_desc = f"VM sandbox '{self.sandbox_name}'"
elif self.container_name is not None:
target_desc = f"container '{self.container_name}'"
else:
target_desc = "agent"

async def _load_one(url: str, filename: str) -> None:
async with semaphore:
dest = f"{destination}/{filename}"
if onto_vm_host:
await self._load_url_onto_vm(sandbox, url, dest)
elif self.container_name is not None:
await self._load_url_into_container(sandbox, sandbox.scoped_name(self.container_name), url, dest)
else:
await sandbox.write_file_from_url(url, dest)
target_desc = f"VM sandbox '{self.sandbox_name}'" if onto_vm_host else "agent"
logger.info(f"Loaded URL into {target_desc}: {url} -> {dest}")

await asyncio.gather(*(_load_one(u, f) for u, f in downloads))
Expand Down
39 changes: 39 additions & 0 deletions tst/integration/task_step/test_run_docker_container_local.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,11 @@
"""

import asyncio
import functools
import http.server
import shutil
import subprocess
import threading
import uuid

import httpx
Expand All @@ -18,6 +21,7 @@
from agent_env.task.teardown import teardown_run
from agent_env.task_step.context import TaskStepContext
from agent_env.task_step.task_steps.deploy_sandbox import DeploySandboxTaskStep
from agent_env.task_step.task_steps.load_artifact import LoadArtifactTaskStep
from agent_env.task_step.task_steps.run_docker_container import RunDockerContainerTaskStep
from tst.util.capabilities import missing_capability_reason

Expand Down Expand Up @@ -122,3 +126,38 @@ async def run(context: TaskStepContext) -> None:
for sandbox_id in sandbox_ids:
assert not _docker("ps", "-aq", "--filter", f"label=agentenv.sandbox={sandbox_id}")
assert not _docker("network", "ls", "-q", "--filter", f"name=^task-net-{sandbox_id}$")


@pytest.mark.asyncio
async def test_urls_load_into_a_run_docker_container_container(local_backends):
"""A URL given with container_name lands in that container, not the sandbox's agent container."""
suffix = uuid.uuid4().hex[:8]
(local_backends / "Dockerfile").write_text("FROM mirror.gcr.io/library/nginx:1.27-bookworm\n")
(local_backends / "data.csv").write_text("a,b\n1,2\n")
server = http.server.ThreadingHTTPServer(
("127.0.0.1", 0), functools.partial(http.server.SimpleHTTPRequestHandler, directory=str(local_backends)))
threading.Thread(target=server.serve_forever, daemon=True).start()
build_context = FileArtifactUniverse.put(id=f"rdc3-ctx-{suffix}", file_artifacts={
"Dockerfile": FileArtifact.put(id=f"rdc3-df-{suffix}", description="Dockerfile",
file_path=str(local_backends / "Dockerfile")),
})
context = TaskStepContext(instance_id=f"rdc3-{suffix}")
try:
await DeploySandboxTaskStep(
id="box", version=None, sandbox_name="box", sandbox_mode="vm", sandbox_type="local",
).execute(context)
await RunDockerContainerTaskStep(
id="ctr", version=None, sandbox_name="box", docker_context_artifact_id=build_context.id,
container_name="worker",
).execute(context)
await LoadArtifactTaskStep(
id="load", version=None, sandbox_name="box", container_name="worker", destination_path="/work",
urls=[f"http://127.0.0.1:{server.server_address[1]}/data.csv"],
).execute(context)
container = f"worker-{context.deployed_sandboxes[0].sandbox_id}"
assert _docker("exec", container, "cat", "/work/data.csv") == "a,b\n1,2"
finally:
server.shutdown()
report = await teardown_run(context)

assert not report.still_up
43 changes: 43 additions & 0 deletions tst/unit/task_step/test_load_artifact_vm_target.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@

from __future__ import annotations

import asyncio

import pytest

from agent_env.artifact.artifact import Artifact
Expand Down Expand Up @@ -54,6 +56,9 @@ async def load_url_file(self, url, destination_path): # pragma: no cover
async def write_file_from_url(self, url, destination_path): # pragma: no cover
calls.append(("container_url", url, destination_path))

async def docker_cp(self, source, destination, *, remove_source=False):
calls.append(("cp", source, destination, remove_source))

from agent_env.providers.sandbox_providers import sandbox_provider as sp_mod

async def _get_sandbox(sandbox_id):
Expand Down Expand Up @@ -209,6 +214,44 @@ async def _fake_into_container(sandbox, container_name, uni, destination):
assert seen == [{"container": "task-container", "destination": "/loaded"}]
assert ctx.metadata["loaded_file_artifact_universes"][0]["container_name"] == "task-container"

@pytest.mark.asyncio
async def test_urls_go_into_the_named_container_not_the_agents(self, vm):
ctx, calls = vm
ctx.metadata["deployed_docker_containers"] = [{"container_name": "task-container", "sandbox_name": "mk"}]
step = LoadArtifactTaskStep(
id="stage", version=None, sandbox_name="mk", container_name="task-container",
urls=["https://example.com/data.csv"], destination_path="/work",
)

await step.execute(ctx)

[curl] = [c[1] for c in calls if c[0] == "exec" and c[1].startswith("curl ")]
assert ("cp", curl.rsplit(" -o ", 1)[1], "task-container:/work/data.csv", True) in calls
assert ("exec", "docker exec -u 0 task-container mkdir -p /work") in calls
assert not [c for c in calls if c[0] == "container_url"]

@pytest.mark.asyncio
@pytest.mark.parametrize("cleanup_fails", [False, True])
@pytest.mark.parametrize("failure", [RuntimeError("curl: (22) 404"), asyncio.CancelledError()], ids=["error", "cancel"])
async def test_a_failed_url_load_reports_its_own_error_and_cleans_up(self, cleanup_fails, failure):
scripts = []

class _Vm:
async def exec_script(self, script, **kw):
scripts.append(script)
if script.startswith("curl "):
raise failure
if script.startswith("rm -f ") and cleanup_fails:
raise RuntimeError("exec transport closed")
return ""

step = LoadArtifactTaskStep(id="s", version=None, sandbox_name="mk", container_name="c", urls=["https://x/y"])

with pytest.raises(type(failure)):
await step._load_url_into_container(_Vm(), "c", "https://x/y", "/work/y")

assert scripts[-1].startswith("rm -f /tmp/_load_url_")


class TestUrlHelper:
@pytest.mark.asyncio
Expand Down
Loading