diff --git a/pymongo/asynchronous/pool.py b/pymongo/asynchronous/pool.py index 69ac14e7e4..b7f7a1f6a0 100644 --- a/pymongo/asynchronous/pool.py +++ b/pymongo/asynchronous/pool.py @@ -106,6 +106,13 @@ _IS_SYNC = False +# Flags recording the counters a checkout has incremented, so a failed +# checkout can restore them exactly once (PYTHON-6136). +_UNDO_OPERATION_COUNT = 1 +_UNDO_REQUESTS = 2 +_UNDO_SOCKETS = 4 +_UNDO_PENDING = 8 + class AsyncConnection(_ConnectionTelemetryInfo): """Store a connection with some metadata. @@ -969,8 +976,12 @@ async def _get_conn( "Attempted to check out a connection from closed connection pool" ) + # Every counter increment sets a flag in ``applied`` so a failed + # checkout can restore them exactly once (PYTHON-6136). + applied = 0 async with self.lock: self.operation_count += 1 + applied |= _UNDO_OPERATION_COUNT # Get a free socket or create one. if _csot.get_timeout(): @@ -980,28 +991,32 @@ async def _get_conn( else: deadline = None - async with self.size_cond: - self._raise_if_not_ready(checkout_started_time, emit_event=True) - while not (self.requests < self.max_pool_size): - timeout = deadline - time.monotonic() if deadline else None - if not await _async_cond_wait(self.size_cond, timeout): - # Timed out, notify the next thread to ensure a - # timeout doesn't consume the condition. - if self.requests < self.max_pool_size: - self.size_cond.notify() - self._raise_wait_queue_timeout(checkout_started_time) + try: + async with self.size_cond: self._raise_if_not_ready(checkout_started_time, emit_event=True) - self.requests += 1 + while not (self.requests < self.max_pool_size): + timeout = deadline - time.monotonic() if deadline else None + if not await _async_cond_wait(self.size_cond, timeout): + # Timed out, notify the next thread to ensure a + # timeout doesn't consume the condition. + if self.requests < self.max_pool_size: + self.size_cond.notify() + self._raise_wait_queue_timeout(checkout_started_time) + self._raise_if_not_ready(checkout_started_time, emit_event=True) + self.requests += 1 + applied |= _UNDO_REQUESTS + except BaseException: + await self._restore_counters(applied) + raise # We've now acquired the semaphore and must release it on error. conn = None - incremented = False emitted_event = False is_new_conn = False try: async with self.lock: self.active_sockets += 1 - incremented = True + applied |= _UNDO_SOCKETS while conn is None: # CMAP: we MUST wait for either maxConnecting OR for a socket # to be checked back into the pool. @@ -1022,6 +1037,7 @@ async def _get_conn( conn = self.conns.popleft() except IndexError: self._pending += 1 + applied |= _UNDO_PENDING if conn: # We got a socket from the pool if await self._perished(conn): conn = None @@ -1033,7 +1049,17 @@ async def _get_conn( finally: async with self._max_connecting_cond: self._pending -= 1 - self._max_connecting_cond.notify() + applied &= ~_UNDO_PENDING + notified = False + try: + self._max_connecting_cond.notify() + notified = True + finally: + if not notified: + # A kill landed inside notify() (a gevent + # yield point); retry so a waiting + # checkout is not stranded (PYTHON-6136). + self._max_connecting_cond.notify() conn.active = True # connect() already adds cancel_context for new connections; only add @@ -1043,27 +1069,14 @@ async def _get_conn( self.active_contexts.add(conn.cancel_context) # Catch KeyboardInterrupt, CancelledError, etc. and cleanup. except BaseException: - if conn: - # We checked out a socket but authentication failed. - await conn.close_conn(ConnectionClosedReason.ERROR) - # Re-apply the accounting if a GreenletExit interrupts - # during the size_cond acquisition; during unwind gevent - # lets the re-acquire complete (PYTHON-6074). - accounted = False try: - async with self.size_cond: - self.requests -= 1 - if incremented: - self.active_sockets -= 1 - accounted = True - self.size_cond.notify() + if conn is not None: + # We checked out a socket but authentication failed. + await conn.close_conn(ConnectionClosedReason.ERROR) finally: - if not accounted: - async with self.size_cond: - self.requests -= 1 - if incremented: - self.active_sockets -= 1 - self.size_cond.notify() + # Restore the counters even if the cleanup above was + # interrupted (PYTHON-6136). + await self._restore_counters(applied) if not emitted_event: self._telemetry.checkout_failed( @@ -1075,6 +1088,51 @@ async def _get_conn( return conn + def _restore_applied(self, applied: int) -> None: + """Restore the counters flagged in ``applied``. Caller holds ``size_cond``.""" + if applied & _UNDO_OPERATION_COUNT: + self.operation_count -= 1 + if applied & _UNDO_REQUESTS: + self.requests -= 1 + if applied & _UNDO_SOCKETS: + self.active_sockets -= 1 + if applied & _UNDO_PENDING: + self._pending -= 1 + + async def _restore_counters(self, applied: int) -> None: + """Restore the counters a failed checkout incremented (PYTHON-6136). + + Gevent grants the re-acquire during unwind (PYTHON-6074). A kill can + also land inside notify() (a yield point), so notifications are + tracked and retried separately from the counter restore. + """ + accounted = False + notified = 0 + try: + async with self.size_cond: + self._restore_applied(applied) + accounted = True + if applied & _UNDO_REQUESTS: + # A pool slot was freed; wake the next waiting thread. + self.size_cond.notify() + notified |= _UNDO_REQUESTS + if applied & _UNDO_PENDING: + # A maxConnecting slot was freed; wake the next waiting thread. + self._max_connecting_cond.notify() + notified |= _UNDO_PENDING + finally: + # Always reacquired: `applied` keeps restore-only flags + # `notified` never tracks, so skipping needs a mask synced + # to the notify sites; uncontended acquires don't yield (PYTHON-6136). + async with self.size_cond: + if not accounted: + self._restore_applied(applied) + missing = applied & ~notified + if missing & _UNDO_REQUESTS: + self.size_cond.notify() + if missing & _UNDO_PENDING: + self._max_connecting_cond.notify() + def _checkin_apply( self, conn: AsyncConnection, txn: bool, cursor: bool, forked: bool ) -> tuple[Optional[str], bool, bool]: @@ -1121,9 +1179,7 @@ async def checkin(self, conn: AsyncConnection) -> None: conn.pinned_cursor = False self._pinned_sockets.discard(conn) forked = self.pid != os.getpid() - # Re-apply the accounting if a gevent GreenletExit interrupts during - # the size_cond acquisition; gevent lets the re-acquire complete while - # unwinding (PYTHON-6074). + # Re-apply the accounting if a BaseException interrupts here (PYTHON-6074). close_conn_reason: Optional[str] = None emit_closed = False accounted = False diff --git a/pymongo/synchronous/pool.py b/pymongo/synchronous/pool.py index c977729ba8..dd562ac68e 100644 --- a/pymongo/synchronous/pool.py +++ b/pymongo/synchronous/pool.py @@ -106,6 +106,13 @@ _IS_SYNC = True +# Flags recording the counters a checkout has incremented, so a failed +# checkout can restore them exactly once (PYTHON-6136). +_UNDO_OPERATION_COUNT = 1 +_UNDO_REQUESTS = 2 +_UNDO_SOCKETS = 4 +_UNDO_PENDING = 8 + class Connection(_ConnectionTelemetryInfo): """Store a connection with some metadata. @@ -965,8 +972,12 @@ def _get_conn( "Attempted to check out a connection from closed connection pool" ) + # Every counter increment sets a flag in ``applied`` so a failed + # checkout can restore them exactly once (PYTHON-6136). + applied = 0 with self.lock: self.operation_count += 1 + applied |= _UNDO_OPERATION_COUNT # Get a free socket or create one. if _csot.get_timeout(): @@ -976,28 +987,32 @@ def _get_conn( else: deadline = None - with self.size_cond: - self._raise_if_not_ready(checkout_started_time, emit_event=True) - while not (self.requests < self.max_pool_size): - timeout = deadline - time.monotonic() if deadline else None - if not _cond_wait(self.size_cond, timeout): - # Timed out, notify the next thread to ensure a - # timeout doesn't consume the condition. - if self.requests < self.max_pool_size: - self.size_cond.notify() - self._raise_wait_queue_timeout(checkout_started_time) + try: + with self.size_cond: self._raise_if_not_ready(checkout_started_time, emit_event=True) - self.requests += 1 + while not (self.requests < self.max_pool_size): + timeout = deadline - time.monotonic() if deadline else None + if not _cond_wait(self.size_cond, timeout): + # Timed out, notify the next thread to ensure a + # timeout doesn't consume the condition. + if self.requests < self.max_pool_size: + self.size_cond.notify() + self._raise_wait_queue_timeout(checkout_started_time) + self._raise_if_not_ready(checkout_started_time, emit_event=True) + self.requests += 1 + applied |= _UNDO_REQUESTS + except BaseException: + self._restore_counters(applied) + raise # We've now acquired the semaphore and must release it on error. conn = None - incremented = False emitted_event = False is_new_conn = False try: with self.lock: self.active_sockets += 1 - incremented = True + applied |= _UNDO_SOCKETS while conn is None: # CMAP: we MUST wait for either maxConnecting OR for a socket # to be checked back into the pool. @@ -1018,6 +1033,7 @@ def _get_conn( conn = self.conns.popleft() except IndexError: self._pending += 1 + applied |= _UNDO_PENDING if conn: # We got a socket from the pool if self._perished(conn): conn = None @@ -1029,7 +1045,17 @@ def _get_conn( finally: with self._max_connecting_cond: self._pending -= 1 - self._max_connecting_cond.notify() + applied &= ~_UNDO_PENDING + notified = False + try: + self._max_connecting_cond.notify() + notified = True + finally: + if not notified: + # A kill landed inside notify() (a gevent + # yield point); retry so a waiting + # checkout is not stranded (PYTHON-6136). + self._max_connecting_cond.notify() conn.active = True # connect() already adds cancel_context for new connections; only add @@ -1039,27 +1065,14 @@ def _get_conn( self.active_contexts.add(conn.cancel_context) # Catch KeyboardInterrupt, CancelledError, etc. and cleanup. except BaseException: - if conn: - # We checked out a socket but authentication failed. - conn.close_conn(ConnectionClosedReason.ERROR) - # Re-apply the accounting if a GreenletExit interrupts - # during the size_cond acquisition; during unwind gevent - # lets the re-acquire complete (PYTHON-6074). - accounted = False try: - with self.size_cond: - self.requests -= 1 - if incremented: - self.active_sockets -= 1 - accounted = True - self.size_cond.notify() + if conn is not None: + # We checked out a socket but authentication failed. + conn.close_conn(ConnectionClosedReason.ERROR) finally: - if not accounted: - with self.size_cond: - self.requests -= 1 - if incremented: - self.active_sockets -= 1 - self.size_cond.notify() + # Restore the counters even if the cleanup above was + # interrupted (PYTHON-6136). + self._restore_counters(applied) if not emitted_event: self._telemetry.checkout_failed( @@ -1071,6 +1084,51 @@ def _get_conn( return conn + def _restore_applied(self, applied: int) -> None: + """Restore the counters flagged in ``applied``. Caller holds ``size_cond``.""" + if applied & _UNDO_OPERATION_COUNT: + self.operation_count -= 1 + if applied & _UNDO_REQUESTS: + self.requests -= 1 + if applied & _UNDO_SOCKETS: + self.active_sockets -= 1 + if applied & _UNDO_PENDING: + self._pending -= 1 + + def _restore_counters(self, applied: int) -> None: + """Restore the counters a failed checkout incremented (PYTHON-6136). + + Gevent grants the re-acquire during unwind (PYTHON-6074). A kill can + also land inside notify() (a yield point), so notifications are + tracked and retried separately from the counter restore. + """ + accounted = False + notified = 0 + try: + with self.size_cond: + self._restore_applied(applied) + accounted = True + if applied & _UNDO_REQUESTS: + # A pool slot was freed; wake the next waiting thread. + self.size_cond.notify() + notified |= _UNDO_REQUESTS + if applied & _UNDO_PENDING: + # A maxConnecting slot was freed; wake the next waiting thread. + self._max_connecting_cond.notify() + notified |= _UNDO_PENDING + finally: + # Always reacquired: `applied` keeps restore-only flags + # `notified` never tracks, so skipping needs a mask synced + # to the notify sites; uncontended acquires don't yield (PYTHON-6136). + with self.size_cond: + if not accounted: + self._restore_applied(applied) + missing = applied & ~notified + if missing & _UNDO_REQUESTS: + self.size_cond.notify() + if missing & _UNDO_PENDING: + self._max_connecting_cond.notify() + def _checkin_apply( self, conn: Connection, txn: bool, cursor: bool, forked: bool ) -> tuple[Optional[str], bool, bool]: @@ -1117,9 +1175,7 @@ def checkin(self, conn: Connection) -> None: conn.pinned_cursor = False self._pinned_sockets.discard(conn) forked = self.pid != os.getpid() - # Re-apply the accounting if a gevent GreenletExit interrupts during - # the size_cond acquisition; gevent lets the re-acquire complete while - # unwinding (PYTHON-6074). + # Re-apply the accounting if a BaseException interrupts here (PYTHON-6074). close_conn_reason: Optional[str] = None emit_closed = False accounted = False diff --git a/test/asynchronous/test_client.py b/test/asynchronous/test_client.py index 56810b1433..563d29cbad 100644 --- a/test/asynchronous/test_client.py +++ b/test/asynchronous/test_client.py @@ -2744,20 +2744,18 @@ def test_gevent_kill_churn_deadlock(self): import gevent.thread as _gthread from gevent import Timeout, spawn - # AMPLIFY_RACE=1 widens gevent's brief sleep on a contended lock - # so a kill lands there reliably. Only the bare sleep() is - # widened; timed sleeps pass through. The unfixed test then - # deadlocks within seconds. + # AMPLIFY_RACE=1 widens the kill windows below; keep AMPLIFY_SECONDS + # small or ops starve. + amplify_seconds = 0.0 if os.environ.get("AMPLIFY_RACE", "0") == "1": - _AMPLIFY_SECONDS = float(os.environ.get("AMPLIFY_SECONDS", "0.02")) + amplify_seconds = float(os.environ.get("AMPLIFY_SECONDS", "0.002")) _orig_thread_sleep = _gthread.sleep def _amplified_sleep(*args): - if not args: # bare sleep(): the courtesy yield on a failed - # non-blocking acquire (Condition.notify -> _is_owned -> acquire(False)) - _orig_thread_sleep(_AMPLIFY_SECONDS) - else: # sleep(0.001), sleep(2), etc.: passthrough + if not args: # bare sleep(): checkin courtesy yield + _orig_thread_sleep(amplify_seconds) + else: # timed sleeps pass through _orig_thread_sleep(*args) _gthread.sleep = _amplified_sleep @@ -2767,6 +2765,31 @@ def _amplified_sleep(*args): coll = client.pymongo_test.coll coll.insert_one({}) + pool = async_get_pool(client) # type:ignore + # Widen the post-gate checkout windows (PYTHON-6136). + if amplify_seconds: + + class _AmplifiedCondition: + def __init__(self, cond, seconds): + self._cond = cond + self._seconds = seconds + + def __enter__(self): + self._cond.__enter__() + return self + + def __exit__(self, *args): + self._cond.__exit__(*args) + time.sleep(self._seconds) + + def __getattr__(self, name): + return getattr(self._cond, name) + + pool.size_cond = _AmplifiedCondition(pool.size_cond, amplify_seconds) + pool._max_connecting_cond = _AmplifiedCondition( + pool._max_connecting_cond, amplify_seconds + ) + op_count = [0] running = [True] workers: list = [] @@ -2780,9 +2803,12 @@ def worker(): except Exception: return + # Scale the kill cadence or ops starve. + reaper_interval = max(0.003, amplify_seconds * 8) + def reaper(): while running[0]: - time.sleep(0.003) + time.sleep(reaper_interval) if not workers: continue idx = random.randrange(len(workers)) @@ -2823,11 +2849,6 @@ def reaper(): coll.find_one({}) except Timeout: self.fail("Pool gate saturated (PYTHON-6074)") - # Deterministic check: a saturated size gate pins the pool's - # checkout counters at maxPoolSize (PYTHON-6074). - pool = async_get_pool(client) # type:ignore - self.assertLess(pool.requests, pool.max_pool_size) - self.assertLess(pool.active_sockets, pool.max_pool_size) self.assertGreater(op_count[0], 0) finally: running[0] = False @@ -2845,6 +2866,14 @@ def reaper(): client.close() except Timeout: pass + # Post-settle: the counters must have fully drained. Note close() + # does not zero operation_count (only a fork does), so a leaked + # increment survives and is caught here. + time.sleep(1.0) + self.assertEqual(pool.requests, 0) + self.assertEqual(pool._pending, 0) + self.assertEqual(pool.active_sockets, 0) + self.assertEqual(pool.operation_count, 0) class TestClientLazyConnect(AsyncIntegrationTest): diff --git a/test/asynchronous/test_pooling.py b/test/asynchronous/test_pooling.py index 661bd4e3d2..25da5134b9 100644 --- a/test/asynchronous/test_pooling.py +++ b/test/asynchronous/test_pooling.py @@ -306,6 +306,7 @@ def notify(): # Accounting was applied exactly once. self.assertEqual(0, cx_pool.requests) self.assertEqual(0, cx_pool.active_sockets) + self.assertEqual(0, cx_pool.operation_count) async def test_checkout_error_accounting_on_kill_during_acquire(self): # PYTHON-6074: an exception delivered while the checkout error @@ -341,6 +342,98 @@ async def __aexit__(self, *args): # The fallback applied the accounting exactly once. self.assertEqual(0, cx_pool.requests) self.assertEqual(0, cx_pool.active_sockets) + self.assertEqual(0, cx_pool.operation_count) + + async def test_checkout_error_accounting_on_connect_keyboard_interrupt(self): + # PYTHON-6136: a KeyboardInterrupt from connect() must roll back the + # size gate, the maxConnecting gate, and active_sockets. + cx_pool = await self.create_pool(max_pool_size=1) + + with patch.object(cx_pool, "connect", side_effect=KeyboardInterrupt()): + with self.assertRaises(KeyboardInterrupt): + async with cx_pool.checkout(): + pass + + self.assertEqual(0, cx_pool.requests) + self.assertEqual(0, cx_pool.active_sockets) + self.assertEqual(0, cx_pool._pending) + self.assertEqual(0, cx_pool.operation_count) + + async def test_checkout_error_accounting_on_kill_during_pending_cleanup(self): + # PYTHON-6136: an interruption during the pending-gate cleanup must + # not skip the counter restore, which must wake threads waiting at + # the maxConnecting gate. + cx_pool = await self.create_pool(max_pool_size=1) + + class _InterruptOnSecondEnter(type(cx_pool._max_connecting_cond)): + def __init__(self, lock): + super().__init__(lock) + self.enters = 0 + self.notifies = 0 + + async def __aenter__(self): + self.enters += 1 + if self.enters == 2: + # First enter is the checkout wait, second is connect()'s + # cleanup. Simulate a kill delivered while waiting there. + raise KeyboardInterrupt() + return await super().__aenter__() + + async def __aexit__(self, *args): + return await super().__aexit__(*args) + + def notify(self, n=1): + # The counter restore wakes the maxConnecting gate, then is + # itself killed. + self.notifies += 1 + raise KeyboardInterrupt() + + cond = _InterruptOnSecondEnter(cx_pool._max_connecting_cond._lock) + cx_pool._max_connecting_cond = cond + + with patch.object(cx_pool, "connect", side_effect=asyncio.CancelledError()): + with self.assertRaises(KeyboardInterrupt): + async with cx_pool.checkout(): + pass + + self.assertEqual(0, cx_pool.requests) + self.assertEqual(0, cx_pool.active_sockets) + self.assertEqual(0, cx_pool._pending) + self.assertEqual(0, cx_pool.operation_count) + # The restore notified the maxConnecting gate; the notify interrupted + # by the kill was retried. + self.assertEqual(2, cond.notifies) + + async def test_checkout_error_accounting_on_kill_during_pending_notify(self): + # PYTHON-6136: a kill landing inside the cleanup notify must not + # strand a checkout waiting at the maxConnecting gate. + cx_pool = await self.create_pool(max_pool_size=1) + + class _InterruptOnFirstNotify(type(cx_pool._max_connecting_cond)): + def __init__(self, lock): + super().__init__(lock) + self.notifies = 0 + + def notify(self, n=1): + self.notifies += 1 + if self.notifies == 1: + # Simulate a kill delivered inside notify(). + raise KeyboardInterrupt() + + cond = _InterruptOnFirstNotify(cx_pool._max_connecting_cond._lock) + cx_pool._max_connecting_cond = cond + + with patch.object(cx_pool, "connect", side_effect=asyncio.CancelledError()): + with self.assertRaises(KeyboardInterrupt): + async with cx_pool.checkout(): + pass + + # The cleanup notify was interrupted, then retried. + self.assertEqual(2, cond.notifies) + self.assertEqual(0, cx_pool.requests) + self.assertEqual(0, cx_pool.active_sockets) + self.assertEqual(0, cx_pool._pending) + self.assertEqual(0, cx_pool.operation_count) async def test_pool_removes_closed_socket(self): # Test that Pool removes explicitly closed socket. @@ -463,6 +556,8 @@ async def test_wait_queue_timeout(self): 1, f"Waited {duration:.2f} seconds for a socket, expected {wait_queue_timeout:f}", ) + # The load metric must not be inflated by the failed checkout. + self.assertEqual(0, pool.operation_count) async def test_no_wait_queue_timeout(self): # Verify get_socket() with no wait_queue_timeout blocks forever. diff --git a/test/test_client.py b/test/test_client.py index ebd1e670ee..0c8d0e7595 100644 --- a/test/test_client.py +++ b/test/test_client.py @@ -2695,20 +2695,18 @@ def test_gevent_kill_churn_deadlock(self): import gevent.thread as _gthread from gevent import Timeout, spawn - # AMPLIFY_RACE=1 widens gevent's brief sleep on a contended lock - # so a kill lands there reliably. Only the bare sleep() is - # widened; timed sleeps pass through. The unfixed test then - # deadlocks within seconds. + # AMPLIFY_RACE=1 widens the kill windows below; keep AMPLIFY_SECONDS + # small or ops starve. + amplify_seconds = 0.0 if os.environ.get("AMPLIFY_RACE", "0") == "1": - _AMPLIFY_SECONDS = float(os.environ.get("AMPLIFY_SECONDS", "0.02")) + amplify_seconds = float(os.environ.get("AMPLIFY_SECONDS", "0.002")) _orig_thread_sleep = _gthread.sleep def _amplified_sleep(*args): - if not args: # bare sleep(): the courtesy yield on a failed - # non-blocking acquire (Condition.notify -> _is_owned -> acquire(False)) - _orig_thread_sleep(_AMPLIFY_SECONDS) - else: # sleep(0.001), sleep(2), etc.: passthrough + if not args: # bare sleep(): checkin courtesy yield + _orig_thread_sleep(amplify_seconds) + else: # timed sleeps pass through _orig_thread_sleep(*args) _gthread.sleep = _amplified_sleep @@ -2718,6 +2716,31 @@ def _amplified_sleep(*args): coll = client.pymongo_test.coll coll.insert_one({}) + pool = get_pool(client) # type:ignore + # Widen the post-gate checkout windows (PYTHON-6136). + if amplify_seconds: + + class _AmplifiedCondition: + def __init__(self, cond, seconds): + self._cond = cond + self._seconds = seconds + + def __enter__(self): + self._cond.__enter__() + return self + + def __exit__(self, *args): + self._cond.__exit__(*args) + time.sleep(self._seconds) + + def __getattr__(self, name): + return getattr(self._cond, name) + + pool.size_cond = _AmplifiedCondition(pool.size_cond, amplify_seconds) + pool._max_connecting_cond = _AmplifiedCondition( + pool._max_connecting_cond, amplify_seconds + ) + op_count = [0] running = [True] workers: list = [] @@ -2731,9 +2754,12 @@ def worker(): except Exception: return + # Scale the kill cadence or ops starve. + reaper_interval = max(0.003, amplify_seconds * 8) + def reaper(): while running[0]: - time.sleep(0.003) + time.sleep(reaper_interval) if not workers: continue idx = random.randrange(len(workers)) @@ -2774,11 +2800,6 @@ def reaper(): coll.find_one({}) except Timeout: self.fail("Pool gate saturated (PYTHON-6074)") - # Deterministic check: a saturated size gate pins the pool's - # checkout counters at maxPoolSize (PYTHON-6074). - pool = get_pool(client) # type:ignore - self.assertLess(pool.requests, pool.max_pool_size) - self.assertLess(pool.active_sockets, pool.max_pool_size) self.assertGreater(op_count[0], 0) finally: running[0] = False @@ -2796,6 +2817,14 @@ def reaper(): client.close() except Timeout: pass + # Post-settle: the counters must have fully drained. Note close() + # does not zero operation_count (only a fork does), so a leaked + # increment survives and is caught here. + time.sleep(1.0) + self.assertEqual(pool.requests, 0) + self.assertEqual(pool._pending, 0) + self.assertEqual(pool.active_sockets, 0) + self.assertEqual(pool.operation_count, 0) class TestClientLazyConnect(IntegrationTest): diff --git a/test/test_pooling.py b/test/test_pooling.py index a3f0eaf589..3a81bc796f 100644 --- a/test/test_pooling.py +++ b/test/test_pooling.py @@ -306,6 +306,7 @@ def notify(): # Accounting was applied exactly once. self.assertEqual(0, cx_pool.requests) self.assertEqual(0, cx_pool.active_sockets) + self.assertEqual(0, cx_pool.operation_count) def test_checkout_error_accounting_on_kill_during_acquire(self): # PYTHON-6074: an exception delivered while the checkout error @@ -341,6 +342,98 @@ def __exit__(self, *args): # The fallback applied the accounting exactly once. self.assertEqual(0, cx_pool.requests) self.assertEqual(0, cx_pool.active_sockets) + self.assertEqual(0, cx_pool.operation_count) + + def test_checkout_error_accounting_on_connect_keyboard_interrupt(self): + # PYTHON-6136: a KeyboardInterrupt from connect() must roll back the + # size gate, the maxConnecting gate, and active_sockets. + cx_pool = self.create_pool(max_pool_size=1) + + with patch.object(cx_pool, "connect", side_effect=KeyboardInterrupt()): + with self.assertRaises(KeyboardInterrupt): + with cx_pool.checkout(): + pass + + self.assertEqual(0, cx_pool.requests) + self.assertEqual(0, cx_pool.active_sockets) + self.assertEqual(0, cx_pool._pending) + self.assertEqual(0, cx_pool.operation_count) + + def test_checkout_error_accounting_on_kill_during_pending_cleanup(self): + # PYTHON-6136: an interruption during the pending-gate cleanup must + # not skip the counter restore, which must wake threads waiting at + # the maxConnecting gate. + cx_pool = self.create_pool(max_pool_size=1) + + class _InterruptOnSecondEnter(type(cx_pool._max_connecting_cond)): + def __init__(self, lock): + super().__init__(lock) + self.enters = 0 + self.notifies = 0 + + def __enter__(self): + self.enters += 1 + if self.enters == 2: + # First enter is the checkout wait, second is connect()'s + # cleanup. Simulate a kill delivered while waiting there. + raise KeyboardInterrupt() + return super().__enter__() + + def __exit__(self, *args): + return super().__exit__(*args) + + def notify(self, n=1): + # The counter restore wakes the maxConnecting gate, then is + # itself killed. + self.notifies += 1 + raise KeyboardInterrupt() + + cond = _InterruptOnSecondEnter(cx_pool._max_connecting_cond._lock) + cx_pool._max_connecting_cond = cond + + with patch.object(cx_pool, "connect", side_effect=asyncio.CancelledError()): + with self.assertRaises(KeyboardInterrupt): + with cx_pool.checkout(): + pass + + self.assertEqual(0, cx_pool.requests) + self.assertEqual(0, cx_pool.active_sockets) + self.assertEqual(0, cx_pool._pending) + self.assertEqual(0, cx_pool.operation_count) + # The restore notified the maxConnecting gate; the notify interrupted + # by the kill was retried. + self.assertEqual(2, cond.notifies) + + def test_checkout_error_accounting_on_kill_during_pending_notify(self): + # PYTHON-6136: a kill landing inside the cleanup notify must not + # strand a checkout waiting at the maxConnecting gate. + cx_pool = self.create_pool(max_pool_size=1) + + class _InterruptOnFirstNotify(type(cx_pool._max_connecting_cond)): + def __init__(self, lock): + super().__init__(lock) + self.notifies = 0 + + def notify(self, n=1): + self.notifies += 1 + if self.notifies == 1: + # Simulate a kill delivered inside notify(). + raise KeyboardInterrupt() + + cond = _InterruptOnFirstNotify(cx_pool._max_connecting_cond._lock) + cx_pool._max_connecting_cond = cond + + with patch.object(cx_pool, "connect", side_effect=asyncio.CancelledError()): + with self.assertRaises(KeyboardInterrupt): + with cx_pool.checkout(): + pass + + # The cleanup notify was interrupted, then retried. + self.assertEqual(2, cond.notifies) + self.assertEqual(0, cx_pool.requests) + self.assertEqual(0, cx_pool.active_sockets) + self.assertEqual(0, cx_pool._pending) + self.assertEqual(0, cx_pool.operation_count) def test_pool_removes_closed_socket(self): # Test that Pool removes explicitly closed socket. @@ -463,6 +556,8 @@ def test_wait_queue_timeout(self): 1, f"Waited {duration:.2f} seconds for a socket, expected {wait_queue_timeout:f}", ) + # The load metric must not be inflated by the failed checkout. + self.assertEqual(0, pool.operation_count) def test_no_wait_queue_timeout(self): # Verify get_socket() with no wait_queue_timeout blocks forever.