|
1 | 1 | import unittest |
| 2 | +from contextlib import contextmanager |
2 | 3 | import asyncio |
3 | 4 | import os |
4 | 5 | import textwrap |
@@ -1437,5 +1438,379 @@ def lines(u, expected_count): |
1437 | 1438 | ) |
1438 | 1439 |
|
1439 | 1440 |
|
| 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 | + |
1440 | 1815 | if __name__ == "__main__": |
1441 | 1816 | unittest.main() |
0 commit comments