diff --git a/routes/shell_routes.py b/routes/shell_routes.py index d63e80ee6..900ac4125 100644 --- a/routes/shell_routes.py +++ b/routes/shell_routes.py @@ -1,6 +1,7 @@ """Shell routes — user-facing command execution endpoint.""" import asyncio +import contextlib import importlib import json import logging @@ -8,6 +9,7 @@ import os import re import shlex import shutil +import signal import subprocess import uuid import tempfile @@ -49,6 +51,7 @@ from core.platform_compat import ( detached_popen_kwargs, find_bash, git_bash_path, + pid_alive, ) @@ -558,6 +561,13 @@ STREAM_TIMEOUT = 120 # default for short commands MAX_OUTPUT = 200_000 # truncate limit TMUX_LOG_DIR = Path(tempfile.gettempdir()) / "odysseus-tmux" PTY_UNSUPPORTED_ERROR = "pty_unsupported" +# PTY teardown. The PTY child leads its own session (os.setsid), so killing it +# has to signal the whole process group and then confirm the group is gone — +# see _terminate_pty_session. +PTY_KILL_ESCALATION = (signal.SIGTERM, signal.SIGKILL) +PTY_KILL_GRACE = 1.0 # seconds a signalled session gets to exit +PTY_KILL_POLL_INTERVAL = 0.05 # re-check interval while waiting for it +PTY_KILL_FAILED_HINT = "; processes it started survived the kill and are still running" class ShellExecRequest(BaseModel): @@ -662,6 +672,124 @@ async def _exec_shell(command: str, timeout: int = EXEC_TIMEOUT) -> Dict[str, An return {"stdout": "", "stderr": str(e), "exit_code": -1} +def _session_pgid(pid: int) -> int | None: + """Process-group id of the session ``pid`` leads, or None if unavailable. + + Read this *before* the leader is reaped: once it is, ``getpgid`` fails and + the group id can no longer be recovered from the pid. + """ + getpgid = getattr(os, "getpgid", None) + if getpgid is None: # no process groups (native Windows) + return None + try: + return getpgid(pid) + except OSError: + return None + + +def _signal_session(pgid: int | None, pid: int, sig: int) -> bool: + """Send ``sig`` to the whole process group, or to the lone process. + + Returns whether anything was signalled, so a caller can tell "the session + is already gone" from "the signal landed". The single-pid fallback + matters: if ``setsid`` did not take effect, or the platform has no process + groups, teardown must still reach the child rather than do nothing. + """ + killpg = getattr(os, "killpg", None) + if pgid is not None and killpg is not None: + try: + killpg(pgid, sig) + return True + except ProcessLookupError: + return False + except OSError: + pass # group signalling refused — fall through to the single pid + try: + os.kill(pid, sig) + return True + except OSError: + return False + + +def _session_alive(pgid: int | None, pid: int) -> bool: + """True while any member of the process group still exists. + + An unreaped zombie is still signallable, so a True here can also mean the + leader has exited but not yet been collected. Without a group id this can + only speak for the child itself, not for anything it spawned. + """ + killpg = getattr(os, "killpg", None) + if pgid is not None and killpg is not None: + try: + killpg(pgid, 0) + return True + except OSError: + return False + return pid_alive(pid) + + +async def _await_session_exit(proc, pgid: int | None, pid: int) -> bool: + """Wait up to the grace period for the leader and its group to go away.""" + loop = asyncio.get_running_loop() + deadline = loop.time() + PTY_KILL_GRACE + while True: + remaining = deadline - loop.time() + if proc.returncode is None and remaining > 0: + # Reap the leader, otherwise its own zombie keeps the group alive + # and the liveness probe below can never come back clean. + with contextlib.suppress(asyncio.TimeoutError): + await asyncio.wait_for(proc.wait(), remaining) + if not _session_alive(pgid, pid): + return True + if loop.time() >= deadline: + return False + await asyncio.sleep(PTY_KILL_POLL_INTERVAL) + + +async def _terminate_pty_session(proc) -> bool: + """Kill the PTY child and every process in the session it leads. + + The child is spawned under ``os.setsid``, so it leads its own session and + process group. Signalling only the leader is strictly worse than never + calling ``setsid`` at all: the descendants are detached from the server's + group too, so nothing will ever reach them, while the caller reports the + command as terminated. Signal the group instead, escalate to SIGKILL if it + outlives the grace period, and return whether the session is actually gone + so the caller can say so rather than assume it. + """ + pid = getattr(proc, "pid", None) + if pid is None: + return True + pgid = _session_pgid(pid) + + for sig in PTY_KILL_ESCALATION: + if not _signal_session(pgid, pid, sig): + break # nothing left to signal + if await _await_session_exit(proc, pgid, pid): + return True + return not _session_alive(pgid, pid) + + +async def _terminate_pty_session_quietly(proc) -> None: + """Best-effort :func:`_terminate_pty_session` for paths with no reader. + + The client-disconnect and exception paths have nowhere left to report a + containment failure to, so they log it instead of raising over the top of + whatever is already going wrong. + """ + pid = getattr(proc, "pid", None) + try: + contained = await _terminate_pty_session(proc) + except Exception: + logger.exception("PTY session teardown failed for pid %s", pid) + return + if not contained: + logger.warning( + "PTY session for pid %s survived teardown; it may still be running", + pid, + ) + + async def _generate_pty(cmd: str, timeout: int, request: Request): """Run command in a pseudo-TTY so tqdm/progress bars work natively.""" if not PTY_SUPPORTED: @@ -702,16 +830,17 @@ async def _generate_pty(cmd: str, timeout: int, request: Request): try: while not process_done.is_set(): if deadline and loop.time() > deadline: - proc.kill() - await proc.wait() - yield f"data: {json.dumps({'stream': 'stderr', 'data': f'Command timed out after {timeout}s'})}\n\n" + contained = await _terminate_pty_session(proc) + msg = f"Command timed out after {timeout}s" + if not contained: + msg += PTY_KILL_FAILED_HINT + yield f"data: {json.dumps({'stream': 'stderr', 'data': msg})}\n\n" yield f"data: {json.dumps({'exit_code': -1})}\n\n" return # Check client disconnect if await request.is_disconnected(): - proc.kill() - await proc.wait() + await _terminate_pty_session_quietly(proc) return # Read available data from PTY @@ -773,11 +902,7 @@ async def _generate_pty(cmd: str, timeout: int, request: Request): yield f"data: {json.dumps({'exit_code': proc.returncode})}\n\n" except Exception as e: - try: - proc.kill() - await proc.wait() - except ProcessLookupError: - pass + await _terminate_pty_session_quietly(proc) yield f"data: {json.dumps({'stream': 'stderr', 'data': str(e)})}\n\n" yield f"data: {json.dumps({'exit_code': -1})}\n\n" finally: diff --git a/tests/test_shell_routes.py b/tests/test_shell_routes.py index 072a13d96..281cc8ea6 100644 --- a/tests/test_shell_routes.py +++ b/tests/test_shell_routes.py @@ -1,16 +1,20 @@ """Tests for shell_routes.py helpers.""" +import asyncio import builtins import importlib import importlib.util import json import os +import shlex +import signal import sys from pathlib import Path from types import SimpleNamespace import pytest +from core.platform_compat import pid_alive from routes.shell_routes import ( _find_line_break, _host_docker_access_enabled, @@ -85,6 +89,231 @@ async def test_generate_pty_reports_explicit_unsupported_error(monkeypatch): ] +pty_session = pytest.mark.skipif( + not hasattr(os, "setsid"), reason="process sessions are POSIX-only" +) + + +async def _spawn_pty_style_session(script: str): + """Spawn `script` the way _generate_pty does: its own session via setsid.""" + return await asyncio.create_subprocess_shell( + script, + stdout=asyncio.subprocess.DEVNULL, + stderr=asyncio.subprocess.DEVNULL, + preexec_fn=os.setsid, + ) + + +def _stubborn_child(pid_file: Path, ignore: tuple[str, ...]) -> str: + """Shell snippet that starts a child ignoring `ignore`, then waits for it. + + Killing a PTY session leader makes the kernel send SIGHUP to the + terminal's foreground process group, so a plain `sleep` child looks + contained even when nothing ever signalled the group. A child that ignores + SIGHUP is what an admin actually runs into — a `nohup`ed job, a daemon, + anything meant to outlive its terminal. + + The child publishes its own pid only after installing the handlers, and + the snippet blocks until it does, so a test can never signal it while it + is still starting up and read that as teardown having worked. + """ + ignores = "".join( + f"signal.signal(signal.{name}, signal.SIG_IGN); " for name in ignore + ) + script = ( + f"import os, signal, time; {ignores}" + f"open({str(pid_file)!r}, 'w').write(str(os.getpid())); " + "time.sleep(120)" + ) + return ( + f"{sys.executable} -c {shlex.quote(script)} & " + f"while [ ! -s {pid_file} ]; do sleep 0.02; done" + ) + + +async def _never_disconnected() -> bool: + return False + + +def _reap_if_alive(pid: int) -> None: + """Clean up a descendant the code under test was supposed to have killed.""" + if pid_alive(pid): + try: + os.kill(pid, signal.SIGKILL) + except OSError: + pass + + +async def _read_pid(path: Path, timeout: float = 5.0) -> int: + """Wait for a child to publish its pid, then return it.""" + loop = asyncio.get_running_loop() + deadline = loop.time() + timeout + while loop.time() < deadline: + if path.exists(): + text = path.read_text().strip() + if text: + return int(text) + await asyncio.sleep(0.01) + raise AssertionError(f"child never wrote its pid to {path}") + + +@pty_session +async def test_terminate_pty_session_kills_descendants(tmp_path): + """Tearing down a PTY command takes its children, not only the shell.""" + import routes.shell_routes as shell_routes + + pid_file = tmp_path / "child.pid" + proc = await _spawn_pty_style_session( + f"sleep 120 & echo $! > {pid_file}; sleep 120" + ) + try: + child_pid = await _read_pid(pid_file) + assert pid_alive(child_pid) + + assert await shell_routes._terminate_pty_session(proc) is True + + assert proc.returncode is not None + assert not pid_alive(child_pid) + finally: + await shell_routes._terminate_pty_session(proc) + + +@pty_session +async def test_terminate_pty_session_escalates_past_ignored_sigterm(tmp_path): + """A child that ignores SIGTERM is still gone when teardown returns.""" + import routes.shell_routes as shell_routes + + pid_file = tmp_path / "child.pid" + child = _stubborn_child(pid_file, ("SIGHUP", "SIGTERM")) + proc = await _spawn_pty_style_session(f"{child}; sleep 120") + try: + child_pid = await _read_pid(pid_file) + assert pid_alive(child_pid) + + assert await shell_routes._terminate_pty_session(proc) is True + + assert not pid_alive(child_pid) + finally: + await shell_routes._terminate_pty_session(proc) + + +@pty_session +async def test_generate_pty_timeout_kills_the_whole_session(tmp_path): + """A timed-out PTY command leaves none of its children running.""" + import routes.shell_routes as shell_routes + + pid_file = tmp_path / "child.pid" + child = _stubborn_child(pid_file, ("SIGHUP",)) + cmd = f"{child}; echo ready; sleep 120" + request = SimpleNamespace(is_disconnected=_never_disconnected) + + events = [ + json.loads(chunk.removeprefix("data: ").strip()) + async for chunk in shell_routes._generate_pty(cmd, 1, request) + ] + + child_pid = await _read_pid(pid_file) + try: + assert events[-1] == {"exit_code": -1} + assert events[-2]["data"].startswith("Command timed out after 1s") + assert not pid_alive(child_pid) + finally: + _reap_if_alive(child_pid) + + +@pty_session +async def test_generate_pty_disconnect_kills_the_whole_session(tmp_path): + """Abandoning the stream kills the command's children too.""" + import routes.shell_routes as shell_routes + + pid_file = tmp_path / "child.pid" + child = _stubborn_child(pid_file, ("SIGHUP",)) + cmd = f"{child}; echo ready; sleep 120" + + polls = [] + + async def disconnect_after_first_poll() -> bool: + polls.append(None) + return len(polls) > 1 + + request = SimpleNamespace(is_disconnected=disconnect_after_first_poll) + + async for _ in shell_routes._generate_pty(cmd, 0, request): + pass + + child_pid = await _read_pid(pid_file) + try: + assert not pid_alive(child_pid) + finally: + _reap_if_alive(child_pid) + + +async def test_terminate_pty_session_reports_a_session_it_could_not_kill( + monkeypatch, +): + """Teardown returns False rather than claiming a surviving session died.""" + import routes.shell_routes as shell_routes + + monkeypatch.setattr(shell_routes, "PTY_KILL_GRACE", 0.01) + monkeypatch.setattr(shell_routes, "_signal_session", lambda *_: True) + monkeypatch.setattr(shell_routes, "_session_alive", lambda *_: True) + monkeypatch.setattr(shell_routes, "_session_pgid", lambda _: 4242) + + proc = SimpleNamespace(pid=4242, returncode=0, wait=None) + assert await shell_routes._terminate_pty_session(proc) is False + + +async def test_terminate_pty_session_escalates_before_giving_up(monkeypatch): + """SIGTERM then SIGKILL — the group is never signalled only once.""" + import routes.shell_routes as shell_routes + + sent = [] + monkeypatch.setattr(shell_routes, "PTY_KILL_GRACE", 0.01) + monkeypatch.setattr(shell_routes, "_session_pgid", lambda _: 4242) + monkeypatch.setattr(shell_routes, "_session_alive", lambda *_: True) + monkeypatch.setattr( + shell_routes, + "_signal_session", + lambda pgid, pid, sig: sent.append(sig) or True, + ) + + proc = SimpleNamespace(pid=4242, returncode=0, wait=None) + await shell_routes._terminate_pty_session(proc) + + assert sent == [signal.SIGTERM, signal.SIGKILL] + + +@pty_session +async def test_generate_pty_timeout_says_so_when_the_session_survives( + monkeypatch, +): + """A timed-out command no longer reports clean termination it didn't get.""" + import routes.shell_routes as shell_routes + + real_terminate = shell_routes._terminate_pty_session + + async def terminate_but_report_failure(proc): + await real_terminate(proc) + return False + + monkeypatch.setattr( + shell_routes, "_terminate_pty_session", terminate_but_report_failure + ) + + request = SimpleNamespace(is_disconnected=_never_disconnected) + events = [ + json.loads(chunk.removeprefix("data: ").strip()) + async for chunk in shell_routes._generate_pty("echo ready; sleep 30", 1, request) + ] + + assert events[-1] == {"exit_code": -1} + timed_out = events[-2] + assert timed_out["stream"] == "stderr" + assert timed_out["data"] == ( + "Command timed out after 1s" + shell_routes.PTY_KILL_FAILED_HINT + ) + + class TestFindLineBreak: """Test line-break detection in byte buffers."""