From 24d7aefe28c9b5fb8aab3fb3fefd89f2dd62f52e Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Fri, 13 Mar 2026 18:47:42 -0500 Subject: [PATCH] fix: shield lifespan teardown from cancellation (#3480) * fix: shield lifespan teardown from cancellation Co-authored-by: Claude * fix: stabilize flaky task and timeout tests under parallel execution Co-authored-by: Claude --------- Co-authored-by: Claude --- src/fastmcp/server/mixins/lifespan.py | 48 +++++----- src/fastmcp/server/mixins/transport.py | 2 +- tests/client/client/test_timeout.py | 10 +-- tests/client/tasks/conftest.py | 15 ++++ tests/server/tasks/conftest.py | 15 ++++ tests/server/test_server_lifespan.py | 117 +++++++++++++++++++++++++ 6 files changed, 178 insertions(+), 29 deletions(-) create mode 100644 tests/client/tasks/conftest.py create mode 100644 tests/server/tasks/conftest.py diff --git a/src/fastmcp/server/mixins/lifespan.py b/src/fastmcp/server/mixins/lifespan.py index 0f2589272..583944d61 100644 --- a/src/fastmcp/server/mixins/lifespan.py +++ b/src/fastmcp/server/mixins/lifespan.py @@ -8,6 +8,7 @@ from collections.abc import AsyncIterator from contextlib import AsyncExitStack, asynccontextmanager, suppress from typing import TYPE_CHECKING, Any +import anyio from uncalled_for import SharedContext import fastmcp @@ -155,30 +156,37 @@ class LifespanMixin: self._lifespan_result = None return + # Use an explicit AsyncExitStack so we can shield teardown from + # cancellation. Without this, Ctrl-C causes CancelledError to + # propagate into lifespan finally blocks, preventing any async + # cleanup (e.g. closing DB connections, flushing buffers). + stack = AsyncExitStack() try: - async with ( - self._lifespan(self) as user_lifespan_result, - self._docket_lifespan(), - ): - self._lifespan_result = user_lifespan_result - self._lifespan_result_set = True + user_lifespan_result = await stack.enter_async_context(self._lifespan(self)) + await stack.enter_async_context(self._docket_lifespan()) - async with AsyncExitStack[bool | None]() as stack: - # Start lifespans for all providers - for provider in self.providers: - await stack.enter_async_context(provider.lifespan()) + self._lifespan_result = user_lifespan_result + self._lifespan_result_set = True - self._started.set() - try: - yield - finally: - self._started.clear() + # Start lifespans for all providers + for provider in self.providers: + await stack.enter_async_context(provider.lifespan()) + + self._started.set() + try: + yield + finally: + self._started.clear() finally: - async with self._lifespan_lock: - self._lifespan_ref_count -= 1 - if self._lifespan_ref_count == 0: - self._lifespan_result_set = False - self._lifespan_result = None + try: + with anyio.CancelScope(shield=True): + await stack.aclose() + finally: + async with self._lifespan_lock: + self._lifespan_ref_count -= 1 + if self._lifespan_ref_count == 0: + self._lifespan_result_set = False + self._lifespan_result = None def _setup_task_protocol_handlers(self: FastMCP) -> None: """Register SEP-1686 task protocol handlers with SDK. diff --git a/src/fastmcp/server/mixins/transport.py b/src/fastmcp/server/mixins/transport.py index 833b5e385..8f94c0b1e 100644 --- a/src/fastmcp/server/mixins/transport.py +++ b/src/fastmcp/server/mixins/transport.py @@ -255,7 +255,7 @@ class TransportMixin: uvicorn_config_from_user = uvicorn_config or {} config_kwargs: dict[str, Any] = { - "timeout_graceful_shutdown": 0, + "timeout_graceful_shutdown": 2, "lifespan": "on", "ws": "websockets-sansio", } diff --git a/tests/client/client/test_timeout.py b/tests/client/client/test_timeout.py index 0b3b0fd3a..5106e7cb6 100644 --- a/tests/client/client/test_timeout.py +++ b/tests/client/client/test_timeout.py @@ -1,7 +1,5 @@ """Client timeout tests.""" -import sys - import pytest from mcp import McpError @@ -36,15 +34,11 @@ class TestTimeout: with pytest.raises(McpError): await client.call_tool("sleep", {"seconds": 0.1}, timeout=0.01) - @pytest.mark.skipif( - sys.platform == "win32", - reason="This test is flaky on Windows. Sometimes the client timeout is respected and sometimes it is not.", - ) async def test_timeout_tool_call_overrides_client_timeout_even_if_lower( self, fastmcp_server: FastMCP ): async with Client( transport=FastMCPTransport(fastmcp_server), - timeout=0.01, + timeout=0.1, ) as client: - await client.call_tool("sleep", {"seconds": 0.1}, timeout=2) + await client.call_tool("sleep", {"seconds": 0.5}, timeout=2) diff --git a/tests/client/tasks/conftest.py b/tests/client/tasks/conftest.py new file mode 100644 index 000000000..d0921e057 --- /dev/null +++ b/tests/client/tasks/conftest.py @@ -0,0 +1,15 @@ +"""Configuration for client task tests. + +Task tests require Docket infrastructure (Redis-backed task queue) which can +take significant time to initialize, especially under parallel test execution. +The default 5s timeout is too tight for these tests. +""" + +import pytest + + +def pytest_collection_modifyitems(items: list[pytest.Item]) -> None: + """Increase timeout for task tests that need Docket infrastructure.""" + for item in items: + if not item.get_closest_marker("timeout"): + item.add_marker(pytest.mark.timeout(15)) diff --git a/tests/server/tasks/conftest.py b/tests/server/tasks/conftest.py new file mode 100644 index 000000000..7d48204bb --- /dev/null +++ b/tests/server/tasks/conftest.py @@ -0,0 +1,15 @@ +"""Configuration for server task tests. + +Task tests require Docket infrastructure (Redis-backed task queue) which can +take significant time to initialize, especially under parallel test execution. +The default 5s timeout is too tight for these tests. +""" + +import pytest + + +def pytest_collection_modifyitems(items: list[pytest.Item]) -> None: + """Increase timeout for task tests that need Docket infrastructure.""" + for item in items: + if not item.get_closest_marker("timeout"): + item.add_marker(pytest.mark.timeout(15)) diff --git a/tests/server/test_server_lifespan.py b/tests/server/test_server_lifespan.py index 2d2b3446e..fe67eb44c 100644 --- a/tests/server/test_server_lifespan.py +++ b/tests/server/test_server_lifespan.py @@ -1,14 +1,17 @@ """Tests for server_lifespan and session_lifespan behavior.""" +import asyncio from collections.abc import AsyncIterator from contextlib import asynccontextmanager from typing import Any +import anyio import pytest from fastmcp import Client, FastMCP from fastmcp.server.context import Context from fastmcp.server.lifespan import ContextManagerLifespan, lifespan +from fastmcp.server.providers import Provider from fastmcp.utilities.lifespan import combine_lifespans @@ -543,3 +546,117 @@ class TestCombineLifespans: "dict_exit", "mapping_exit", ] + + +class TestLifespanTeardownShielding: + """Test that async operations in lifespan teardown complete under cancellation. + + When a server shuts down (e.g. Ctrl-C), the cancel scope becomes active. + Lifespan teardown must be shielded from cancellation so that async cleanup + (closing DB connections, flushing buffers, etc.) can actually run. + """ + + async def test_server_lifespan_async_teardown_under_cancellation(self): + """Async operations in server lifespan finally block complete even when cancelled.""" + events: list[str] = [] + + @asynccontextmanager + async def server_lifespan(mcp: FastMCP) -> AsyncIterator[dict[str, Any]]: + events.append("setup") + try: + yield {} + finally: + events.append("teardown_start") + await asyncio.sleep(0) + events.append("teardown_complete") + + mcp = FastMCP("TestServer", lifespan=server_lifespan) + + async with anyio.create_task_group() as tg: + + async def run_and_cancel() -> None: + async with mcp._lifespan_manager(): + tg.cancel_scope.cancel() + + tg.start_soon(run_and_cancel) + + assert "setup" in events + assert "teardown_start" in events + assert "teardown_complete" in events + + async def test_provider_lifespan_async_teardown_under_cancellation(self): + """Async operations in provider lifespan finally block complete when cancelled.""" + events: list[str] = [] + + class TestProvider(Provider): + @asynccontextmanager + async def lifespan(self) -> AsyncIterator[None]: + events.append("provider_setup") + try: + yield + finally: + events.append("provider_teardown_start") + await asyncio.sleep(0) + events.append("provider_teardown_complete") + + mcp = FastMCP("TestServer", providers=[TestProvider()]) + + async with anyio.create_task_group() as tg: + + async def run_and_cancel() -> None: + async with mcp._lifespan_manager(): + tg.cancel_scope.cancel() + + tg.start_soon(run_and_cancel) + + assert "provider_setup" in events + assert "provider_teardown_start" in events + assert "provider_teardown_complete" in events + + async def test_composed_lifespans_async_teardown_under_cancellation(self): + """Both server and provider async teardown completes under cancellation.""" + events: list[str] = [] + + @asynccontextmanager + async def server_lifespan(mcp: FastMCP) -> AsyncIterator[dict[str, Any]]: + events.append("server_setup") + try: + yield {} + finally: + events.append("server_teardown_start") + await asyncio.sleep(0) + events.append("server_teardown_complete") + + class TestProvider(Provider): + @asynccontextmanager + async def lifespan(self) -> AsyncIterator[None]: + events.append("provider_setup") + try: + yield + finally: + events.append("provider_teardown_start") + await asyncio.sleep(0) + events.append("provider_teardown_complete") + + mcp = FastMCP( + "TestServer", + lifespan=server_lifespan, + providers=[TestProvider()], + ) + + async with anyio.create_task_group() as tg: + + async def run_and_cancel() -> None: + async with mcp._lifespan_manager(): + tg.cancel_scope.cancel() + + tg.start_soon(run_and_cancel) + + assert events == [ + "server_setup", + "provider_setup", + "provider_teardown_start", + "provider_teardown_complete", + "server_teardown_start", + "server_teardown_complete", + ]