diff --git a/studio/backend/main.py b/studio/backend/main.py index 2435efba21..08c7bc6826 100644 --- a/studio/backend/main.py +++ b/studio/backend/main.py @@ -296,6 +296,7 @@ from utils.hardware import ( import utils.hardware.hardware as _hw_module from utils.cache_cleanup import clear_unsloth_compiled_cache +from utils.lifespan_shutdown import run_lifespan_shutdown from utils.native_path_leases import native_path_leases_supported from utils.update_status import ( get_studio_install_source_status, @@ -463,9 +464,12 @@ async def lifespan(app: FastAPI): else: app.state.bootstrap_password = storage.get_bootstrap_password() yield - await asyncio.to_thread(terminate_hub_downloads) - _hw_module.DEVICE = None - clear_unsloth_compiled_cache() + + await run_lifespan_shutdown( + terminate_hub_downloads, + clear_unsloth_compiled_cache, + _hw_module, + ) app = FastAPI( diff --git a/studio/backend/tests/test_lifespan_shutdown.py b/studio/backend/tests/test_lifespan_shutdown.py new file mode 100644 index 0000000000..0e3a34441d --- /dev/null +++ b/studio/backend/tests/test_lifespan_shutdown.py @@ -0,0 +1,134 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Regression tests for run_lifespan_shutdown: a dead default executor (the +abrupt-shutdown teardown race) must not abort the remaining cleanup. The helper +is dependency-injected, so these need only structlog.""" + +import asyncio +import contextvars +import types + +from utils.lifespan_shutdown import run_lifespan_shutdown + + +def _counter(): + box = {"n": 0} + + def _fn(): + box["n"] += 1 + + return box, _fn + + +def test_run_lifespan_shutdown_survives_dead_default_executor(): + term_box, terminate = _counter() + clear_box, clear = _counter() + hw = types.SimpleNamespace(DEVICE = "cuda:0") + + async def _drive(): + loop = asyncio.get_running_loop() + # Kill the default executor to mimic the teardown race. + await asyncio.to_thread(lambda: None) + loop._default_executor.shutdown(wait = True) + await run_lifespan_shutdown(terminate, clear, hw) + + asyncio.run(_drive()) + + assert term_box["n"] == 1, "terminate must run via inline fallback" + assert clear_box["n"] == 1, "clear must still run after the to_thread failure" + assert hw.DEVICE is None + + +def test_run_lifespan_shutdown_survives_shutdown_default_executor(): + """Production path: loop.shutdown_default_executor() makes run_in_executor raise + 'Executor shutdown has been called'; the helper must still recover inline.""" + term_box, terminate = _counter() + clear_box, clear = _counter() + hw = types.SimpleNamespace(DEVICE = "cuda:0") + + async def _drive(): + await asyncio.get_running_loop().shutdown_default_executor() + await run_lifespan_shutdown(terminate, clear, hw) + + asyncio.run(_drive()) + + assert term_box["n"] == 1, "terminate must run via inline fallback" + assert clear_box["n"] == 1 + assert hw.DEVICE is None + + +def test_run_lifespan_shutdown_normal_path(): + """Healthy executor: each step runs exactly once.""" + term_box, terminate = _counter() + clear_box, clear = _counter() + hw = types.SimpleNamespace(DEVICE = "cuda:0") + + asyncio.run(run_lifespan_shutdown(terminate, clear, hw)) + + assert term_box["n"] == 1 + assert clear_box["n"] == 1 + assert hw.DEVICE is None + + +def test_run_lifespan_shutdown_swallows_terminate_errors(): + """A terminate failure must not block later cleanup.""" + clear_box, clear = _counter() + hw = types.SimpleNamespace(DEVICE = "cuda:0") + + def _boom(): + raise ValueError("boom") + + asyncio.run(run_lifespan_shutdown(_boom, clear, hw)) + + assert clear_box["n"] == 1, "later cleanup must run even when terminate raises" + assert hw.DEVICE is None + + +def test_run_lifespan_shutdown_swallows_clear_errors(): + """A clear failure must not raise out of shutdown.""" + term_box, terminate = _counter() + hw = types.SimpleNamespace(DEVICE = "cuda:0") + + def _boom(): + raise ValueError("boom") + + asyncio.run(run_lifespan_shutdown(terminate, _boom, hw)) + + assert term_box["n"] == 1 + assert hw.DEVICE is None + + +def test_run_lifespan_shutdown_does_not_retry_body_runtime_error(): + """A body-side RuntimeError (healthy executor) must run terminate once, not retry inline.""" + term_box, _ = _counter() + clear_box, clear = _counter() + hw = types.SimpleNamespace(DEVICE = "cuda:0") + + def _boom(): + term_box["n"] += 1 + raise RuntimeError("body failed") + + asyncio.run(run_lifespan_shutdown(_boom, clear, hw)) + + assert term_box["n"] == 1, "body RuntimeError must not be retried inline" + assert clear_box["n"] == 1, "later cleanup must still run" + assert hw.DEVICE is None + + +def test_run_lifespan_shutdown_preserves_contextvars(): + """terminate runs in a copy of the caller's context (parity with asyncio.to_thread).""" + cv = contextvars.ContextVar("unsloth_test_cv") + cv.set("bound-value") + seen = [] + clear_box, clear = _counter() + hw = types.SimpleNamespace(DEVICE = "cuda:0") + + def terminate(): + seen.append(cv.get("UNSET")) + + asyncio.run(run_lifespan_shutdown(terminate, clear, hw)) + + assert seen == ["bound-value"], "terminate must run with the caller's contextvars" + assert clear_box["n"] == 1 + assert hw.DEVICE is None diff --git a/studio/backend/utils/lifespan_shutdown.py b/studio/backend/utils/lifespan_shutdown.py new file mode 100644 index 0000000000..6e8371e53d --- /dev/null +++ b/studio/backend/utils/lifespan_shutdown.py @@ -0,0 +1,56 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Resilient FastAPI lifespan shutdown cleanup. + +On an abrupt shutdown (Windows console-close, interpreter teardown racing +uvicorn) the loop's default executor may already be dead, so an unguarded +``asyncio.to_thread`` raise here would abort the nested-lifespan unwind and +surface as "Application shutdown failed". Dependency-injected so it can be +unit-tested without the heavy backend import graph. +""" + +import asyncio +import contextvars +import types +from typing import Callable + +import structlog + +logger = structlog.get_logger(__name__) + + +async def run_lifespan_shutdown( + terminate_downloads: Callable[[], None], + clear_compiled_cache: Callable[[], None], + hw_module: types.ModuleType, +) -> None: + """Run each shutdown step guarded so one failure can't skip the others; never raise.""" + loop = asyncio.get_running_loop() + # Copy context for parity with asyncio.to_thread. Schedule and await + # separately so a dead executor (raises at submit) runs inline, while a + # body exception (raised at await) is logged, not re-run. + ctx = contextvars.copy_context() + try: + future = loop.run_in_executor(None, ctx.run, terminate_downloads) + except RuntimeError: + # Executor gone: run inline on the loop thread. + try: + ctx.run(terminate_downloads) + except Exception as exc: + logger.warning("terminate_downloads (inline) failed at shutdown: %s", exc) + else: + try: + await future + except Exception as exc: + logger.warning("terminate_downloads failed at shutdown: %s", exc) + + try: + hw_module.DEVICE = None + except Exception as exc: + logger.warning("clearing hardware DEVICE failed at shutdown: %s", exc) + + try: + clear_compiled_cache() + except Exception as exc: + logger.warning("clear_compiled_cache failed at shutdown: %s", exc)