diff --git a/src/schematic/client.py b/src/schematic/client.py index 45d0bcf..63f6957 100644 --- a/src/schematic/client.py +++ b/src/schematic/client.py @@ -1651,6 +1651,12 @@ async def _prewarm_one(self, company_id: str, credit_type_id: str) -> None: manager = self._lease_manager if manager is None: return + if self._is_shutting_down: + # shutdown() only cancels the prewarms it spawned; a caller + # awaiting prewarm() directly would otherwise install a lease + # after the release has already listed the store. + self.logger.debug("prewarm: client is shutting down, skipping acquire") + return try: await manager.acquire_if_needed(company_id, credit_type_id) except Exception as e: @@ -2035,6 +2041,18 @@ async def shutdown(self) -> None: try: if self._lease_manager is not None: self._lease_manager.stop() + # A prewarm is worthless to a process that is exiting, and + # waiting out its company-resolve poll would stall shutdown for + # seconds. Cancel it, then drain what it already put on the + # wire, so a lease installed mid-shutdown is one + # release_all_local_leases() can see. Both run for a shared + # backend too: the work must not outlive the client. + pending = list(self._background_tasks) + for task in pending: + task.cancel() + if pending: + await asyncio.gather(*pending, return_exceptions=True) + await self._lease_manager.drain() if not self._lease_backend_shared: # Per-process leases have no sibling drawing on them, so # releasing hands the unspent remainder back to the company diff --git a/src/schematic/leases/lease_manager.py b/src/schematic/leases/lease_manager.py index 141b93b..b3e088e 100644 --- a/src/schematic/leases/lease_manager.py +++ b/src/schematic/leases/lease_manager.py @@ -19,6 +19,7 @@ from .reservation_store import ReservationStore from .types import ( DEFAULT_SWEEP_INTERVAL, + SHUTDOWN_DRAIN_TIMEOUT, Clock, LeaseConfig, LeaseState, @@ -147,7 +148,10 @@ def __init__( # or the other way round. self._inflight_acquire: Dict[str, "asyncio.Future[Optional[LeaseState]]"] = {} self._inflight_extend: Dict[str, "asyncio.Future[Optional[LeaseState]]"] = {} - self._background: Set["asyncio.Task[None]"] = set() + # Every task shutdown has to wait out, whatever it resolves to: the + # fire-and-forget work from `_spawn` and the single-flight acquires and + # extends, which resolve to a LeaseState. + self._background: Set["asyncio.Task[Any]"] = set() self._sweep_task: Optional["asyncio.Task[None]"] = None self._stopped = False @@ -392,6 +396,13 @@ async def _single_flight( ) -> Optional[LeaseState]: task = asyncio.ensure_future(coro) registry[key] = task + # The registry dedupes concurrent callers and the drain set waits the + # wire call out; they have different lifetimes. Cancelling a caller + # cancels its `shield`, not the task, and drops the registry entry the + # instant it lands, so without this the acquire would be tracked + # nowhere and could install a lease after shutdown released the store. + self._background.add(task) + task.add_done_callback(self._background.discard) try: return await asyncio.shield(task) finally: @@ -406,6 +417,13 @@ async def _release(self, lease_id: str) -> None: def _spawn(self, coro: Awaitable[None]) -> None: """Run a fire-and-forget step, holding a reference so it is not collected.""" + if self._stopped: + # Past stop() the drain has run or is running; work started now + # would install or extend a lease nothing is left to release. + logger.debug("Lease manager is stopped; skipping background lease work") + if asyncio.iscoroutine(coro): + coro.close() + return try: task = asyncio.get_running_loop().create_task(_never_raises(coro)) except RuntimeError: @@ -419,6 +437,22 @@ async def _drain_background(self) -> None: while self._background: await asyncio.gather(*list(self._background), return_exceptions=True) + async def drain(self) -> None: + """Wait out in-flight lease work, so a close can release what it installed. + + Bounded: whatever has not landed by ``SHUTDOWN_DRAIN_TIMEOUT`` is + cancelled rather than stalling the caller's shutdown, and a grant the + server issued for it falls back to server-side expiry. + """ + try: + await asyncio.wait_for(self._drain_background(), SHUTDOWN_DRAIN_TIMEOUT) + except asyncio.TimeoutError: + logger.warning( + "Timed out after %ss draining in-flight credit lease work; " + "any credits it holds will be released by server-side expiry", + SHUTDOWN_DRAIN_TIMEOUT, + ) + async def _never_raises(coro: Awaitable[None]) -> None: try: diff --git a/src/schematic/leases/types.py b/src/schematic/leases/types.py index 4a9d8a8..16f406a 100644 --- a/src/schematic/leases/types.py +++ b/src/schematic/leases/types.py @@ -32,6 +32,10 @@ # How long a prewarm waits for a freshly identified company to surface in the # datastream cache before giving up. DEFAULT_PREWARM_RESOLVE_TIMEOUT = 5.0 +# How long shutdown waits for in-flight lease work to land before giving up on +# it. Bounded on purpose: a shutdown that hangs is worse than a hold the server +# expires in DEFAULT_LEASE_DURATION. +SHUTDOWN_DRAIN_TIMEOUT = 5.0 @dataclass diff --git a/tests/custom/test_client.py b/tests/custom/test_client.py index 9f76216..b02cc7d 100644 --- a/tests/custom/test_client.py +++ b/tests/custom/test_client.py @@ -2903,6 +2903,41 @@ async def test_shutdown_stops_the_sweep_and_releases_a_per_process_lease(self): assert client._lease_manager._sweep_task is None client.credits.release_credit_lease.assert_awaited_once_with("lse_1", request_options=None) + async def test_shutdown_releases_a_lease_an_in_flight_prewarm_installs(self): + # The prewarm's acquire is on the wire when shutdown starts. Cancelling + # the prewarm cancels its shield, not the acquire, so the lease still + # lands: shutdown has to drain it before listing the store, or nothing + # releases it and the credits stay held until server-side expiry. + client = _async_lease_client() + client._datastream_client = _lease_datastream([]) + on_the_wire = asyncio.Event() + + async def slow_acquire(**kwargs): + on_the_wire.set() + await asyncio.sleep(0.05) + return _lease_grant() + + client.credits.acquire_credit_lease = AsyncMock(side_effect=slow_acquire) + + client._spawn_prewarm({"id": "co_1"}, ["bilcr_inference"]) + await asyncio.wait_for(on_the_wire.wait(), 1) + await client.shutdown() + # An untracked acquire would install here, behind the release. + await asyncio.sleep(0.1) + + assert client._lease_store.list_leases() == [] + client.credits.release_credit_lease.assert_awaited_once_with("lse_1", request_options=None) + + async def test_prewarm_started_during_shutdown_acquires_nothing(self): + client = _async_lease_client() + client._datastream_client = _lease_datastream([]) + try: + client._is_shutting_down = True + await client.prewarm({"id": "co_1"}, ["bilcr_inference"]) + client.credits.acquire_credit_lease.assert_not_awaited() + finally: + await self._drain(client) + async def test_shutdown_leaves_a_shared_lease_for_the_pods_still_drawing_on_it(self): redis_client = make_fake_redis() client = _async_lease_client( diff --git a/tests/leases/test_lease_manager.py b/tests/leases/test_lease_manager.py index a4c0dbf..a507ad9 100644 --- a/tests/leases/test_lease_manager.py +++ b/tests/leases/test_lease_manager.py @@ -352,6 +352,65 @@ async def counting_sweep(now: Optional[float] = None) -> int: assert len(swept) == ticks +async def test_a_cancelled_acquire_is_still_drained_to_completion(clock: VirtualClock) -> None: + # Cancelling the caller cancels its shield, not the wire call underneath, + # so the acquire goes on to install a lease. The drain set has to hold it, + # or a close releases the store before that write lands. + manager, store, wire = _make_manager(clock) + wire.acquire_responses.append(_lease(clock)) + landed = asyncio.Event() + wire.during_acquire = landed.wait + + caller = asyncio.ensure_future(manager.acquire_if_needed("co_1", "ct_1")) + await asyncio.sleep(0) + caller.cancel() + with pytest.raises(asyncio.CancelledError): + await caller + # The registry entry is gone the instant the caller unwinds; the drain set + # is what is left holding the acquire. + assert manager._inflight_acquire == {} + assert manager._background + + landed.set() + await manager.drain() + + assert not manager._background + installed = await store.get("co_1", "ct_1") + assert installed is not None and installed.lease_id == "lse_1" + + +async def test_drain_gives_up_on_work_that_will_not_land(clock: VirtualClock, monkeypatch) -> None: + monkeypatch.setattr("schematic.leases.lease_manager.SHUTDOWN_DRAIN_TIMEOUT", 0.01) + manager, store, wire = _make_manager(clock) + wire.acquire_responses.append(_lease(clock)) + wire.during_acquire = asyncio.Event().wait + + caller = asyncio.ensure_future(manager.acquire_if_needed("co_1", "ct_1")) + await asyncio.sleep(0) + caller.cancel() + with pytest.raises(asyncio.CancelledError): + await caller + + # Bounded: a shutdown that hangs is worse than a hold the server expires, + # so the acquire is cancelled and never installs. + await manager.drain() + assert await store.get("co_1", "ct_1") is None + + +async def test_stop_keeps_a_background_extend_from_starting(clock: VirtualClock) -> None: + manager, store, wire = _make_manager(clock) + wire.acquire_responses.append(_lease(clock)) + await manager.acquire_if_needed("co_1", "ct_1") + await manager.drain() + await store.try_reserve("co_1", "ct_1", 900) + + manager.stop() + manager.extend_in_background("co_1", "ct_1") + await manager.drain() + + assert wire.extend_calls == [] + + async def test_sweep_loop_survives_a_failing_sweep(clock: VirtualClock) -> None: leases = InMemoryLeaseStore(clock=clock) reservations = InMemoryReservationStore(leases, clock=clock)