Skip to content
Merged
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
375 changes: 375 additions & 0 deletions Lib/test/test_external_inspection.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import unittest
from contextlib import contextmanager
import asyncio
import os
import textwrap
Expand Down Expand Up @@ -1437,5 +1438,379 @@ def lines(u, expected_count):
)




TRANSIENT_ERRORS = (OSError, RuntimeError, UnicodeDecodeError)

def _create_server_socket(port, backlog=1):
"""Create and configure a server socket for test communication."""
server_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
server_socket.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
server_socket.bind(("localhost", port))
server_socket.settimeout(SHORT_TIMEOUT)
server_socket.listen(backlog)
return server_socket


def _wait_for_signal(sock, expected_signals, timeout=SHORT_TIMEOUT):
"""
Wait for expected signal(s) from a socket with proper timeout and EOF handling.

Args:
sock: Connected socket to read from
expected_signals: Single bytes object or list of bytes objects to wait for
timeout: Socket timeout in seconds

Returns:
bytes: Complete accumulated response buffer

Raises:
RuntimeError: If connection closed before signal received or timeout
"""
if isinstance(expected_signals, bytes):
expected_signals = [expected_signals]

sock.settimeout(timeout)
buffer = b""

while True:
# Check if all expected signals are in buffer
if all(sig in buffer for sig in expected_signals):
return buffer

try:
chunk = sock.recv(4096)
if not chunk:
# EOF - connection closed
raise RuntimeError(
f"Connection closed before receiving expected signals. "
f"Expected: {expected_signals}, Got: {buffer[-200:]!r}"
)
buffer += chunk
except socket.timeout:
raise RuntimeError(
f"Timeout waiting for signals. "
f"Expected: {expected_signals}, Got: {buffer[-200:]!r}"
)


@contextmanager
def _managed_subprocess(args, timeout=SHORT_TIMEOUT):
"""
Context manager for subprocess lifecycle management.

Ensures process is properly terminated and cleaned up even on exceptions.
Uses graceful termination first, then forceful kill if needed.
"""
p = subprocess.Popen(args)
try:
yield p
finally:
try:
p.terminate()
try:
p.wait(timeout=timeout)
except subprocess.TimeoutExpired:
p.kill()
try:
p.wait(timeout=timeout)
except subprocess.TimeoutExpired:
pass # Process refuses to die, nothing more we can do
except OSError:
pass # Process already dead


def _cleanup_sockets(*sockets):
"""Safely close multiple sockets, ignoring errors."""
for sock in sockets:
if sock is not None:
try:
sock.close()
except OSError:
pass



class RemoteInspectionTestBase(unittest.TestCase):
@contextmanager
def _target_process(self, script_body):
"""Context manager for running a target process with socket sync."""
port = find_unused_port()
script = f"""\
import socket
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
sock.connect(('localhost', {port}))
{textwrap.dedent(script_body)}
"""

with os_helper.temp_dir() as work_dir:
script_dir = os.path.join(work_dir, "script_pkg")
os.mkdir(script_dir)

server_socket = _create_server_socket(port)
script_name = _make_test_script(script_dir, "script", script)
client_socket = None

try:
with _managed_subprocess([sys.executable, script_name]) as p:
client_socket, _ = server_socket.accept()
server_socket.close()
server_socket = None

def make_unwinder():
try:
return RemoteUnwinder(p.pid, all_threads=True)
except PermissionError:
self.skipTest("Insufficient permissions to read the stack trace")

yield p, client_socket, make_unwinder
finally:
_cleanup_sockets(client_socket, server_socket)


def _get_task_id_map(self, stack_trace):
"""Create task_id -> task mapping from async stack trace."""
return {task.task_id: task for task in stack_trace[0].awaited_by}


def _get_awaited_by_relationships(self, stack_trace):
"""Extract task name to awaited_by set mapping."""
id_to_task = self._get_task_id_map(stack_trace)
return {
task.task_name: set(
id_to_task[awaited.task_name].task_name
for awaited in task.awaited_by
)
for task in stack_trace[0].awaited_by
}



class TestFrameChainLimits(RemoteInspectionTestBase):
"""Frame chain walks abort instead of looping/overflowing on deep chains."""

# Limits plus one, to exceed them (must match MAX_FRAME_CHAIN_DEPTH /
# MAX_TASK_WAITER_WALK_TASKS from _remote_debugging_module.c)
FRAME_CHAIN_DEPTH = 1024 + 1
TASK_WAITER_WALK_TASKS = 2**14 + 1

def _assert_unwinder_limit_error(self, unwind, expected_substring):
"""Call unwind() until it raises the frame chain limit error.

unwind must construct the RemoteUnwinder and call it, so that
transient RuntimeErrors from either step are retried; a successful
call means the limit never triggered and fails immediately.
"""
last_error = None
for _ in busy_retry(SHORT_TIMEOUT, error=False):
try:
unwind()
except PermissionError:
self.skipTest("Insufficient permissions to read the stack trace")
except TRANSIENT_ERRORS as e:
if expected_substring in str(e):
return
last_error = e
continue
self.fail(
"frame chain limit did not trigger; call returned a result"
)
self.fail(
f"frame chain limit never raised; last transient error: "
f"{last_error!r}"
)

