diff --git a/src/google/adk/code_executors/unsafe_local_code_executor.py b/src/google/adk/code_executors/unsafe_local_code_executor.py index f4b52145..1f5c552a 100644 --- a/src/google/adk/code_executors/unsafe_local_code_executor.py +++ b/src/google/adk/code_executors/unsafe_local_code_executor.py @@ -14,15 +14,12 @@ from __future__ import annotations -from contextlib import redirect_stdout -import io import logging -import multiprocessing import os -import queue import re import signal -import traceback +import subprocess +import sys from typing import Any from pydantic import Field @@ -39,82 +36,83 @@ logger = logging.getLogger('google_adk.' + __name__) # escalating to SIGKILL, so that the timeout itself cannot block forever. _TERMINATE_GRACE_SECONDS = 5 +# Runs one program in the child interpreter. +# +# The program arrives on stdin rather than in argv because a single argument is +# capped (at 128 KiB on Linux) and generated programs can carry their own data. +# Reading it leaves stdin at end-of-file, which is what the program would have +# seen before. +# +# The traceback is printed here, without the frame this wrapper contributes, so +# that a failure shows the model its own code and not a file inside this +# package -- a frame it can do nothing about, which then stays in the +# conversation for every later request. +_RUNNER = """ +import sys, traceback -def _execute_in_process( - code: str, - globals_: dict[str, Any], - result_queue: multiprocessing.Queue[tuple[str, str | None]], -) -> None: - """Executes code in a separate process and puts result in queue.""" - # Detach into a new session/process group before running anything, so that a - # timed-out execution can be killed together with everything it spawned. - if hasattr(os, 'setsid'): - try: - os.setsid() - except OSError: - logger.debug('Could not detach the execution process group.') +_run_name = sys.argv[1] +del sys.argv[1:] - stdout = io.StringIO() - error = None - try: - with redirect_stdout(stdout): - exec(code, globals_, globals_) - except BaseException: - error = traceback.format_exc() - result_queue.put((stdout.getvalue(), error)) +_globals = {'__name__': _run_name} if _run_name else {} +_source = sys.stdin.buffer.read().decode('utf-8') + +try: + exec(compile(_source, '', 'exec'), _globals, _globals) +except SystemExit: + # The program chose its own exit status, so let it stand rather than + # reporting a deliberate clean exit as a failure. + raise +except BaseException as exc: + _tb = exc.__traceback__ + traceback.print_exception( + type(exc), exc, _tb.tb_next if _tb else None, file=sys.stderr + ) + sys.exit(1) +""" -def _execution_group( - process: multiprocessing.process.BaseProcess, -) -> int | None: - """Returns the group the execution detached into, or None if it has not.""" - if process.pid is None or not hasattr(os, 'killpg'): - return None - try: - group = os.getpgid(process.pid) - # Only report the group once the execution has detached into its own; - # otherwise the group is still ours and signalling it would take down the - # agent along with the code it is running. - return group if group != os.getpgid(0) else None - except OSError: - return None +def _run_name(code: str) -> str: + """Returns the `__name__` the code should run under, or '' for none.""" + if re.search(r"if\s+__name__\s*==\s*['\"]__main__['\"]", code): + return '__main__' + return '' def _signal_group(group: int, sig: int) -> None: """Signals every process left in a group, tolerating an empty one.""" + if not hasattr(os, 'killpg'): + return try: os.killpg(group, sig) except OSError: logger.debug('Could not signal the execution process group.') -def _kill_execution(process: multiprocessing.process.BaseProcess) -> None: - """Kills a timed-out execution along with any process it spawned.""" - # Resolved up front: once the execution process has been reaped its group can - # no longer be looked up through it, and the group is what holds whatever the - # code spawned. - group = _execution_group(process) +def _kill_execution(process: subprocess.Popen[str]) -> tuple[str, str]: + """Kills a timed-out execution, returning what it wrote before it died.""" + # The execution leads its own process group, so its pid is the group that + # holds whatever it spawned. `terminate` and `kill` below reach the execution + # itself on the platforms that have no group to signal. + group = process.pid # SIGTERM first, so the code and its children get the same grace period the # execution process itself gets before anything is killed outright. - if group is not None: - _signal_group(group, signal.SIGTERM) + _signal_group(group, signal.SIGTERM) process.terminate() - process.join(_TERMINATE_GRACE_SECONDS) + try: + process.wait(_TERMINATE_GRACE_SECONDS) + except subprocess.TimeoutExpired: + pass # Escalate unconditionally: the execution process exiting says nothing about - # a child of it that is ignoring SIGTERM. - if group is not None: - _signal_group(group, signal.SIGKILL) - if process.is_alive(): - process.kill() - process.join() - - -def _prepare_globals(code: str, globals_: dict[str, Any]) -> None: - """Prepare globals for code execution, injecting __name__ if needed.""" - if re.search(r"if\s+__name__\s*==\s*['\"]__main__['\"]", code): - globals_['__name__'] = '__main__' + # a child of it that is ignoring SIGTERM, and such a child holds the output + # pipes open besides. + _signal_group(group, signal.SIGKILL) + process.kill() + try: + return process.communicate(timeout=_TERMINATE_GRACE_SECONDS) + except subprocess.TimeoutExpired: + return '', '' class UnsafeLocalCodeExecutor(BaseCodeExecutor): @@ -144,33 +142,56 @@ class UnsafeLocalCodeExecutor(BaseCodeExecutor): code_execution_input: CodeExecutionInput, ) -> CodeExecutionResult: logger.debug('Executing code:\n```\n%s\n```', code_execution_input.code) - # Execute the code. - globals_: dict[str, Any] = {} - _prepare_globals(code_execution_input.code, globals_) - ctx = multiprocessing.get_context('spawn') - result_queue: multiprocessing.Queue[tuple[str, str | None]] = ctx.Queue() - process = ctx.Process( - target=_execute_in_process, - args=(code_execution_input.code, globals_, result_queue), - daemon=True, + env = dict(os.environ) + # A plain child interpreter starts with its own import path, so pass the + # agent's along: code importing a module from the application resolved + # before this and still does. + env['PYTHONPATH'] = os.pathsep.join(p for p in sys.path if p) + # Pin the child's output encoding. Its stdout is a pipe, so it would + # otherwise encode with the host locale and a program printing non-ASCII + # under a non-UTF-8 locale would die encoding its own output. + env['PYTHONIOENCODING'] = 'utf-8' + + process = subprocess.Popen( + [sys.executable, '-c', _RUNNER, _run_name(code_execution_input.code)], + stdin=subprocess.PIPE, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + encoding='utf-8', + errors='replace', + env=env, + # Its own session, so a timeout can take down everything the code + # started and nothing else. Ignored on Windows, as `os.killpg` is. + start_new_session=True, ) - process.start() - output = '' - error = '' + timed_out = False try: - output, err = result_queue.get(timeout=self.timeout_seconds) - process.join() - if err: - error = err - except queue.Empty: - _kill_execution(process) - error = f'Code execution timed out after {self.timeout_seconds} seconds.' + output, error = process.communicate( + input=code_execution_input.code, timeout=self.timeout_seconds + ) + except subprocess.TimeoutExpired: + output, error = _kill_execution(process) + timed_out = True + + if timed_out: + # Appended rather than assigned: whatever the code wrote before it was + # killed is still the useful diagnostic, but on its own it would hide + # the fact that the run was cut short. + note = f'Code execution timed out after {self.timeout_seconds} seconds.' + error = f'{error}\n{note}' if error else note + elif process.returncode == 0: + # A non-empty stderr is what marks the result failed and drives the retry + # counter, so it has to mean the program failed rather than that the + # program wrote a warning. The exit status is what says which happened. + error = '' + elif not error: + # The code died without saying why: a signal, or a call to `os._exit`. + # Reporting nothing would show the model a clean run. + error = f'Code execution exited with status {process.returncode}.' # Collect the final result. - result_queue.close() - result_queue.join_thread() return CodeExecutionResult( stdout=output, stderr=error, diff --git a/tests/unittests/code_executors/test_unsafe_local_code_executor.py b/tests/unittests/code_executors/test_unsafe_local_code_executor.py index 49b54d89..b5256dfe 100644 --- a/tests/unittests/code_executors/test_unsafe_local_code_executor.py +++ b/tests/unittests/code_executors/test_unsafe_local_code_executor.py @@ -12,9 +12,9 @@ # See the License for the specific language governing permissions and # limitations under the License. -import multiprocessing import os import signal +import subprocess import textwrap import time from unittest.mock import MagicMock @@ -49,6 +49,31 @@ def _is_alive(pid: int) -> bool: return state != "Z" +def _execute_within( + executor: UnsafeLocalCodeExecutor, + invocation_context: InvocationContext, + code: str, + seconds: float, +) -> CodeExecutionResult: + """Executes `code` under a wall-clock bound. + + Several callers below cover code that dies without reporting a result, which + is precisely the case the old executor could wait on forever. The bound is + the executor's own timeout rather than a watchdog around it: a child that + dies closes its pipes, so the wait ends on its own, and passing the timeout + through also exercises the path that enforces it. + """ + started = time.monotonic() + result = executor.execute_code( + invocation_context, CodeExecutionInput(code=code) + ) + elapsed = time.monotonic() - started + assert ( + elapsed < seconds * 4 + ), f"the executor took {elapsed:.1f}s against a {seconds}s timeout" + return result + + @pytest.fixture def mock_invocation_context() -> InvocationContext: """Provides a mock InvocationContext.""" @@ -156,6 +181,185 @@ class TestUnsafeLocalCodeExecutor: assert result.stdout == "" assert "Code execution timed out after 1 seconds." in result.stderr + def test_execute_code_main_guard_runs( + self, mock_invocation_context: InvocationContext + ): + """Code guarded on `__main__` runs, as it did in the previous child.""" + executor = UnsafeLocalCodeExecutor() + code_input = CodeExecutionInput(code=textwrap.dedent(""" + if __name__ == '__main__': + print('guarded') + """)) + + result = executor.execute_code(mock_invocation_context, code_input) + + assert result.stdout == "guarded\n" + assert result.stderr == "" + + def test_execute_code_without_main_guard_is_not_main( + self, mock_invocation_context: InvocationContext + ): + """Code that does not ask to be `__main__` is not given the name.""" + executor = UnsafeLocalCodeExecutor() + code_input = CodeExecutionInput(code="print(globals().get('__name__'))") + + result = executor.execute_code(mock_invocation_context, code_input) + + assert result.stdout == "None\n" + + def test_execute_code_separates_stdout_from_stderr( + self, mock_invocation_context: InvocationContext + ): + """Each stream lands in its own field rather than being interleaved.""" + executor = UnsafeLocalCodeExecutor() + code_input = CodeExecutionInput(code=textwrap.dedent(""" + import sys + sys.stdout.write('to out') + sys.stderr.write('to err') + sys.exit(0) + """)) + + result = executor.execute_code(mock_invocation_context, code_input) + + assert result.stdout == "to out" + # The program wrote to stderr but succeeded. A non-empty stderr is what + # marks a result failed and drives the retry counter, so a warning must not + # be reported as a failure. + assert result.stderr == "" + + def test_execute_code_nonzero_exit_is_reported_as_a_failure( + self, mock_invocation_context: InvocationContext + ): + """A program that exits non-zero fails even when it says nothing.""" + executor = UnsafeLocalCodeExecutor() + code_input = CodeExecutionInput(code="import sys\nsys.exit(3)") + + result = executor.execute_code(mock_invocation_context, code_input) + + assert result.stdout == "" + assert result.stderr == "Code execution exited with status 3." + + def test_execute_code_reports_a_crash_rather_than_hanging( + self, mock_invocation_context: InvocationContext + ): + """Code that exits without reporting anything still returns a result.""" + executor = UnsafeLocalCodeExecutor() + + result = _execute_within( + executor, + mock_invocation_context, + "import os\nprint('before', flush=True)\nos._exit(3)", + seconds=30, + ) + + assert result.stdout == "before\n" + assert result.stderr == "Code execution exited with status 3." + + @pytest.mark.skipif( + not hasattr(signal, "SIGKILL"), + reason="Death by signal is checked on POSIX only.", + ) + def test_execute_code_reports_death_by_signal( + self, mock_invocation_context: InvocationContext + ): + """Code killed outright -- by a segfault or the OOM killer -- returns.""" + executor = UnsafeLocalCodeExecutor() + + result = _execute_within( + executor, + mock_invocation_context, + "import os, signal\nos.kill(os.getpid(), signal.SIGKILL)", + seconds=30, + ) + + assert result.stdout == "" + assert result.stderr == ( + f"Code execution exited with status {-signal.SIGKILL}." + ) + + def test_execute_code_traceback_omits_this_module( + self, mock_invocation_context: InvocationContext + ): + """A failure shows the model its own code, not a frame from this package.""" + executor = UnsafeLocalCodeExecutor() + code_input = CodeExecutionInput( + code="def divide():\n return 1 / 0\n\ndivide()" + ) + + result = executor.execute_code(mock_invocation_context, code_input) + + assert "ZeroDivisionError" in result.stderr + assert "unsafe_local_code_executor" not in result.stderr + # Both frames of the executed code survive; only the wrapper's is dropped. + assert result.stderr.count('File ""') == 2 + + def test_execute_code_preserves_unicode( + self, mock_invocation_context: InvocationContext + ): + executor = UnsafeLocalCodeExecutor() + code_input = CodeExecutionInput(code="print('你好, café')") + + result = executor.execute_code(mock_invocation_context, code_input) + + assert result.stdout == "你好, café\n" + assert result.stderr == "" + + def test_execute_code_output_encoding_does_not_follow_the_host_locale( + self, mock_invocation_context: InvocationContext + ): + """Otherwise a program printing non-ASCII dies encoding its own output.""" + executor = UnsafeLocalCodeExecutor() + code_input = CodeExecutionInput( + code="import sys\nprint(sys.stdout.encoding)" + ) + + result = executor.execute_code(mock_invocation_context, code_input) + + assert result.stdout.strip().lower() == "utf-8" + + def test_execute_code_large_output( + self, mock_invocation_context: InvocationContext + ): + """Output far larger than a pipe buffer is read out rather than deadlocked.""" + executor = UnsafeLocalCodeExecutor() + code_input = CodeExecutionInput(code="print('x' * 1000000)") + + result = _execute_within( + executor, mock_invocation_context, code_input.code, seconds=60 + ) + + assert result.stdout == "x" * 1000000 + "\n" + assert result.stderr == "" + + def test_execute_code_large_program( + self, mock_invocation_context: InvocationContext + ): + """A program carrying its own data is not capped by an argument limit.""" + executor = UnsafeLocalCodeExecutor() + payload = "a" * 300000 + code_input = CodeExecutionInput( + code=f"data = {payload!r}\nprint(len(data))" + ) + + result = executor.execute_code(mock_invocation_context, code_input) + + assert result.stdout == f"{len(payload)}\n" + assert result.stderr == "" + + def test_execute_code_imports_resolve_from_the_agent_path( + self, mock_invocation_context: InvocationContext + ): + """Code importing what the application can import still resolves.""" + executor = UnsafeLocalCodeExecutor() + code_input = CodeExecutionInput(code="import google.adk\nprint('imported')") + + result = _execute_within( + executor, mock_invocation_context, code_input.code, seconds=120 + ) + + assert result.stdout == "imported\n" + assert result.stderr == "" + def test_kill_execution_signals_group_before_killing_it(self, monkeypatch): """The group gets SIGTERM and its grace period before SIGKILL.""" signalled = [] @@ -163,35 +367,51 @@ class TestUnsafeLocalCodeExecutor: unsafe_local_code_executor.os, "killpg", lambda group, sig: signalled.append((group, sig)), - ) - monkeypatch.setattr( - unsafe_local_code_executor, - "_execution_group", - lambda process: 4321, + raising=False, ) process = MagicMock() - process.is_alive.return_value = False + process.pid = 4321 process.terminate.side_effect = lambda: signalled.append(("child", "term")) + process.kill.side_effect = lambda: signalled.append(("child", "kill")) + process.communicate.return_value = ("out", "err") - unsafe_local_code_executor._kill_execution(process) + assert unsafe_local_code_executor._kill_execution(process) == ("out", "err") assert signalled == [ (4321, signal.SIGTERM), ("child", "term"), (4321, signal.SIGKILL), + ("child", "kill"), ] - process.join.assert_any_call( + process.wait.assert_called_once_with( unsafe_local_code_executor._TERMINATE_GRACE_SECONDS ) + def test_kill_execution_gives_up_on_pipes_that_never_close(self, monkeypatch): + """A pipe held open by something unkillable does not block the agent.""" + monkeypatch.setattr( + unsafe_local_code_executor.os, + "killpg", + lambda group, sig: None, + raising=False, + ) + process = MagicMock() + process.pid = 4321 + process.wait.side_effect = subprocess.TimeoutExpired("cmd", 5) + process.communicate.side_effect = subprocess.TimeoutExpired("cmd", 5) + + assert unsafe_local_code_executor._kill_execution(process) == ("", "") + @pytest.mark.skipif( not hasattr(os, "killpg") or not hasattr(os, "fork") or not os.path.isdir("/proc"), reason="Process-group teardown is checked on POSIX with /proc only.", ) - def test_kill_execution_kills_what_the_code_spawned(self, tmp_path): - """Killing a live execution takes the processes it spawned with it.""" + def test_timeout_kills_what_the_code_spawned( + self, mock_invocation_context: InvocationContext, tmp_path + ): + """A timed-out execution takes the processes it spawned with it.""" pid_file = tmp_path / "spawned.pid" # Forked rather than spawned through `sys.executable`, so the descendant # exists within milliseconds and the test never waits on interpreter @@ -208,29 +428,18 @@ class TestUnsafeLocalCodeExecutor: f.write(str(spawned)) time.sleep(60) """) - ctx = multiprocessing.get_context("spawn") - result_queue = ctx.Queue() - process = ctx.Process( - target=unsafe_local_code_executor._execute_in_process, - args=(code, {}, result_queue), - daemon=True, - ) - process.start() + executor = UnsafeLocalCodeExecutor(timeout_seconds=5) + spawned_pid = None try: - # Waiting for the pid to be written rather than for a fixed duration: - # the only thing that has to have happened is the fork. The file exists - # from the moment it is opened, so its content is what is polled for. - deadline = time.time() + 30 - while time.time() < deadline and not _written_pid(pid_file): - time.sleep(0.05) + result = _execute_within( + executor, mock_invocation_context, code, seconds=60 + ) + assert "Code execution timed out after 5 seconds." in result.stderr + spawned_pid = _written_pid(pid_file) if spawned_pid is None: pytest.skip("this environment could not start the execution process") - - unsafe_local_code_executor._kill_execution(process) - - assert not process.is_alive() deadline = time.time() + 10 while time.time() < deadline and _is_alive(spawned_pid): time.sleep(0.05) @@ -241,7 +450,3 @@ class TestUnsafeLocalCodeExecutor: os.kill(spawned_pid, signal.SIGKILL) except OSError: pass - if process.is_alive(): - process.kill() - process.join() - result_queue.close()