92 lines
4.3 KiB
Python
92 lines
4.3 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import io
|
|
import os
|
|
import tarfile
|
|
from dataclasses import replace
|
|
|
|
import pytest
|
|
|
|
import docker
|
|
from agent_platform.gateway.provider import DockerExecutionProvider, workspace_ref
|
|
|
|
|
|
@pytest.mark.skipif(os.getenv("RUN_DOCKER_INTEGRATION") != "1", reason="set RUN_DOCKER_INTEGRATION=1")
|
|
async def test_real_docker_workspaces_are_isolated(settings) -> None:
|
|
client = docker.from_env()
|
|
settings = replace(settings, workspace_network_enabled=True)
|
|
provider = DockerExecutionProvider(settings, client=client)
|
|
users = ("integration-user-a", "integration-user-b")
|
|
try:
|
|
assert (await provider.write_file(users[0], "secret.txt", "only-a")).ok
|
|
assert (await provider.write_file(users[1], "secret.txt", "only-b")).ok
|
|
first = await provider.read_file(users[0], "secret.txt", 1, 10)
|
|
second = await provider.read_file(users[1], "secret.txt", 1, 10)
|
|
assert "only-a" in first.output and "only-b" not in first.output
|
|
assert "only-b" in second.output and "only-a" not in second.output
|
|
listing = await provider.browse_files(users[0], ".")
|
|
assert {entry.name for entry in listing.entries} >= {"secret.txt"}
|
|
download = await provider.download_file(users[0], "secret.txt")
|
|
assert download.content == b"only-a"
|
|
archive = await provider.archive_files(users[0], ".")
|
|
archive_bytes = b"".join(archive.chunks or ())
|
|
with tarfile.open(fileobj=io.BytesIO(archive_bytes), mode="r:gz") as workspace_tar:
|
|
assert any(name.endswith("secret.txt") for name in workspace_tar.getnames())
|
|
escaped = await provider.exec(users[0], "test ! -e /var/run/docker.sock", ".", 10)
|
|
assert escaped.ok
|
|
|
|
git_init = "git init -q && git config user.email test@example.invalid && git config user.name Integration"
|
|
assert (await provider.exec(users[0], git_init, ".", 10)).ok
|
|
assert (await provider.write_file(users[0], "tracked.txt", "before\n")).ok
|
|
assert (await provider.exec(users[0], "git add tracked.txt && git commit -qm initial", ".", 10)).ok
|
|
patch = (
|
|
"diff --git a/tracked.txt b/tracked.txt\n"
|
|
"index 90be1a7..2a140c2 100644\n"
|
|
"--- a/tracked.txt\n"
|
|
"+++ b/tracked.txt\n"
|
|
"@@ -1 +1 @@\n"
|
|
"-before\n"
|
|
"+after\n"
|
|
)
|
|
assert (await provider.apply_patch(users[0], patch, ".")).ok
|
|
assert "after" in (await provider.read_file(users[0], "tracked.txt", 1, 10)).output
|
|
|
|
started = await provider.start_process(users[0], "sleep 0.2; echo background-done", ".")
|
|
process_id = started.metadata["process_id"]
|
|
polled = await provider.poll_process(users[0], process_id)
|
|
for _ in range(20):
|
|
if not polled.metadata.get("running"):
|
|
break
|
|
await asyncio.sleep(0.1)
|
|
polled = await provider.poll_process(users[0], process_id)
|
|
assert not polled.metadata.get("running")
|
|
assert polled.exit_code == 0
|
|
assert "background-done" in polled.output
|
|
|
|
container = client.containers.get(workspace_ref(users[0]).container_name)
|
|
container.reload()
|
|
second_container = client.containers.get(workspace_ref(users[1]).container_name)
|
|
second_container.reload()
|
|
assert container.attrs["Config"]["User"] == "1000:1000"
|
|
assert container.attrs["HostConfig"]["ReadonlyRootfs"] is True
|
|
assert container.attrs["HostConfig"]["CapDrop"] == ["ALL"]
|
|
assert container.attrs["HostConfig"]["PidsLimit"] == settings.workspace_pids_limit
|
|
assert set(container.attrs["NetworkSettings"]["Networks"]) == {workspace_ref(users[0]).network_name}
|
|
assert set(second_container.attrs["NetworkSettings"]["Networks"]) == {workspace_ref(users[1]).network_name}
|
|
finally:
|
|
for user in users:
|
|
ref = workspace_ref(user)
|
|
try:
|
|
client.containers.get(ref.container_name).remove(force=True)
|
|
except docker.errors.NotFound:
|
|
pass
|
|
try:
|
|
client.volumes.get(ref.volume_name).remove(force=True)
|
|
except docker.errors.NotFound:
|
|
pass
|
|
try:
|
|
client.networks.get(ref.network_name).remove()
|
|
except docker.errors.NotFound:
|
|
pass
|