@skip_if_not_supported
@unittest.skipIf(
sys.platform == "linux" and not PROCESS_VM_READV_SUPPORTED,
"Test only runs on Linux with process_vm_readv support",
)
def test_get_stack_trace_deep_frame_chain_aborts(self):
"""Test that a frame chain deeper than the limit aborts the
synchronous stack walk instead of walking it indefinitely."""
script_body = f"""\
import sys
sys.setrecursionlimit({self.FRAME_CHAIN_DEPTH * 2})

def recurse(n):
if n <= 0:
sock.sendall(b"ready")
sock.recv(16)
return
recurse(n - 1)

recurse({self.FRAME_CHAIN_DEPTH})
"""
with self._target_process(script_body) as (p, client_socket, _):
_wait_for_signal(client_socket, b"ready")
self._assert_unwinder_limit_error(
lambda: RemoteUnwinder(p.pid).get_stack_trace(),
"Too many stack frames",
)
client_socket.sendall(b"done")

@skip_if_not_supported
@unittest.skipIf(
sys.platform == "linux" and not PROCESS_VM_READV_SUPPORTED,
"Test only runs on Linux with process_vm_readv support",
)
def test_get_async_stack_trace_deep_task_waiter_chain_aborts(self):
"""Test that a task waiter chain deeper than the limit aborts
the walk instead of overflowing the C stack."""
script_body = f"""\
import asyncio

async def chain(n):
if n <= 0:
sock.sendall(b"ready")
sock.recv(16)
return

task = asyncio.create_task(chain(n - 1))
await task

asyncio.run(chain({self.TASK_WAITER_WALK_TASKS}))
"""
with self._target_process(script_body) as (p, client_socket, _):
_wait_for_signal(client_socket, b"ready")
self._assert_unwinder_limit_error(
lambda: RemoteUnwinder(p.pid).get_async_stack_trace(),
"Too many task waiters",
)
client_socket.sendall(b"done")

@skip_if_not_supported
@unittest.skipIf(
sys.platform == "linux" and not PROCESS_VM_READV_SUPPORTED,
"Test only runs on Linux with process_vm_readv support",
)
def test_get_async_stack_trace_deep_frame_chain_aborts(self):
"""Test that a frame chain deeper than the limit aborts the async
stack walk instead of walking it indefinitely."""
script_body = f"""\
import sys, asyncio
sys.setrecursionlimit({self.FRAME_CHAIN_DEPTH * 2})

def recurse(n):
if n <= 0:
sock.sendall(b"ready")
sock.recv(16)
return
recurse(n - 1)

async def deep():
recurse({self.FRAME_CHAIN_DEPTH})

asyncio.run(deep())
"""
with self._target_process(script_body) as (p, client_socket, _):
_wait_for_signal(client_socket, b"ready")
self._assert_unwinder_limit_error(
lambda: RemoteUnwinder(p.pid).get_async_stack_trace(),
"Too many async stack frames",
)
client_socket.sendall(b"done")

@skip_if_not_supported
@unittest.skipIf(
sys.platform == "linux" and not PROCESS_VM_READV_SUPPORTED,
"Test only runs on Linux with process_vm_readv support",
)
def test_get_all_awaited_by_deep_coro_chain_aborts(self):
"""Test that a coroutine await chain deeper than the limit aborts
the walk instead of overflowing the C stack."""
script_body = f"""\
import sys, asyncio
sys.setrecursionlimit({self.FRAME_CHAIN_DEPTH * 2})

async def chain(n):
if n <= 0:
await asyncio.sleep(10_000)
return
await chain(n - 1)

async def main():
task = asyncio.create_task(chain({self.FRAME_CHAIN_DEPTH}))
await asyncio.sleep(0)
sock.sendall(b"ready")
await task

asyncio.run(main())
"""
with self._target_process(script_body) as (p, client_socket, _):
_wait_for_signal(client_socket, b"ready")
self._assert_unwinder_limit_error(
lambda: RemoteUnwinder(p.pid).get_all_awaited_by(),
"Too many coroutine frames",
)


@skip_if_not_supported
@unittest.skipIf(
sys.platform == "linux" and not PROCESS_VM_READV_SUPPORTED,
"Test only runs on Linux with process_vm_readv support",
)
def test_async_awaited_by_skips_set_tombstones(self):
script_body = """\
import asyncio

class RemovedTask(asyncio.Task):
def __hash__(self):
return 0

class RemainingTask(asyncio.Task):
def __hash__(self):
return 1

async def main():
victim = asyncio.current_task()
victim.set_name("victim")
removed = RemovedTask(
asyncio.sleep(10_000), name="removed"
)
remaining = RemainingTask(
asyncio.sleep(10_000), name="remaining"
)

asyncio.future_add_to_awaited_by(victim, removed)
asyncio.future_add_to_awaited_by(victim, remaining)

# Removing hash 0 leaves a dummy in slot 0 before the only
# active entry in slot 1. It must not count toward the set's
# used entries.
asyncio.future_discard_from_awaited_by(victim, removed)

sock.sendall(b"ready")
sock.recv(16)

asyncio.run(main())
"""

with self._target_process(script_body) as (
_,
client_socket,
make_unwinder,
):
_wait_for_signal(client_socket, b"ready")

for method_name in (
"get_async_stack_trace",
"get_all_awaited_by",
):
with self.subTest(method=method_name):
unwinder = make_unwinder()
stack_trace = getattr(unwinder, method_name)()
relationships = self._get_awaited_by_relationships(
stack_trace
)
self.assertEqual(
relationships["victim"],
{"remaining"},
)

client_socket.sendall(b"done")



if __name__ == "__main__":
unittest.main()
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
Make frame, coroutine and task-waiter walks iterative and bounded, avoiding
potential hangs and stack overflows. Fix asyncio task inspection when
awaited-by sets contain removed entries. Patch by Maurycy Pawłowski-Wieroński.
Loading
Loading