Skip to content
Open
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
138 changes: 138 additions & 0 deletions arc/job/local.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
import os
import re
import shutil
import signal
import subprocess
import time

Expand All @@ -28,20 +29,37 @@ def execute_command(command: str | list[str],
shell: bool = True,
no_fail: bool = False,
executable: str | None = None,
timeout: float | None = None,
) -> tuple[list | None, list | None]:
"""
Execute a command.

Notes:
If ``no_fail`` is ``True``, then a warning is logged and ``False`` is returned
so that the calling function can debug the situation.
A command that exceeds ``timeout`` is not retried: the point of a deadline is to bound
the call, so the timeout is reported through the same channels as a terminal failure.

Args:
command (str | list[str]): An array of string commands to send.
shell (bool, optional): Specifies whether the command should be executed using bash instead of Python.
no_fail (bool, optional): If ``True`` then ARC will not crash if an error is encountered.
executable (str, optional): Select a specific shell to run with, e.g., '/bin/bash'.
Default shell of the subprocess command is '/bin/sh'.
timeout (float, optional): The number of seconds to wait for the command to complete.
If the command does not complete in time, it is killed along with
every process it spawned that stayed in its process group. A
descendant that puts itself in a group of its own, e.g. via
``setsid``, is outside what signalling a group can reach and
survives. ``None``, the default, waits forever.
Note that the deadline is not exact: a timing out call returns
after ``timeout`` plus up to twice the grace period of
``_kill_process_group()``, i.e. up to 10 seconds late by default.

Raises:
SettingsError: If the command timed out and ``no_fail`` is ``False``. A non-zero exit status
is not an error here: the command is run without ``check=True``, so a failing
command's stderr is returned as output rather than raised.

Returns: tuple[list, list]:
- A list of lines of standard output stream.
Expand All @@ -55,12 +73,28 @@ def execute_command(command: str | list[str],
sleep_time = 60 # Seconds
while i < max_times_to_try:
try:
if timeout is not None:
stdout, stderr = _run_with_timeout(command=command, shell=shell,
executable=executable, timeout=timeout)
return _format_stdout(stdout), _format_stdout(stderr)
if executable is None:
completed_process = subprocess.run(command, shell=shell, capture_output=True)
else:
completed_process = subprocess.run(command, shell=shell, capture_output=True, executable=executable)
return _format_stdout(completed_process.stdout), _format_stdout(completed_process.stderr)
except subprocess.TimeoutExpired:
message = f'The command "{command}" did not complete within {timeout} seconds ' \
f'and was terminated along with all of the processes it spawned.'
if no_fail:
logger.warning(message)
return None, None
logger.error(message)
raise SettingsError(f'{message}\nConsider increasing the timeout, or check whether this is a '
f'server issue by executing the command manually on the server.') from None
except subprocess.CalledProcessError as e:
# Note: ``subprocess.run()`` is called without ``check=True`` and therefore never
# raises this, so this handler and the retry loop it drives are dead code. They are
# left as they are; the ``TimeoutExpired`` handler above is the only live one.
error = e # Store the error so we can raise a SettingsError if needed.
if no_fail:
_output_command_error_message(command, e, logger.warning)
Expand All @@ -82,6 +116,110 @@ def execute_command(command: str | list[str],
f'sbatch path required in the submit_command dictionary.')


def _run_with_timeout(command: list[str],
shell: bool,
executable: str | None,
timeout: float,
) -> tuple[bytes, bytes]:
"""
Run a command under a deadline, killing its process group if the deadline expires.

The command is started in its own session (hence its own process group) so that the whole tree
can be signalled at once. Killing only the direct child is not enough: with ``shell=True`` the
direct child is a shell that commonly ``exec``s its last command, and anything it backgrounded
would survive as an orphan.

The containment this buys reaches exactly as far as the process group does. A descendant that
leaves the group, e.g. by calling ``setsid`` itself, receives neither signal and keeps running;
holding it too would take a cgroup, which is beyond what this function attempts.

Note that a timing out call returns only after the grace periods of ``_kill_process_group()``,
which sleeps for one before escalating to SIGKILL and then waits up to another for the child.
The call is therefore bounded by ``timeout + 2 * grace_period``, i.e. up to 10 seconds beyond
the deadline with the default grace period, rather than by ``timeout`` alone.

