From 16e71cc4c8fcb9874425fbb96e17ad199374dc0a Mon Sep 17 00:00:00 2001 From: Brian O'Kelley Date: Fri, 25 Sep 2026 11:15:51 +0000 Subject: [PATCH 1/2] fix(server): settle lifespan cleanup after startup failures --- src/adcp/server/serve.py | 192 ++++++++---- tests/test_serve_lifespan_hooks.py | 469 +++++++++++++++++++++++++++++ 2 files changed, 603 insertions(+), 58 deletions(-) diff --git a/src/adcp/server/serve.py b/src/adcp/server/serve.py index 257a811bc..88ca66d61 100644 --- a/src/adcp/server/serve.py +++ b/src/adcp/server/serve.py @@ -18,11 +18,13 @@ async def get_adcp_capabilities(self, params, context=None): from __future__ import annotations +import asyncio +import contextlib import logging import os import sys import warnings -from collections.abc import Awaitable, Callable, Sequence +from collections.abc import AsyncIterator, Awaitable, Callable, Sequence from contextvars import ContextVar from dataclasses import dataclass from types import MethodType @@ -122,6 +124,122 @@ class RequestMetadata: """ +async def _run_shutdown_hooks(hooks: tuple[LifespanHook, ...]) -> BaseException | None: + """Attempt every callback in its startup task, without formatting adopter data.""" + first_error: BaseException | None = None + for index, hook in enumerate(hooks): + try: + await hook() + except BaseException as exc: # noqa: BLE001 + if first_error is None: + first_error = exc + logger.error("on_shutdown hook failed; continuing cleanup (hook_index=%d)", index) + return first_error + + +async def _settle_lifespan_task( + task: asyncio.Task[BaseException | None], +) -> tuple[BaseException | None, asyncio.CancelledError | None]: + """Join the hook task despite repeated cancellation of its framework waiter.""" + import anyio + + cancellation: asyncio.CancelledError | None = None + # This scope belongs to the framework waiter, never above an adopter's + # scope in the hook task: shutdown may need to exit that adopter scope. + with anyio.CancelScope(shield=True): + while not task.done(): + try: + await asyncio.shield(task) + except asyncio.CancelledError as exc: + if cancellation is None: + cancellation = exc + return task.result(), cancellation + + +@contextlib.asynccontextmanager +async def _user_lifespan_hooks( + startup: tuple[LifespanHook, ...], shutdown: tuple[LifespanHook, ...] +) -> AsyncIterator[None]: + """Keep paired hooks in one retained task inside the framework lifespans.""" + if not startup and not shutdown: + yield + return + + ready: asyncio.Future[BaseException | None] = asyncio.get_running_loop().create_future() + stop = asyncio.Event() + started = False + abort_startup = False + hook_error: BaseException | None = None + ended_early = False + owner = asyncio.current_task() + + async def run() -> BaseException | None: + nonlocal started, hook_error + started = True + try: + if not abort_startup: + for hook in startup: + await hook() + ready.set_result(None) + await stop.wait() + except BaseException as exc: # noqa: BLE001 + # Includes an adopter-owned scope ending this task after startup. + hook_error = exc + if not ready.done(): + ready.set_result(exc) + cleanup_error = await _run_shutdown_hooks(shutdown) + if ( + ready.result() is None + and isinstance(hook_error, asyncio.CancelledError) + and cleanup_error is not None + ): + # An adopter TaskGroup cancels its host to wake it, then exposes + # the causal child failure when the shutdown hook exits the group. + return cleanup_error + return hook_error if hook_error is not None else cleanup_error + + task = asyncio.create_task(run()) + + def completed(_task: asyncio.Task[BaseException | None]) -> None: + nonlocal ended_early + if ready.done() and ready.result() is None and not stop.is_set(): + ended_early = True + if owner is not None: + owner.cancel() + + task.add_done_callback(completed) + try: + error = await asyncio.shield(ready) + if error is not None: + raise error + yield + finally: + primary_error = sys.exc_info()[1] + if not ready.done(): + # Interrupt in-flight startup once. If the task has not run yet, + # let it skip startup and still attempt partial-startup cleanup. + abort_startup = True + if started: + task.cancel() + stop.set() + error, cancellation = await _settle_lifespan_task(task) + if ( + ended_early + and error is not None + and (primary_error is None or isinstance(primary_error, asyncio.CancelledError)) + ): + if isinstance(error, asyncio.CancelledError): + # AnyIO framework scopes can suppress their cancellation + # signals. An unexpectedly ended hook lifecycle is a failure. + raise RuntimeError("Server lifespan hook task stopped unexpectedly") from None + raise error + if primary_error is None: + if cancellation is not None: + raise cancellation + if error is not None: + raise error + + @dataclass(frozen=True) class ServeConfig: """Configuration bundle for :func:`serve`. @@ -973,10 +1091,20 @@ def resolver(request): dropping the hook. See ``examples/scheduler_lifespan.py``. on_shutdown: Optional sequence of :data:`LifespanHook` zero-arg async callables fired before either inner lifespan tears - down. Every hook runs on a best-effort basis even if an - earlier one raised; the first failure re-raises so - Starlette surfaces it, later failures land in - ``logger.error``. Same ``transport="both"`` restriction + down, including when an adopter startup hook fails. Every + hook runs once in registration order, even after an earlier + error or cancellation. Hooks must tolerate partial startup; + register dependent resources' shutdowns in reverse dependency + order. Startup and shutdown share one retained task on the server + loop, preserving paired ContextVar tokens and AnyIO scopes. + Its initial context is copied from the framework task; startup + ContextVar mutations stay local to the hook lifecycle. Cleanup + settles before cancellation propagates or transports close. + There is no SDK cleanup timeout: a hook that never completes + can hold shutdown indefinitely. Startup/body failures take + precedence; otherwise cancellation or the first cleanup failure + propagates. SDK cleanup diagnostics omit exception text and + callable representations. Same ``transport="both"`` restriction as ``on_startup``. Example (MCP): @@ -1931,7 +2059,6 @@ def _build_mcp_and_a2a_app( Returns the size-limit-wrapped ASGI app. Wire to uvicorn / Starlette / your test harness as you would any other ASGI app. """ - import contextlib from starlette.applications import Starlette from starlette.types import ASGIApp, Receive, Scope, Send @@ -2049,59 +2176,8 @@ def _build_mcp_and_a2a_app( async def _composed_lifespan(_app): # type: ignore[no-untyped-def] async with mcp_inner.router.lifespan_context(mcp_inner): async with a2a_inner.router.lifespan_context(a2a_inner): - for hook in user_startup: - await hook() - try: + async with _user_lifespan_hooks(user_startup, user_shutdown): yield - finally: - # Run every shutdown hook even if an earlier one - # raised — adopters that wire multiple cleanup - # hooks (close DB pool, stop scheduler, drain - # queue) want all of them attempted on a - # best-effort basis. Re-raise the first failure - # so Starlette surfaces it; log later failures - # without ``exc_info`` so adopter closure state - # (DB DSNs, tokens stashed in hook captures) - # doesn't end up verbatim in shutdown logs that - # downstream aggregators attach locals to. - # - # Catch ``Exception`` only — ``CancelledError`` / - # ``KeyboardInterrupt`` / ``SystemExit`` are the - # exact signals uvicorn uses to drive shutdown, - # and we want them to propagate immediately - # rather than getting collected into ``first_error``. - first_error: Exception | None = None - for hook in user_shutdown: - try: - await hook() - except Exception as exc: # noqa: BLE001 - if first_error is None: - first_error = exc - else: - logger.error( - "on_shutdown hook %r raised: %s " - "(suppressed; earlier hook also " - "raised)", - getattr(hook, "__name__", hook), - exc, - ) - if first_error is not None: - # If we reached the ``finally`` because the - # body raised (framework lifespan teardown, - # request handler escaping), don't overwrite - # that propagation with our shutdown error — - # the operator wants to see the upstream - # cause, not a secondary cleanup failure. - # Log the shutdown error so it isn't lost, - # let the original exception keep propagating. - if sys.exc_info()[0] is None: - raise first_error - logger.error( - "on_shutdown hook raised during exception " - "unwinding: %s (suppressed; the upstream " - "exception takes precedence)", - first_error, - ) parent = Starlette(lifespan=_composed_lifespan) diff --git a/tests/test_serve_lifespan_hooks.py b/tests/test_serve_lifespan_hooks.py index c48689125..8fae20145 100644 --- a/tests/test_serve_lifespan_hooks.py +++ b/tests/test_serve_lifespan_hooks.py @@ -7,6 +7,12 @@ from __future__ import annotations +import asyncio +import importlib +from contextlib import asynccontextmanager +from contextvars import ContextVar, Token +from unittest.mock import Mock + import pytest starlette = pytest.importorskip("starlette") @@ -97,6 +103,62 @@ async def shutdown() -> None: assert events == ["startup", "shutdown"] +def test_paired_hooks_share_contextvar_tokens() -> None: + value: ContextVar[str] = ContextVar("hook_value", default="initial") + token: Token[str] | None = None + events: list[str] = [] + + async def startup() -> None: + nonlocal token + assert value.get() == "parent" + token = value.set("started") + + async def shutdown() -> None: + assert value.get() == "started" + assert token is not None + value.reset(token) + events.append(value.get()) + + parent_token = value.set("parent") + try: + with TestClient(_build_app(on_startup=[startup], on_shutdown=[shutdown])): + assert value.get() == "parent" + assert value.get() == "parent" + finally: + value.reset(parent_token) + assert events == ["parent"] + + +@pytest.mark.parametrize("kind", ["cancel_scope", "task_group"]) +def test_paired_hooks_can_enter_and_exit_an_anyio_scope(kind: str) -> None: + import anyio + + scope = None + events: list[str] = [] + + async def startup() -> None: + nonlocal scope + if kind == "cancel_scope": + scope = anyio.CancelScope() + scope.__enter__() + else: + scope = anyio.create_task_group() + await scope.__aenter__() + events.append("entered") + + async def shutdown() -> None: + assert scope is not None + if kind == "cancel_scope": + scope.__exit__(None, None, None) + else: + await scope.__aexit__(None, None, None) + events.append("exited") + + with TestClient(_build_app(on_startup=[startup], on_shutdown=[shutdown])): + pass + assert events == ["entered", "exited"] + + # ----- Failure modes ---------------------------------------------------- @@ -120,6 +182,119 @@ async def boom() -> None: assert "boot-time wiring broke" in _flatten_exception_text(exc_info.value) +def test_later_startup_failure_closes_started_resources_once() -> None: + events: list[str] = [] + + async def start_resource() -> None: + events.append("resource_started") + + async def fail_later_startup() -> None: + events.append("later_startup") + raise RuntimeError("later startup failed") + + async def close_resource() -> None: + events.append("resource_closed") + + app = _build_app( + on_startup=[start_resource, fail_later_startup], + on_shutdown=[close_resource], + ) + with pytest.raises(BaseException) as raised: + with TestClient(app): + pytest.fail("failed startup must not admit requests") + + assert "later startup failed" in _flatten_exception_text(raised.value) + assert events == ["resource_started", "later_startup", "resource_closed"] + + +def test_startup_failure_remains_primary_after_all_cleanup_hooks(caplog) -> None: + events: list[str] = [] + + async def fail_startup() -> None: + raise ValueError("primary startup failure") + + async def fail_cleanup() -> None: + events.append("first_cleanup") + raise RuntimeError("secret-provider-body") + + async def finish_cleanup() -> None: + events.append("last_cleanup") + + app = _build_app(on_startup=[fail_startup], on_shutdown=[fail_cleanup, finish_cleanup]) + with pytest.raises(BaseException) as raised: + with TestClient(app): + pytest.fail("failed startup must not admit requests") + + assert "primary startup failure" in _flatten_exception_text(raised.value) + assert events == ["first_cleanup", "last_cleanup"] + assert "secret-provider-body" not in caplog.text + + +@pytest.mark.parametrize("startup_failure", ["cancelled", "error"]) +async def test_startup_failure_settles_cleanup_despite_repeated_cancellation( + startup_failure: str, +) -> None: + events: list[str] = [] + later_started = asyncio.Event() + fail_startup = asyncio.Event() + cleanup_started = asyncio.Event() + allow_cleanup = asyncio.Event() + incoming: asyncio.Queue[dict] = asyncio.Queue() + outgoing: asyncio.Queue[dict] = asyncio.Queue() + + async def first() -> None: + events.append("started") + + async def later() -> None: + later_started.set() + try: + await fail_startup.wait() + except asyncio.CancelledError: + events.append("startup_cancelled") + raise + raise ValueError("primary startup failure") + + async def close() -> None: + cleanup_started.set() + await allow_cleanup.wait() + events.append("closed") + + app = _build_app(on_startup=[first, later], on_shutdown=[close]) + task = asyncio.create_task( + app( + {"type": "lifespan", "asgi": {"version": "3.0"}, "state": {}}, + incoming.get, + outgoing.put, + ) + ) + try: + await incoming.put({"type": "lifespan.startup"}) + await asyncio.wait_for(later_started.wait(), 5) + if startup_failure == "cancelled": + task.cancel() + else: + fail_startup.set() + await asyncio.wait_for(cleanup_started.wait(), 5) + task.cancel() + await asyncio.sleep(0) + task.cancel() + await asyncio.sleep(0) + assert not task.done() + assert "closed" not in events + finally: + allow_cleanup.set() + with pytest.raises(BaseException) as raised: + await asyncio.wait_for(task, 5) + + if startup_failure == "cancelled": + assert isinstance(raised.value, asyncio.CancelledError) + assert events == ["started", "startup_cancelled", "closed"] + else: + assert "primary startup failure" in _flatten_exception_text(raised.value) + assert events == ["started", "closed"] + assert (await outgoing.get())["type"] == "lifespan.startup.failed" + + def _flatten_exception_text(exc: BaseException) -> str: """Collect ``str(exc)`` plus every cause / context / group leaf.""" parts: list[str] = [] @@ -178,6 +353,300 @@ async def last_ok() -> None: assert "scheduler stop failed" in _flatten_exception_text(raised) +def test_secondary_shutdown_diagnostics_do_not_format_adopter_values(caplog) -> None: + class SecretClosure: + async def __call__(self) -> None: + raise ValueError("secret-secondary-error") + + def __repr__(self) -> str: + return "secret-callable-repr" + + async def first_error() -> None: + raise RuntimeError("first cleanup failure") + + app = _build_app(on_shutdown=[first_error, SecretClosure()]) + with pytest.raises(BaseException) as raised: + with TestClient(app): + pass + + assert "first cleanup failure" in _flatten_exception_text(raised.value) + assert "on_shutdown hook failed" in caplog.text + assert "secret-secondary-error" not in caplog.text + assert "secret-callable-repr" not in caplog.text + assert all(record.exc_info is None for record in caplog.records if record.name == "adcp.server") + + +async def test_cancelled_shutdown_hook_does_not_skip_later_callbacks() -> None: + from adcp.server.serve import _user_lifespan_hooks + + events: list[str] = [] + + async def cancel_first() -> None: + events.append("cancelled") + raise asyncio.CancelledError() + + async def close_second() -> None: + events.append("closed") + + with pytest.raises(asyncio.CancelledError): + async with _user_lifespan_hooks((), (cancel_first, close_second)): + pass + assert events == ["cancelled", "closed"] + + +async def test_anyio_cancelled_scope_still_settles_shutdown_hooks() -> None: + import anyio + + from adcp.server.serve import _user_lifespan_hooks + + events: list[str] = [] + + async def close() -> None: + await anyio.sleep(0) + events.append("closed") + + with anyio.CancelScope() as scope: + scope.cancel() + async with _user_lifespan_hooks((), (close,)): + pass + + assert events == ["closed"] + + +async def test_hook_lifecycle_failure_does_not_leave_framework_waiting() -> None: + import anyio + + crash = asyncio.Event() + closed = asyncio.Event() + group = None + incoming: asyncio.Queue[dict] = asyncio.Queue() + outgoing: asyncio.Queue[dict] = asyncio.Queue() + + async def worker() -> None: + await crash.wait() + raise RuntimeError("worker failed") + + async def start() -> None: + nonlocal group + group = anyio.create_task_group() + await group.__aenter__() + group.start_soon(worker) + + async def close() -> None: + assert group is not None + try: + await group.__aexit__(None, None, None) + finally: + closed.set() + + app = _build_app(on_startup=[start], on_shutdown=[close]) + task = asyncio.create_task( + app( + {"type": "lifespan", "asgi": {"version": "3.0"}, "state": {}}, + incoming.get, + outgoing.put, + ) + ) + try: + await incoming.put({"type": "lifespan.startup"}) + assert (await asyncio.wait_for(outgoing.get(), 5))["type"] == "lifespan.startup.complete" + crash.set() + with pytest.raises(BaseException) as raised: + await asyncio.wait_for(task, 5) + assert not isinstance(raised.value, asyncio.TimeoutError) + assert closed.is_set() + assert (await outgoing.get())["type"] == "lifespan.shutdown.failed" + finally: + if not task.done(): + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + +async def test_body_failure_remains_primary_after_hook_lifecycle_failure() -> None: + import anyio + + from adcp.server.serve import _user_lifespan_hooks + + crash = asyncio.Event() + group = None + + async def worker() -> None: + await crash.wait() + raise RuntimeError("secondary worker failure") + + async def start() -> None: + nonlocal group + group = anyio.create_task_group() + await group.__aenter__() + group.start_soon(worker) + + async def close() -> None: + assert group is not None + await group.__aexit__(None, None, None) + + with pytest.raises(ValueError, match="primary body failure"): + async with _user_lifespan_hooks((start,), (close,)): + crash.set() + try: + await asyncio.wait_for(asyncio.Event().wait(), 5) + except asyncio.CancelledError: + raise ValueError("primary body failure") from None + + +@pytest.mark.parametrize("cancel_body", [False, True]) +async def test_repeated_cancellation_settles_cleanup_before_transport_teardown( + monkeypatch, cancel_body: bool +) -> None: + events: list[str] = [] + close_requested = asyncio.Event() + close_started = asyncio.Event() + allow_close = asyncio.Event() + worker: asyncio.Task[None] | None = None + incoming: asyncio.Queue[dict] = asyncio.Queue() + outgoing: asyncio.Queue[dict] = asyncio.Queue() + + async def work() -> None: + await close_requested.wait() + close_started.set() + await allow_close.wait() + events.append("worker_settled") + + async def start() -> None: + nonlocal worker + worker = asyncio.create_task(work()) + + async def close_worker() -> None: + close_requested.set() + assert worker is not None + await worker + events.append("worker_joined") + + async def close_pool() -> None: + assert worker is not None and worker.done() and not worker.cancelled() + events.append("pool_closed") + + # Observe the real inner transports without replacing their lifecycle. + def observe(inner, name): + original = inner.router.lifespan_context + + @asynccontextmanager + async def lifespan(app): + try: + async with original(app): + events.append(name + "_started") + try: + yield + finally: + events.append(name + "_closing") + finally: + events.append(name + "_closed") + + inner.router.lifespan_context = lifespan + return inner + + serve_module = importlib.import_module("adcp.server.serve") + a2a_module = importlib.import_module("adcp.server.a2a_server") + original_mcp = serve_module.create_mcp_server + original_a2a = a2a_module.create_a2a_server + + def create_mcp(*args, **kwargs): + server = original_mcp(*args, **kwargs) + original_app = server.streamable_http_app + monkeypatch.setattr(server, "streamable_http_app", lambda: observe(original_app(), "mcp")) + return server + + monkeypatch.setattr(serve_module, "create_mcp_server", create_mcp) + monkeypatch.setattr( + a2a_module, "create_a2a_server", lambda *a, **kw: observe(original_a2a(*a, **kw), "a2a") + ) + app = _build_app(on_startup=[start], on_shutdown=[close_worker, close_pool]) + task = asyncio.create_task( + app( + {"type": "lifespan", "asgi": {"version": "3.0"}, "state": {}}, + incoming.get, + outgoing.put, + ) + ) + try: + await incoming.put({"type": "lifespan.startup"}) + startup = await asyncio.wait_for(outgoing.get(), 5) + assert startup["type"] == "lifespan.startup.complete" + if cancel_body: + task.cancel() + else: + await incoming.put({"type": "lifespan.shutdown"}) + await asyncio.wait_for(close_started.wait(), 5) + task.cancel() + await asyncio.sleep(0) + task.cancel() + await asyncio.sleep(0) + assert not task.done() + assert worker is not None and not worker.done() + assert events == ["mcp_started", "a2a_started"] + finally: + allow_close.set() + try: + await asyncio.wait_for(task, 5) + except asyncio.CancelledError: + pass + + assert task.cancelled() + assert events[:5] == [ + "mcp_started", + "a2a_started", + "worker_settled", + "worker_joined", + "pool_closed", + ] + assert events[5:] == ["a2a_closing", "a2a_closed", "mcp_closing", "mcp_closed"] + shutdown = await outgoing.get() + assert shutdown["type"] == "lifespan.shutdown.failed" + + +def test_synchronous_serve_owns_the_hook_event_loop(monkeypatch) -> None: + import uvicorn + + from adcp.server import serve + + events: list[str] = [] + hook_loops = [] + + async def start() -> None: + hook_loops.append(asyncio.get_running_loop()) + events.append("start") + + async def close() -> None: + hook_loops.append(asyncio.get_running_loop()) + await asyncio.sleep(0) + events.append("close") + + async def drive_lifespan(server, sockets) -> None: + messages = iter([{"type": "lifespan.startup"}, {"type": "lifespan.shutdown"}]) + results = [] + + async def receive(): + return next(messages) + + async def send(message): + results.append(message["type"]) + + await server.config.app( + {"type": "lifespan", "asgi": {"version": "3.0"}, "state": {}}, receive, send + ) + assert results == ["lifespan.startup.complete", "lifespan.shutdown.complete"] + + sock = Mock() + module = importlib.import_module("adcp.server.serve") + monkeypatch.setattr(module, "_bind_reusable_socket", lambda *args: sock) + monkeypatch.setattr(uvicorn.Server, "serve", drive_lifespan) + serve(_Handler(), transport="both", on_startup=[start], on_shutdown=[close]) + assert events == ["start", "close"] + assert hook_loops[0] is hook_loops[1] + assert hook_loops[0].is_closed() + sock.close.assert_called_once() + + # ----- Boot-time validation -------------------------------------------- From a1355c5efa59118fd3a9fa91f842b625ad50705c Mon Sep 17 00:00:00 2001 From: Brian O'Kelley Date: Sat, 26 Sep 2026 14:48:39 +0000 Subject: [PATCH 2/2] fix(server): preserve startup errors through cancellation --- src/adcp/server/serve.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/src/adcp/server/serve.py b/src/adcp/server/serve.py index 88ca66d61..c5a46495a 100644 --- a/src/adcp/server/serve.py +++ b/src/adcp/server/serve.py @@ -223,6 +223,14 @@ def completed(_task: asyncio.Task[BaseException | None]) -> None: task.cancel() stop.set() error, cancellation = await _settle_lifespan_task(task) + startup_error = ready.result() if ready.done() else None + if startup_error is not None and ( + primary_error is None or isinstance(primary_error, asyncio.CancelledError) + ): + # A cancellation can win the wait for ``ready`` after the hook + # task has already recorded a startup failure. Keep that failure + # primary once cleanup has settled. + raise startup_error if ( ended_early and error is not None