from __future__ import annotations import asyncio import hashlib import io import json import mimetypes import shlex import tarfile import uuid import zlib from collections.abc import Iterator from dataclasses import dataclass from pathlib import PurePosixPath from typing import Protocol from docker.errors import ImageNotFound, NotFound import docker from agent_platform.config import Settings from agent_platform.gateway.schemas import ToolResult, WorkspaceListing, WorkspaceStatus from agent_platform.text import truncate_middle def normalize_workspace_path(raw_path: str) -> PurePosixPath: if "\x00" in raw_path: raise ValueError("Path contains a null byte") raw = raw_path.strip() or "." path = PurePosixPath(raw) if path.is_absolute(): try: path = path.relative_to("/workspace") except ValueError as exc: raise ValueError("Absolute paths must be under /workspace") from exc if any(part in {"..", ""} for part in path.parts): raise ValueError("Path traversal is not allowed") return PurePosixPath(".") if str(path) in {"", "."} else path def absolute_workspace_path(raw_path: str) -> str: relative = normalize_workspace_path(raw_path) return "/workspace" if relative == PurePosixPath(".") else f"/workspace/{relative}" @dataclass(frozen=True, slots=True) class WorkspaceRef: workspace_id: str container_name: str volume_name: str network_name: str @dataclass(slots=True) class WorkspaceDownload: filename: str media_type: str content: bytes | None = None chunks: Iterator[bytes] | None = None def workspace_ref(user_id: str) -> WorkspaceRef: digest = hashlib.sha256(user_id.encode("utf-8")).hexdigest()[:20] return WorkspaceRef( workspace_id=digest, container_name=f"k1412-ws-{digest}", volume_name=f"k1412-ws-data-{digest}", network_name=f"k1412-ws-net-{digest}", ) class ExecutionProvider(Protocol): provider_name: str async def status(self, user_id: str) -> WorkspaceStatus: ... async def exec(self, user_id: str, command: str, cwd: str, timeout_seconds: int) -> ToolResult: ... async def list_files(self, user_id: str, path: str, max_depth: int, limit: int) -> ToolResult: ... async def read_file(self, user_id: str, path: str, start_line: int, max_lines: int) -> ToolResult: ... async def write_file(self, user_id: str, path: str, content: str) -> ToolResult: ... async def search_files( self, user_id: str, query: str, path: str, glob: str | None, limit: int, ) -> ToolResult: ... async def apply_patch(self, user_id: str, patch: str, cwd: str) -> ToolResult: ... async def start_process(self, user_id: str, command: str, cwd: str) -> ToolResult: ... async def poll_process(self, user_id: str, process_id: str) -> ToolResult: ... async def cancel_process(self, user_id: str, process_id: str) -> ToolResult: ... async def browse_files(self, user_id: str, path: str) -> WorkspaceListing: ... async def download_file(self, user_id: str, path: str) -> WorkspaceDownload: ... async def archive_files(self, user_id: str, path: str) -> WorkspaceDownload: ... class DockerExecutionProvider: provider_name = "local-docker" def __init__(self, settings: Settings, client: docker.DockerClient | None = None) -> None: self.settings = settings self.client = client or docker.from_env() self._locks: dict[str, asyncio.Lock] = {} self._lock_guard = asyncio.Lock() async def _user_lock(self, user_id: str) -> asyncio.Lock: async with self._lock_guard: return self._locks.setdefault(user_id, asyncio.Lock()) async def _container(self, user_id: str): ref = workspace_ref(user_id) def ensure(): try: container = self.client.containers.get(ref.container_name) except NotFound: try: self.client.images.get(self.settings.workspace_image) except ImageNotFound as exc: raise RuntimeError(f"Workspace image {self.settings.workspace_image!r} is not installed") from exc if self.settings.workspace_network_enabled: try: self.client.networks.get(ref.network_name) except NotFound: self.client.networks.create( ref.network_name, driver="bridge", check_duplicate=True, labels={ "app.k1412.component": "user-workspace-network", "app.k1412.workspace-id": ref.workspace_id, }, ) container = self.client.containers.create( image=self.settings.workspace_image, name=ref.container_name, hostname=f"workspace-{ref.workspace_id}", command=["sleep", "infinity"], user="1000:1000", working_dir="/workspace", volumes={ref.volume_name: {"bind": "/workspace", "mode": "rw"}}, labels={ "app.k1412.component": "user-workspace", "app.k1412.workspace-id": ref.workspace_id, }, detach=True, read_only=True, tmpfs={ "/tmp": "rw,noexec,nosuid,size=256m", # noqa: S108 - isolated container tmpfs "/run": "rw,noexec,nosuid,size=16m", }, cap_drop=["ALL"], security_opt=["no-new-privileges:true"], mem_limit=self.settings.workspace_memory_limit, nano_cpus=int(self.settings.workspace_cpu_limit * 1_000_000_000), pids_limit=self.settings.workspace_pids_limit, network_mode=ref.network_name if self.settings.workspace_network_enabled else "none", ) container.reload() if container.status != "running": container.start() container.reload() return container lock = await self._user_lock(user_id) async with lock: return await asyncio.to_thread(ensure) async def status(self, user_id: str) -> WorkspaceStatus: container = await self._container(user_id) ref = workspace_ref(user_id) return WorkspaceStatus( workspace_id=ref.workspace_id, provider=self.provider_name, container_name=ref.container_name, state=container.status, ) async def _exec_argv( self, user_id: str, argv: list[str], *, cwd: str = ".", max_output: int = 128_000, ) -> ToolResult: container = await self._container(user_id) workdir = absolute_workspace_path(cwd) def execute() -> ToolResult: result = container.exec_run( argv, workdir=workdir, user="1000:1000", demux=True, ) stdout, stderr = result.output if isinstance(result.output, tuple) else (result.output, b"") output = ((stdout or b"") + (stderr or b"")).decode("utf-8", errors="replace") truncated = len(output) > max_output if truncated: output = truncate_middle(output, max_output, "\n… output truncated; middle omitted …\n") return ToolResult( ok=result.exit_code == 0, output=output, exit_code=result.exit_code, truncated=truncated, ) return await asyncio.to_thread(execute) async def exec(self, user_id: str, command: str, cwd: str, timeout_seconds: int) -> ToolResult: normalize_workspace_path(cwd) bounded_timeout = max(1, min(timeout_seconds, 1800)) shell = f"timeout --signal=TERM {bounded_timeout}s sh -lc {shlex.quote(command)}" return await self._exec_argv(user_id, ["sh", "-lc", shell], cwd=cwd) async def list_files(self, user_id: str, path: str, max_depth: int, limit: int) -> ToolResult: absolute = absolute_workspace_path(path) script = ( "import os,sys\n" "root=sys.argv[1]; max_depth=int(sys.argv[2]); limit=int(sys.argv[3]); out=[]\n" "base_depth=root.rstrip('/').count('/')\n" "for current, dirs, files in os.walk(root):\n" " depth=current.rstrip('/').count('/')-base_depth\n" " dirs[:]=sorted(d for d in dirs if d not in {'.git','node_modules','.venv','__pycache__'})\n" " if depth>=max_depth: dirs[:]=[]\n" " rel=os.path.relpath(current,'/workspace')\n" " for name in sorted(dirs): out.append((os.path.join(rel,name) if rel!='.' else name)+'/')\n" " for name in sorted(files): out.append(os.path.join(rel,name) if rel!='.' else name)\n" " if len(out)>=limit: break\n" "print('\\n'.join(out[:limit]))\n" ) return await self._exec_argv( user_id, ["python3", "-c", script, absolute, str(max_depth), str(limit)], ) async def browse_files(self, user_id: str, path: str) -> WorkspaceListing: relative = normalize_workspace_path(path) absolute = absolute_workspace_path(path) script = ( "import json,os,pathlib,sys\n" "root=pathlib.Path(sys.argv[1])\n" "if not root.is_dir(): raise SystemExit(f'Not a directory: {root}')\n" "items=[]\n" "for item in root.iterdir():\n" " if item.name=='.agent': continue\n" " try: info=item.lstat()\n" " except OSError: continue\n" " kind='symlink' if item.is_symlink() else ('directory' if item.is_dir() else 'file')\n" " rel=item.relative_to('/workspace').as_posix()\n" " items.append({'name':item.name,'path':rel,'kind':kind,'size':info.st_size," "'modified_at':int(info.st_mtime)})\n" "items.sort(key=lambda value:(value['kind']!='directory',value['name'].casefold()))\n" "print(json.dumps(items,ensure_ascii=False))\n" ) raw = await self._exec_argv(user_id, ["python3", "-c", script, absolute]) if not raw.ok: raise FileNotFoundError(raw.output.strip() or f"Directory not found: {relative}") try: entries = json.loads(raw.output) except json.JSONDecodeError as exc: raise RuntimeError("Workspace returned an invalid directory listing") from exc return WorkspaceListing(path="." if relative == PurePosixPath(".") else str(relative), entries=entries) async def download_file(self, user_id: str, path: str) -> WorkspaceDownload: relative = normalize_workspace_path(path) if relative == PurePosixPath("."): raise ValueError("A file path is required") container = await self._container(user_id) absolute = absolute_workspace_path(path) def collect() -> WorkspaceDownload: chunks, metadata = container.get_archive(absolute) size = int(metadata.get("size", 0)) if size > self.settings.workspace_download_max_bytes: raise ValueError( f"File exceeds the {self.settings.workspace_download_max_bytes // (1024 * 1024)} MB download limit" ) archive = io.BytesIO() archive_limit = self.settings.workspace_download_max_bytes + 2 * 1024 * 1024 for chunk in chunks: archive.write(chunk) if archive.tell() > archive_limit: raise ValueError("File archive exceeded the download limit") archive.seek(0) with tarfile.open(fileobj=archive, mode="r:*") as tar: members = tar.getmembers() member = members[0] if members else None if member is None: raise FileNotFoundError(f"File not found: {relative}") if not member.isfile(): raise IsADirectoryError(f"Not a regular file: {relative}") source = tar.extractfile(member) if source is None: raise FileNotFoundError(f"File not found: {relative}") content = source.read(self.settings.workspace_download_max_bytes + 1) if len(content) > self.settings.workspace_download_max_bytes: raise ValueError("File exceeds the download limit") media_type = mimetypes.guess_type(relative.name)[0] or "application/octet-stream" return WorkspaceDownload(filename=relative.name, media_type=media_type, content=content) return await asyncio.to_thread(collect) async def archive_files(self, user_id: str, path: str) -> WorkspaceDownload: relative = normalize_workspace_path(path) container = await self._container(user_id) chunks, _ = await asyncio.to_thread(container.get_archive, absolute_workspace_path(path)) def compressed_chunks() -> Iterator[bytes]: compressor = zlib.compressobj(level=6, method=zlib.DEFLATED, wbits=31) for chunk in chunks: compressed = compressor.compress(chunk) if compressed: yield compressed final = compressor.flush() if final: yield final basename = "workspace" if relative == PurePosixPath(".") else relative.name return WorkspaceDownload( filename=f"{basename}.tar.gz", media_type="application/gzip", chunks=compressed_chunks(), ) async def read_file(self, user_id: str, path: str, start_line: int, max_lines: int) -> ToolResult: absolute = absolute_workspace_path(path) script = ( "import pathlib,sys\n" "p=pathlib.Path(sys.argv[1]); start=int(sys.argv[2]); count=int(sys.argv[3])\n" "if not p.is_file(): raise SystemExit(f'Not a file: {p}')\n" "with p.open('r',encoding='utf-8',errors='replace') as f:\n" " lines=f.readlines()\n" "for i,line in enumerate(lines[start-1:start-1+count],start): print(f'{i:>6} {line}',end='')\n" ) return await self._exec_argv( user_id, ["python3", "-c", script, absolute, str(start_line), str(max_lines)], ) async def write_file(self, user_id: str, path: str, content: str) -> ToolResult: relative = normalize_workspace_path(path) if relative == PurePosixPath("."): raise ValueError("A file path is required") parent = str(relative.parent) container = await self._container(user_id) if parent not in {"", "."}: await self._exec_argv(user_id, ["mkdir", "-p", absolute_workspace_path(parent)]) archive = io.BytesIO() encoded = content.encode("utf-8") with tarfile.open(fileobj=archive, mode="w") as tar: info = tarfile.TarInfo(name=relative.name) info.size = len(encoded) info.mode = 0o644 info.uid = 1000 info.gid = 1000 tar.addfile(info, io.BytesIO(encoded)) archive.seek(0) destination = "/workspace" if parent in {"", "."} else absolute_workspace_path(parent) await asyncio.to_thread(container.put_archive, destination, archive.getvalue()) return ToolResult(ok=True, output=f"Wrote {len(encoded)} bytes to {relative}") async def search_files( self, user_id: str, query: str, path: str, glob: str | None, limit: int, ) -> ToolResult: absolute = absolute_workspace_path(path) argv = ["rg", "--line-number", "--color=never", "--max-count", str(limit), "--", query, absolute] if glob: argv[1:1] = ["--glob", glob] result = await self._exec_argv(user_id, argv) if result.exit_code == 1: return ToolResult(ok=True, output="No matches.", exit_code=0) return result async def apply_patch(self, user_id: str, patch: str, cwd: str) -> ToolResult: process_id = uuid.uuid4().hex patch_path = f".agent/tmp/{process_id}.patch" await self.write_file(user_id, patch_path, patch) result = await self.exec( user_id, f"git apply --whitespace=nowarn {shlex.quote(absolute_workspace_path(patch_path))}", cwd, 120, ) await self._exec_argv(user_id, ["rm", "-f", absolute_workspace_path(patch_path)]) return result async def start_process(self, user_id: str, command: str, cwd: str) -> ToolResult: normalize_workspace_path(cwd) process_id = uuid.uuid4().hex process_dir = absolute_workspace_path(f".agent/processes/{process_id}") worker = ( f"sh -lc {shlex.quote(command)}; " "code=$?; " f'echo "$code" > {shlex.quote(process_dir + "/exit_code")}; ' 'exit "$code"' ) wrapped = ( f"mkdir -p {shlex.quote(process_dir)}; " f"setsid sh -lc {shlex.quote(worker)} " f"> {shlex.quote(process_dir + '/output.log')} 2>&1 & " f'pid=$!; echo "$pid" > {shlex.quote(process_dir + "/pid")}; ' f"echo {shlex.quote(json.dumps({'process_id': process_id}))}" ) result = await self._exec_argv(user_id, ["sh", "-lc", wrapped], cwd=cwd) result.metadata = {"process_id": process_id} return result async def poll_process(self, user_id: str, process_id: str) -> ToolResult: process_dir = absolute_workspace_path(f".agent/processes/{process_id}") script = ( "import json,os,pathlib,sys\n" "d=pathlib.Path(sys.argv[1]); pid=(d/'pid').read_text().strip() if (d/'pid').exists() else ''\n" "code=(d/'exit_code').read_text().strip() if (d/'exit_code').exists() else None\n" "log=(d/'output.log').read_text(errors='replace')[-50000:] if (d/'output.log').exists() else ''\n" "print(json.dumps({'running': code is None and bool(pid), 'pid': pid, 'exit_code': code, 'output': log}))\n" ) raw = await self._exec_argv(user_id, ["python3", "-c", script, process_dir]) if not raw.ok: return raw try: data = json.loads(raw.output) except json.JSONDecodeError: return ToolResult(ok=False, output="Invalid process state", exit_code=1) exit_code = int(data["exit_code"]) if data["exit_code"] is not None else None return ToolResult( ok=exit_code in {None, 0}, output=data["output"], exit_code=exit_code, metadata={"process_id": process_id, "running": data["running"], "pid": data["pid"]}, ) async def cancel_process(self, user_id: str, process_id: str) -> ToolResult: process_dir = absolute_workspace_path(f".agent/processes/{process_id}") command = ( f"pid=$(cat {shlex.quote(process_dir + '/pid')} 2>/dev/null) || exit 1; " 'kill -TERM -- "-$pid" 2>/dev/null || kill -TERM "$pid" 2>/dev/null || true; ' f"echo 143 > {shlex.quote(process_dir + '/exit_code')}" ) return await self._exec_argv(user_id, ["sh", "-lc", command]) class SSHDockerExecutionProvider(DockerExecutionProvider): provider_name = "ssh-docker" def __init__(self, settings: Settings) -> None: client = docker.DockerClient( base_url=f"ssh://{settings.workspace_ssh_host}", use_ssh_client=True, ) super().__init__(settings, client=client) def create_execution_provider(settings: Settings) -> ExecutionProvider: if settings.execution_provider == "ssh-docker": return SSHDockerExecutionProvider(settings) return DockerExecutionProvider(settings)