Skip to content

Commit f660816

Browse files
pablogsalmaurycy
andauthored
[3.14] gh-153364: Make frame, coroutine, and task-waiter chain walks iterative and bounded (GH-153365) (#158814)
gh-153364: Make frame, coroutine, and task-waiter chain walks iterative and bounded (#153365) * let me declare single limit * use our new limit in process_frame_chain() * add it in parse_async_frame_chain() * parse_coro_chain() * NEWS * async in the message? * test * no race * process_task_awaited_by * process_task_awaited_by limit test * NEWS * MAX_TASK_WAITER_CHAIN_DEPTH * TASK_WAITER_CHAIN_DEPTH in test * TASK_WAITER_CHAIN_DEPTH 256 * prevent the drift with the comment * better naming, better style * MAX_TASK_WAITER_CHAIN_DEPTH comment * task-waiter iterative bfs walk * iterative coro-walk * nicer news * 1 << 14 * comment * unused read_Py_ssize_t * fix tombstones * simplify * correct msg * better test * news for tombstones * left-over from when testing buggy version * redundant new line (cherry picked from commit e0861c6) Co-authored-by: Maurycy Pawłowski-Wieroński <maurycy@maurycy.com>
1 parent 3c6de41 commit f660816

3 files changed

Lines changed: 490 additions & 113 deletions

File tree

‎Lib/test/test_external_inspection.py‎

Lines changed: 375 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import unittest
2+
from contextlib import contextmanager
23
import asyncio
34
import os
45
import textwrap
@@ -1437,5 +1438,379 @@ def lines(u, expected_count):
14371438
)
14381439

14391440

1441+
1442+
1443+
TRANSIENT_ERRORS = (OSError, RuntimeError, UnicodeDecodeError)
1444+
1445+
def _create_server_socket(port, backlog=1):
1446+
"""Create and configure a server socket for test communication."""
1447+
server_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
1448+
server_socket.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
1449+
server_socket.bind(("localhost", port))
1450+
server_socket.settimeout(SHORT_TIMEOUT)
1451+
server_socket.listen(backlog)
1452+
return server_socket
1453+
1454+
1455+
def _wait_for_signal(sock, expected_signals, timeout=SHORT_TIMEOUT):
1456+
"""
1457+
Wait for expected signal(s) from a socket with proper timeout and EOF handling.
1458+
1459+
Args:
1460+
sock: Connected socket to read from
1461+
expected_signals: Single bytes object or list of bytes objects to wait for
1462+
timeout: Socket timeout in seconds
1463+
1464+
Returns:
1465+
bytes: Complete accumulated response buffer
1466+
1467+
Raises:
1468+
RuntimeError: If connection closed before signal received or timeout
1469+
"""
1470+
if isinstance(expected_signals, bytes):
1471+
expected_signals = [expected_signals]
1472+
1473+
sock.settimeout(timeout)
1474+
buffer = b""
1475+
1476+
while True:
1477+
# Check if all expected signals are in buffer
1478+
if all(sig in buffer for sig in expected_signals):
1479+
return buffer
1480+
1481+
try:
1482+
chunk = sock.recv(4096)
1483+
if not chunk:
1484+
# EOF - connection closed
1485+
raise RuntimeError(
1486+
f"Connection closed before receiving expected signals. "
1487+
f"Expected: {expected_signals}, Got: {buffer[-200:]!r}"
1488+
)
1489+
buffer += chunk
1490+
except socket.timeout:
1491+
raise RuntimeError(
1492+
f"Timeout waiting for signals. "
1493+
f"Expected: {expected_signals}, Got: {buffer[-200:]!r}"
1494+
)
1495+
1496+
1497+
@contextmanager
1498+
def _managed_subprocess(args, timeout=SHORT_TIMEOUT):
1499+
"""
1500+
Context manager for subprocess lifecycle management.
1501+
1502+
Ensures process is properly terminated and cleaned up even on exceptions.
1503+
Uses graceful termination first, then forceful kill if needed.
1504+
"""
1505+
p = subprocess.Popen(args)
1506+
try:
1507+
yield p
1508+
finally:
1509+
try:
1510+
p.terminate()
1511+
try:
1512+
p.wait(timeout=timeout)
1513+
except subprocess.TimeoutExpired:
1514+
p.kill()
1515+
try:
1516+
p.wait(timeout=timeout)
1517+
except subprocess.TimeoutExpired:
1518+
pass # Process refuses to die, nothing more we can do
1519+
except OSError:
1520+
pass # Process already dead
1521+
1522+
1523+
def _cleanup_sockets(*sockets):
1524+
"""Safely close multiple sockets, ignoring errors."""
1525+
for sock in sockets:
1526+
if sock is not None:
1527+
try:
1528+
sock.close()
1529+
except OSError:
1530+
pass
1531+
1532+
1533+
1534+
class RemoteInspectionTestBase(unittest.TestCase):
1535+
@contextmanager
1536+
def _target_process(self, script_body):
1537+
"""Context manager for running a target process with socket sync."""
1538+
port = find_unused_port()
1539+
script = f"""\
1540+
import socket
1541+
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
1542+
sock.connect(('localhost', {port}))
1543+
{textwrap.dedent(script_body)}
1544+
"""
1545+
1546+
with os_helper.temp_dir() as work_dir:
1547+
script_dir = os.path.join(work_dir, "script_pkg")
1548+
os.mkdir(script_dir)
1549+
1550+
server_socket = _create_server_socket(port)
1551+
script_name = _make_test_script(script_dir, "script", script)
1552+
client_socket = None
1553+
1554+
try:
1555+
with _managed_subprocess([sys.executable, script_name]) as p:
1556+
client_socket, _ = server_socket.accept()
1557+
server_socket.close()
1558+
server_socket = None
1559+
1560+
def make_unwinder():
1561+
try:
1562+
return RemoteUnwinder(p.pid, all_threads=True)
1563+
except PermissionError:
1564+
self.skipTest("Insufficient permissions to read the stack trace")
1565+
1566+
yield p, client_socket, make_unwinder
1567+
finally:
1568+
_cleanup_sockets(client_socket, server_socket)
1569+
1570+
1571+
def _get_task_id_map(self, stack_trace):
1572+
"""Create task_id -> task mapping from async stack trace."""
1573+
return {task.task_id: task for task in stack_trace[0].awaited_by}
1574+
1575+
1576+
def _get_awaited_by_relationships(self, stack_trace):
1577+
"""Extract task name to awaited_by set mapping."""
1578+
id_to_task = self._get_task_id_map(stack_trace)
1579+
return {
1580+
task.task_name: set(
1581+
id_to_task[awaited.task_name].task_name
1582+
for awaited in task.awaited_by
1583+
)
1584+
for task in stack_trace[0].awaited_by
1585+
}
1586+
1587+
1588+
1589+
class TestFrameChainLimits(RemoteInspectionTestBase):
1590+
"""Frame chain walks abort instead of looping/overflowing on deep chains."""
1591+
1592+
# Limits plus one, to exceed them (must match MAX_FRAME_CHAIN_DEPTH /
1593+
# MAX_TASK_WAITER_WALK_TASKS from _remote_debugging_module.c)
1594+
FRAME_CHAIN_DEPTH = 1024 + 1
1595+
TASK_WAITER_WALK_TASKS = 2**14 + 1
1596+
1597+
def _assert_unwinder_limit_error(self, unwind, expected_substring):
1598+
"""Call unwind() until it raises the frame chain limit error.
1599+
1600+
unwind must construct the RemoteUnwinder and call it, so that
1601+
transient RuntimeErrors from either step are retried; a successful
1602+
call means the limit never triggered and fails immediately.
1603+
"""
1604+
last_error = None
1605+
for _ in busy_retry(SHORT_TIMEOUT, error=False):
1606+
try:
1607+
unwind()
1608+
except PermissionError:
1609+
self.skipTest("Insufficient permissions to read the stack trace")
1610+
except TRANSIENT_ERRORS as e:
1611+
if expected_substring in str(e):
1612+
return
1613+
last_error = e
1614+
continue
1615+
self.fail(
1616+
"frame chain limit did not trigger; call returned a result"
1617+
)
1618+
self.fail(
1619+
f"frame chain limit never raised; last transient error: "
1620+
f"{last_error!r}"
1621+
)
1622+
1623+
@skip_if_not_supported
1624+
@unittest.skipIf(
1625+
sys.platform == "linux" and not PROCESS_VM_READV_SUPPORTED,
1626+
"Test only runs on Linux with process_vm_readv support",
1627+
)
1628+
def test_get_stack_trace_deep_frame_chain_aborts(self):
1629+
"""Test that a frame chain deeper than the limit aborts the
1630+
synchronous stack walk instead of walking it indefinitely."""
1631+
script_body = f"""\
1632+
import sys
1633+
sys.setrecursionlimit({self.FRAME_CHAIN_DEPTH * 2})
1634+
1635+
def recurse(n):
1636+
if n <= 0:
1637+
sock.sendall(b"ready")
1638+
sock.recv(16)
1639+
return
1640+
recurse(n - 1)
1641+
1642+
recurse({self.FRAME_CHAIN_DEPTH})
1643+
"""
1644+
with self._target_process(script_body) as (p, client_socket, _):
1645+
_wait_for_signal(client_socket, b"ready")
1646+
self._assert_unwinder_limit_error(
1647+
lambda: RemoteUnwinder(p.pid).get_stack_trace(),
1648+
"Too many stack frames",
1649+
)
1650+
client_socket.sendall(b"done")
1651+
1652+
@skip_if_not_supported
1653+
@unittest.skipIf(
1654+
sys.platform == "linux" and not PROCESS_VM_READV_SUPPORTED,
1655+
"Test only runs on Linux with process_vm_readv support",
1656+
)
1657+
def test_get_async_stack_trace_deep_task_waiter_chain_aborts(self):
1658+
"""Test that a task waiter chain deeper than the limit aborts
1659+
the walk instead of overflowing the C stack."""
1660+
script_body = f"""\
1661+
import asyncio
1662+
1663+
async def chain(n):
1664+
if n <= 0:
1665+
sock.sendall(b"ready")
1666+
sock.recv(16)
1667+
return
1668+
1669+
task = asyncio.create_task(chain(n - 1))
1670+
await task
1671+
1672+
asyncio.run(chain({self.TASK_WAITER_WALK_TASKS}))
1673+
"""
1674+
with self._target_process(script_body) as (p, client_socket, _):
1675+
_wait_for_signal(client_socket, b"ready")
1676+
self._assert_unwinder_limit_error(
1677+
lambda: RemoteUnwinder(p.pid).get_async_stack_trace(),
1678+
"Too many task waiters",
1679+
)
1680+
client_socket.sendall(b"done")
1681+
1682+
@skip_if_not_supported
1683+
@unittest.skipIf(
1684+
sys.platform == "linux" and not PROCESS_VM_READV_SUPPORTED,
1685+
"Test only runs on Linux with process_vm_readv support",
1686+
)
1687+
def test_get_async_stack_trace_deep_frame_chain_aborts(self):
1688+
"""Test that a frame chain deeper than the limit aborts the async
1689+
stack walk instead of walking it indefinitely."""
1690+
script_body = f"""\
1691+
import sys, asyncio
1692+
sys.setrecursionlimit({self.FRAME_CHAIN_DEPTH * 2})
1693+
1694+
def recurse(n):
1695+
if n <= 0:
1696+
sock.sendall(b"ready")
1697+
sock.recv(16)
1698+
return
1699+
recurse(n - 1)
1700+
1701+
async def deep():
1702+
recurse({self.FRAME_CHAIN_DEPTH})
1703+
1704+
asyncio.run(deep())
1705+
"""
1706+
with self._target_process(script_body) as (p, client_socket, _):
1707+
_wait_for_signal(client_socket, b"ready")
1708+
self._assert_unwinder_limit_error(
1709+
lambda: RemoteUnwinder(p.pid).get_async_stack_trace(),
1710+
"Too many async stack frames",
1711+
)
1712+
client_socket.sendall(b"done")
1713+
1714+
@skip_if_not_supported
1715+
@unittest.skipIf(
1716+
sys.platform == "linux" and not PROCESS_VM_READV_SUPPORTED,
1717+
"Test only runs on Linux with process_vm_readv support",
1718+
)
1719+
def test_get_all_awaited_by_deep_coro_chain_aborts(self):
1720+
"""Test that a coroutine await chain deeper than the limit aborts
1721+
the walk instead of overflowing the C stack."""
1722+
script_body = f"""\
1723+
import sys, asyncio
1724+
sys.setrecursionlimit({self.FRAME_CHAIN_DEPTH * 2})
1725+
1726+
async def chain(n):
1727+
if n <= 0:
1728+
await asyncio.sleep(10_000)
1729+
return
1730+
await chain(n - 1)
1731+
1732+
async def main():
1733+
task = asyncio.create_task(chain({self.FRAME_CHAIN_DEPTH}))
1734+
await asyncio.sleep(0)
1735+
sock.sendall(b"ready")
1736+
await task
1737+
1738+
asyncio.run(main())
1739+
"""
1740+
with self._target_process(script_body) as (p, client_socket, _):
1741+
_wait_for_signal(client_socket, b"ready")
1742+
self._assert_unwinder_limit_error(
1743+
lambda: RemoteUnwinder(p.pid).get_all_awaited_by(),
1744+
"Too many coroutine frames",
1745+
)
1746+
1747+
1748+
@skip_if_not_supported
1749+
@unittest.skipIf(
1750+
sys.platform == "linux" and not PROCESS_VM_READV_SUPPORTED,
1751+
"Test only runs on Linux with process_vm_readv support",
1752+
)
1753+
def test_async_awaited_by_skips_set_tombstones(self):
1754+
script_body = """\
1755+
import asyncio
1756+
1757+
class RemovedTask(asyncio.Task):
1758+
def __hash__(self):
1759+
return 0
1760+
1761+
class RemainingTask(asyncio.Task):
1762+
def __hash__(self):
1763+
return 1
1764+
1765+
async def main():
1766+
victim = asyncio.current_task()
1767+
victim.set_name("victim")
1768+
removed = RemovedTask(
1769+
asyncio.sleep(10_000), name="removed"
1770+
)
1771+
remaining = RemainingTask(
1772+
asyncio.sleep(10_000), name="remaining"
1773+
)
1774+
1775+
asyncio.future_add_to_awaited_by(victim, removed)
1776+
asyncio.future_add_to_awaited_by(victim, remaining)
1777+
1778+
# Removing hash 0 leaves a dummy in slot 0 before the only
1779+
# active entry in slot 1. It must not count toward the set's
1780+
# used entries.
1781+
asyncio.future_discard_from_awaited_by(victim, removed)
1782+
1783+
sock.sendall(b"ready")
1784+
sock.recv(16)
1785+
1786+
asyncio.run(main())
1787+
"""
1788+
1789+
with self._target_process(script_body) as (
1790+
_,
1791+
client_socket,
1792+
make_unwinder,
1793+
):
1794+
_wait_for_signal(client_socket, b"ready")
1795+
1796+
for method_name in (
1797+
"get_async_stack_trace",
1798+
"get_all_awaited_by",
1799+
):
1800+
with self.subTest(method=method_name):
1801+
unwinder = make_unwinder()
1802+
stack_trace = getattr(unwinder, method_name)()
1803+
relationships = self._get_awaited_by_relationships(
1804+
stack_trace
1805+
)
1806+
self.assertEqual(
1807+
relationships["victim"],
1808+
{"remaining"},
1809+
)
1810+
1811+
client_socket.sendall(b"done")
1812+
1813+
1814+
14401815
if __name__ == "__main__":
14411816
unittest.main()
Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,3 @@
1+
Make frame, coroutine and task-waiter walks iterative and bounded, avoiding
2+
potential hangs and stack overflows. Fix asyncio task inspection when
3+
awaited-by sets contain removed entries. Patch by Maurycy Pawłowski-Wieroński.

0 commit comments

Comments
 (0)