Args:
command (list[str]): The command to run, already normalized by ``execute_command()`` into
the single-element list that ``subprocess`` expects.
shell (bool): Specifies whether the command should be executed using bash instead of Python.
executable (str | None): Select a specific shell to run with, e.g., '/bin/bash'.
timeout (float): The number of seconds to wait for the command to complete.

Raises:
subprocess.TimeoutExpired: If the command did not complete within ``timeout`` seconds.

Returns: tuple[bytes, bytes]:
- The standard output stream.
- The standard error stream.
"""
process = subprocess.Popen(command,
shell=shell,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
executable=executable,
start_new_session=True,
)
try:
return process.communicate(timeout=timeout)
except subprocess.TimeoutExpired:
_kill_process_group(process=process)
raise
finally:
# ``Popen`` is deliberately not used as a context manager here: its ``__exit__()`` waits for
# the child without a timeout, which would defeat the very deadline this function enforces
# whenever a child cannot be reaped promptly. Closing the pipes is therefore done here, and
# reaping is left to ``communicate()`` on the normal path and to the bounded wait of
# ``_kill_process_group()`` on the timeout path.
for stream in (process.stdout, process.stderr):
if stream is not None:
stream.close()


def _kill_process_group(process: subprocess.Popen,
grace_period: float = 5,
) -> None:
"""
Kill a process and every process it spawned that is still in its process group.

SIGTERM is sent first to let the tree shut down cleanly, then SIGKILL after ``grace_period``.
Falls back to killing just the given process if it does not have a process group of its own,
which also guards against ever signalling the group ARC itself runs in.

Args:
process (subprocess.Popen): The process to kill, started with ``start_new_session=True``.
grace_period (float, optional): The number of seconds to allow the tree to exit on SIGTERM.
"""
try:
pgid = os.getpgid(process.pid)
except OSError:
pgid = None
if pgid is None or pgid == os.getpgid(0):
process.kill()
else:
try:
os.killpg(pgid, signal.SIGTERM)
except OSError:
Comment thread
github-advanced-security[bot] marked this conversation as resolved.
Fixed
pass # The group is already gone, which is the outcome this call is after anyway.
Comment thread
alongd marked this conversation as resolved.
# The child is deliberately not reaped here. It is the group leader, so while it remains
# unreaped its pid cannot be recycled and ``pgid`` stays pinned to this group, which is
# what makes the SIGKILL below safe to send. Reaping first would open a window in which
# the pid is free and the SIGKILL could land on an unrelated process group.
# The escalation also cannot be conditioned on the child having survived SIGTERM: with
# ``shell=True`` the direct child is a shell that dies on SIGTERM, while a grandchild
# that ignores SIGTERM is precisely the process that gets orphaned.
time.sleep(grace_period)
try:
os.killpg(pgid, signal.SIGKILL)
except OSError:
Comment thread
github-advanced-security[bot] marked this conversation as resolved.
Fixed
pass # The group exited on SIGTERM, which is the preferred outcome.
try:
process.wait(timeout=grace_period)
except subprocess.TimeoutExpired:
logger.warning(f'Could not terminate process {process.pid} after it timed out. '
f'It is left unreaped rather than waited for indefinitely, so that a child '
f'which cannot be killed does not also hang ARC.')


def _output_command_error_message(command: list[str],
error: subprocess.CalledProcessError,
logging_func,
Expand Down
126 changes: 126 additions & 0 deletions arc/job/local_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,11 +8,71 @@
import datetime
import os
import shutil
import signal
import sys
import tempfile
import time
import unittest
from unittest.mock import patch

import arc.job.local as local
from arc.common import ARC_PATH
from arc.exceptions import SettingsError


def process_is_alive(pid: int) -> bool:
"""
Whether a process with the given pid exists and is not a zombie.

Args:
pid (int): The process ID to check.

Returns:
bool: Whether the process is alive.
"""
try:
os.kill(pid, 0)
except OSError:
return False
try:
with open(f'/proc/{pid}/stat', 'r') as f:
return f.read().rsplit(')', 1)[-1].split()[0] != 'Z'
except (OSError, IndexError):
return True


def get_parent_pid(pid: int) -> int | None:
"""
Get the parent pid of a process, used to report whether a process was reparented to init.

Args:
pid (int): The process ID to check.

