diff --git a/tartarus/background.py b/tartarus/background.py new file mode 100644 index 0000000..e1e3334 --- /dev/null +++ b/tartarus/background.py @@ -0,0 +1,241 @@ +"""BackgroundRegistry: track detached capability runs (PLAN.md §6.9). + +A capability declared `kind = "background"` is launched detached by the jail and +handed here. The registry assigns it a short id (`bg-1`, `bg-2`, …), tails its +combined output, and reports liveness to the control-plane capabilities +(`bg_status` / `bg_output` / `bg_stop`). + +This object bridges two worlds: the **sync broker** registers a task from a +worker thread (`Broker.handle` runs under `asyncio.to_thread`), while the +**async loop** owns the completion notifications. `register` is therefore +thread-safe and schedules the per-task monitor onto the captured event loop; +`status`/`output` refresh liveness directly so they work even with no loop +attached (e.g. unit tests that poll). +""" + +import asyncio +import os +import signal +import threading +from dataclasses import dataclass +from datetime import UTC, datetime + +from tartarus.audit import AuditEvent, AuditSink, NullAuditLog +from tartarus.jail import BackgroundHandle, ExecResult +from tartarus.models import ToolResult + +DEFAULT_OUTPUT_TRUNCATE_CHARS = 10_000 +MAX_UTF8_BYTES_PER_CHAR = 4 + + +class BackgroundError(Exception): + """Raised when a control op references an unknown task.""" + + +@dataclass(frozen=True) +class Notice: + """A completion notice the loop turns into a transcript message.""" + + task_id: str + capability: str + exit_code: int + output_tail: str + network_summary: str | None = None + + +@dataclass +class _Task: + task_id: str + capability: str + handle: BackgroundHandle + started_at: str + status: str = "running" # "running" | "exited" + exit_code: int | None = None + monitor: asyncio.Task | None = None + _finalized: bool = False + + +class BackgroundRegistry: + def __init__( + self, + notices: "asyncio.Queue[Notice] | None" = None, + loop: asyncio.AbstractEventLoop | None = None, + output_truncate: int = DEFAULT_OUTPUT_TRUNCATE_CHARS, + audit: AuditSink | None = None, + ): + self._tasks: dict[str, _Task] = {} + self._counter = 0 + self._notices = notices + self._loop = loop + self._output_truncate = output_truncate + self._audit = audit if audit is not None else NullAuditLog() + self._lock = threading.RLock() + + def register(self, capability: str, handle: BackgroundHandle) -> str: + """Track a freshly launched task and return its id. Thread-safe.""" + with self._lock: + self._counter += 1 + bg_id = f"bg-{self._counter}" + task = _Task( + task_id=bg_id, + capability=capability, + handle=handle, + started_at=datetime.now(UTC).isoformat(), + ) + self._tasks[bg_id] = task + if self._loop is not None: + # register runs in the broker's worker thread; hop to the loop thread. + self._loop.call_soon_threadsafe(self._start_monitor, task) + return bg_id + + # --- control-plane operations (called by the broker) -------------------- + + def status(self, task_id: str | None) -> str: + task = self._get(task_id) + self._refresh(task) + if task.status == "running": + return ( + f"{task.task_id} ({task.capability}): running since {task.started_at}" + ) + return ( + f"{task.task_id} ({task.capability}): exited with code {task.exit_code} " + f"(started {task.started_at})" + ) + + def output(self, task_id: str | None, offset: int = 0) -> str: + task = self._get(task_id) + self._refresh(task) + text = self._read_log(task, offset) + if not text: + return "(no output)" + return text + + def stop(self, task_id: str | None) -> str: + task = self._get(task_id) + self._refresh(task) + if task.status == "exited": + return f"{task.task_id} already exited with code {task.exit_code}" + self._kill(task) + return f"{task.task_id} sent SIGTERM" + + # --- teardown ----------------------------------------------------------- + + def shutdown_all(self) -> None: + """Kill every running task, stop its proxy, cancel its monitor.""" + with self._lock: + tasks = list(self._tasks.values()) + for task in tasks: + if task.status == "running": + self._kill(task) + if task.handle.proxy is not None: + task.handle.proxy.stop() + if task.monitor is not None: + task.monitor.cancel() + + @property + def has_running(self) -> bool: + return any(self._is_running(task) for task in self._tasks.values()) + + # --- internals ---------------------------------------------------------- + + def _start_monitor(self, task: _Task) -> None: + assert self._loop is not None + task.monitor = self._loop.create_task(self._monitor(task)) + + async def _monitor(self, task: _Task) -> None: + code = await asyncio.to_thread(task.handle.proc.wait) + summary = self._finalize(task, code) + if self._notices is not None: + await self._notices.put( + Notice( + task_id=task.task_id, + capability=task.capability, + exit_code=code, + output_tail=self._read_log(task, 0), + network_summary=summary, + ) + ) + + def _finalize(self, task: _Task, code: int) -> str | None: + """Record the exit once and tear down the task's proxy. Idempotent.""" + with self._lock: + if task._finalized: + return None + task._finalized = True + task.status = "exited" + task.exit_code = code + + summary = None + if task.handle.proxy is not None: + summary = task.handle.proxy.summary() + task.handle.proxy.stop() + + self._audit_completion(task, code, summary) + return summary + + def _audit_completion(self, task: _Task, code: int, summary: str | None) -> None: + """Append the second audit record for a background task: its completion. + + The launch was already audited by the broker; this closes the loop with + the exit code and output size, reusing the same JSONL sink. + """ + tail = self._read_log(task, 0) + self._audit.record( + AuditEvent( + call_id=task.task_id, + tool_name=task.capability, + arguments={}, + command="(background completion)", + exec_result=ExecResult(code, tail, "", network_summary=summary), + result=ToolResult(task.task_id, tail, is_error=code != 0), + ) + ) + + def _refresh(self, task: _Task) -> None: + """Pick up an exit even when no async monitor is running (poll path).""" + if task.status != "running": + return + code = task.handle.proc.poll() + if code is not None: + self._finalize(task, code) + + def _is_running(self, task: _Task) -> bool: + self._refresh(task) + return task.status == "running" + + def _kill(self, task: _Task) -> None: + try: + os.killpg(task.handle.pgid, signal.SIGTERM) + except (ProcessLookupError, PermissionError): + pass + + def _read_log(self, task: _Task, offset: int) -> str: + try: + with open(task.handle.log_path, "rb") as log_file: + size = log_file.seek(0, os.SEEK_END) + # Read only the final window that can survive truncation, not the + # whole (possibly huge) log. output_truncate counts characters but + # the log is bytes, so budget the worst case: a UTF-8 character is + # at most MAX_UTF8_BYTES_PER_CHAR bytes. A seek may split a leading + # character; decode(errors="replace") absorbs it and the tail slice + # below discards it. + tail_byte_budget = self._output_truncate * MAX_UTF8_BYTES_PER_CHAR + window_start = max(offset, size - tail_byte_budget) + log_file.seek(window_start) + # Bound the read explicitly: a concurrent writer may have grown + # the file past the sampled size, and read() alone would follow + # it to the new EOF, defeating the memory bound. + data = log_file.read(tail_byte_budget) + except OSError: + return "" + text = data.decode("utf-8", "replace") + if window_start > offset or len(text) > self._output_truncate: + # Keep the tail: for a long-running task the recent output matters most. + return "...(truncated)\n" + text[-self._output_truncate :] + return text + + def _get(self, task_id: str | None) -> _Task: + with self._lock: + if task_id not in self._tasks: + raise BackgroundError(f"unknown background task {task_id!r}") + return self._tasks[task_id] diff --git a/tests/test_background.py b/tests/test_background.py new file mode 100644 index 0000000..6c6b601 --- /dev/null +++ b/tests/test_background.py @@ -0,0 +1,183 @@ +"""BackgroundRegistry tests. + +These launch real short-lived `sh` subprocesses directly (no bwrap), so they +exercise the registry's tracking, log tailing, signalling, and async monitor in +isolation from the jail. The jailed launch path is covered by test_jail.py. +""" + +import asyncio +import os +import subprocess +import threading + +import pytest + +from tartarus.background import BackgroundError, BackgroundRegistry +from tartarus.jail import BackgroundHandle + + +def _launch(args: list[str], log_path: str) -> BackgroundHandle: + log_file = open(log_path, "wb") + proc = subprocess.Popen( + args, stdout=log_file, stderr=subprocess.STDOUT, start_new_session=True + ) + log_file.close() + return BackgroundHandle(proc=proc, pgid=os.getpgid(proc.pid), log_path=log_path) + + +class MemorySink: + def __init__(self): + self.records = [] + + def record(self, event) -> None: + self.records.append(event) + + +def test_register_assigns_sequential_ids(tmp_path): + registry = BackgroundRegistry() + one = _launch(["sh", "-c", "true"], str(tmp_path / "1.log")) + two = _launch(["sh", "-c", "true"], str(tmp_path / "2.log")) + + assert registry.register("cap", one) == "bg-1" + assert registry.register("cap", two) == "bg-2" + + +def test_status_reflects_exit_without_a_loop(tmp_path): + # No event loop/monitor attached: status must refresh liveness itself. + registry = BackgroundRegistry() + handle = _launch(["sh", "-c", "sleep 0.2; exit 0"], str(tmp_path / "s.log")) + registry.register("cap", handle) + + assert "running" in registry.status("bg-1") + handle.proc.wait(timeout=5) + assert "exited with code 0" in registry.status("bg-1") + + +def test_output_reads_log_from_offset(tmp_path): + registry = BackgroundRegistry() + handle = _launch(["sh", "-c", "printf abcdef"], str(tmp_path / "o.log")) + handle.proc.wait(timeout=5) + registry.register("cap", handle) + + assert registry.output("bg-1") == "abcdef" + assert registry.output("bg-1", 3) == "def" + + +def test_output_truncates_logs_larger_than_the_tail_window(tmp_path): + # A log far larger than the budget must come back as only its final window + # (with the truncation marker), exercising the seek-to-tail read path. + registry = BackgroundRegistry(output_truncate=10) + payload = "HEAD" + "x" * 200 + "TAIL" + handle = _launch(["sh", "-c", f"printf '%s' {payload}"], str(tmp_path / "big.log")) + handle.proc.wait(timeout=5) + registry.register("cap", handle) + + output = registry.output("bg-1") + + assert output.startswith("...(truncated)\n") + assert output.endswith("xxxxxxTAIL") # the last output_truncate characters + assert "HEAD" not in output + + +def test_stop_kills_a_running_task(tmp_path): + registry = BackgroundRegistry() + handle = _launch(["sh", "-c", "sleep 30"], str(tmp_path / "k.log")) + registry.register("cap", handle) + + assert "SIGTERM" in registry.stop("bg-1") + handle.proc.wait(timeout=5) + assert "exited" in registry.status("bg-1") + + +def test_shutdown_all_reaps_running_tasks(tmp_path): + registry = BackgroundRegistry() + handle = _launch(["sh", "-c", "sleep 30"], str(tmp_path / "r.log")) + registry.register("cap", handle) + + registry.shutdown_all() + + handle.proc.wait(timeout=5) + assert handle.proc.returncode is not None + + +def test_unknown_task_raises(): + registry = BackgroundRegistry() + with pytest.raises(BackgroundError, match="unknown background task"): + registry.status("bg-404") + + +def test_completion_is_audited(tmp_path): + sink = MemorySink() + registry = BackgroundRegistry(audit=sink) + handle = _launch(["sh", "-c", "printf hi; exit 2"], str(tmp_path / "a.log")) + handle.proc.wait(timeout=5) + registry.register("cap", handle) + + registry.status("bg-1") # refresh → finalize → one completion record + + assert len(sink.records) == 1 + assert sink.records[0].tool_name == "cap" + assert sink.records[0].exec_result.code == 2 + # Idempotent: refreshing again does not double-record. + registry.status("bg-1") + assert len(sink.records) == 1 + + +def test_monitor_enqueues_notice_on_exit(tmp_path): + async def run(): + notices: asyncio.Queue = asyncio.Queue() + registry = BackgroundRegistry(notices=notices, loop=asyncio.get_running_loop()) + handle = _launch(["sh", "-c", "printf hello; exit 3"], str(tmp_path / "n.log")) + bg_id = registry.register("cap", handle) + notice = await asyncio.wait_for(notices.get(), timeout=5) + return bg_id, notice + + bg_id, notice = asyncio.run(run()) + + assert bg_id == "bg-1" + assert notice.task_id == "bg-1" + assert notice.capability == "cap" + assert notice.exit_code == 3 + assert "hello" in notice.output_tail + + +def test_register_is_thread_safe(tmp_path): + registry = BackgroundRegistry() + handles = [ + _launch(["sh", "-c", "true"], str(tmp_path / f"{i}.log")) for i in range(50) + ] + ids: list[str] = [] + errors: list[Exception] = [] + + def register_one(handle): + try: + ids.append(registry.register("cap", handle)) + except Exception as exc: # pragma: no cover - failures should fail the test + errors.append(exc) + + threads = [threading.Thread(target=register_one, args=(h,)) for h in handles] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + assert not errors + assert len(ids) == 50 + assert len(set(ids)) == 50 + assert sorted(ids, key=lambda value: int(value.split("-", 1)[1])) == [ + f"bg-{i}" for i in range(1, 51) + ] + + +def test_finalize_is_idempotent_under_lock(tmp_path): + sink = MemorySink() + registry = BackgroundRegistry(audit=sink) + handle = _launch(["sh", "-c", "printf hi; exit 2"], str(tmp_path / "a.log")) + handle.proc.wait(timeout=5) + registry.register("cap", handle) + task = registry._get("bg-1") + + # The first call finalizes; the second must not record again or stop proxy twice. + assert registry._finalize(task, 2) is None + assert registry._finalize(task, 2) is None + assert len(sink.records) == 1