Returns:
int | None: The parent process ID, or ``None`` if it could not be determined.
"""
try:
with open(f'/proc/{pid}/stat', 'r') as f:
return int(f.read().rsplit(')', 1)[-1].split()[1])
except (OSError, IndexError, ValueError):
return None


def _read_pid(path: str) -> int | None:
"""
Read a pid from a file, returning ``None`` if it is absent or unreadable.

Args:
path (str): The path of the file holding the pid.

Returns:
int | None: The pid, or ``None``.
"""
try:
with open(path, 'r') as f:
return int(f.read().strip())
except (OSError, ValueError):
return None


class TestLocal(unittest.TestCase):
Expand All @@ -37,6 +97,72 @@ def test_execute_command(self):
self.assertIn('adapter.py', out1[0])
self.assertIn('ssh.py', out1[0])

def test_execute_command_without_a_timeout_is_unchanged(self):
"""Test that not passing a timeout leaves the subprocess call exactly as it was"""
with patch('arc.job.local.subprocess.run') as mock_run:
mock_run.return_value.stdout, mock_run.return_value.stderr = b'ok\n', b''
local.execute_command('ls')
local.execute_command('ls', executable='/bin/bash')
self.assertEqual(len(mock_run.call_args_list), 2)
for call in mock_run.call_args_list:
self.assertNotIn('timeout', call.kwargs)
self.assertNotIn('start_new_session', call.kwargs)
self.assertEqual(mock_run.call_args_list[0].kwargs, {'shell': True, 'capture_output': True})
self.assertEqual(mock_run.call_args_list[1].kwargs,
{'shell': True, 'capture_output': True, 'executable': '/bin/bash'})

def test_execute_command_with_a_generous_timeout(self):
"""Test that a command which completes in time is unaffected by a timeout"""
stdout, stderr = local.execute_command('echo hello', timeout=60)
self.assertEqual(stdout, ['hello'])
self.assertEqual(stderr, [])

def test_execute_command_timeout_kills_spawned_children(self):
"""Test that a timing out command is killed along with the processes it spawned"""
# The grandchild ignores SIGTERM, so this pins both halves of the kill:
# signalling the whole process group rather than only the direct child, and escalating
# to SIGKILL. A grandchild that dies on SIGTERM cannot tell the escalation apart.
# The direct child here is a plain ``sleep`` that does die on SIGTERM, which is exactly
# why the escalation must not be conditioned on the direct child having survived.
temp_dir = tempfile.mkdtemp()
pid_path = os.path.join(temp_dir, 'grandchild.pid')
script_path = os.path.join(temp_dir, 'stubborn_grandchild.py')
with open(script_path, 'w') as f:
f.write('import os, signal, sys, time\n'
'signal.signal(signal.SIGTERM, signal.SIG_IGN)\n'
"with open(sys.argv[1], 'w') as f:\n"
" f.write(str(os.getpid()))\n"
'time.sleep(300)\n')
# The shell waits for the grandchild to record its pid before starting the long sleep, so a
# loaded worker cannot have the timeout fire while the grandchild is still starting up and
# fail the assertion below even though process group cleanup worked correctly.
command = f'{sys.executable} {script_path} {pid_path} & ' \
f'while [ ! -s {pid_path} ]; do sleep 0.05; done; sleep 300'
try:
with self.assertRaises(SettingsError):
local.execute_command(command, timeout=3)
self.assertTrue(os.path.isfile(pid_path), 'The grandchild never recorded its pid.')
with open(pid_path, 'r') as f:
grandchild_pid = int(f.read().strip())
for _ in range(200):
if not process_is_alive(grandchild_pid):
break
time.sleep(0.1)
self.assertFalse(process_is_alive(grandchild_pid),
f'The spawned process {grandchild_pid} was orphaned instead of killed '
f'(parent pid {get_parent_pid(grandchild_pid)}).')
finally:
for pid in [_read_pid(pid_path)]:
if pid is not None and process_is_alive(pid):
os.kill(pid, signal.SIGKILL)
shutil.rmtree(temp_dir, ignore_errors=True)

def test_execute_command_timeout_with_no_fail(self):
"""Test that a timing out command returns None, None when no_fail is True"""
stdout, stderr = local.execute_command('sleep 300', no_fail=True, timeout=2)
self.assertIsNone(stdout)
self.assertIsNone(stderr)

def test_determine_job_id(self):
"""Test determining a job ID from the stdout of a job submission command."""
# Slurm
Expand Down
Loading
Loading