diff --git a/install.sh b/install.sh index 7a3c490368..7a2e18e374 100755 --- a/install.sh +++ b/install.sh @@ -3126,10 +3126,11 @@ echo "" if [ -t 1 ]; then echo "" printf " Start Unsloth Studio now? [Y/n] " + # No readable answer (closed/EOF tty) defaults to no; Enter is still yes. if [ -r /dev/tty ]; then - read -r _reply None: + for stream in (sys.stdout, sys.stderr): + try: + stream.flush() + except Exception: + pass + + +def _wait_for_server_shutdown(timeout: Optional[float] = _SERVER_SHUTDOWN_JOIN_TIMEOUT) -> None: + """Join the uvicorn thread so the prompt returns only after its shutdown logs + flush. Skip the self-join when called from the server thread.""" + import threading + + thread = _server_thread + if thread is None or thread is threading.current_thread(): + _flush_standard_streams() + return + thread.join(timeout = timeout) + if thread.is_alive(): + logger.warning("Timed out waiting for uvicorn server thread to stop") + _flush_standard_streams() + + # The uvicorn server instance -- set by run_server(), used by callers # that tell the server to exit (e.g. signal handlers). _server = None +_server_thread = None # Shutdown event -- wakes the main loop on signal. _shutdown_event = None @@ -909,7 +937,7 @@ def run_server( Signal handlers are NOT registered here so embedders (e.g. Colab) keep their own interrupt semantics; standalone callers register them after. """ - global _server, _shutdown_event + global _server, _server_thread, _shutdown_event # Reap every child if the parent dies abnormally (terminal close, Task # Manager kill, SIGKILL); must run before any child can spawn. @@ -1102,6 +1130,7 @@ def run_server( startup_failed.set() thread = Thread(target = _run, daemon = True) + _server_thread = thread thread.start() # Wait until uvicorn finishes lifespan startup and binds sockets, or until it @@ -1310,6 +1339,11 @@ if __name__ == "__main__": # Signal handler -- ensures subprocess cleanup on Ctrl+C. def _signal_handler(signum, frame): + # Restore defaults so a second signal force-quits if shutdown stalls. + signal.signal(signal.SIGINT, signal.SIG_DFL) + signal.signal(signal.SIGTERM, signal.SIG_DFL) + if hasattr(signal, "SIGBREAK"): + signal.signal(signal.SIGBREAK, signal.SIG_DFL) _graceful_shutdown(_server) _shutdown_event.set() @@ -1325,3 +1359,4 @@ if __name__ == "__main__": # lets the interpreter process pending signals. while not _shutdown_event.is_set(): _shutdown_event.wait(timeout = 1) + _wait_for_server_shutdown() diff --git a/tests/test_studio_shutdown_thread_wait.py b/tests/test_studio_shutdown_thread_wait.py new file mode 100644 index 0000000000..4ec2afc0f0 --- /dev/null +++ b/tests/test_studio_shutdown_thread_wait.py @@ -0,0 +1,133 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 + +"""Source-level regression tests for terminal shutdown ordering (no backend import).""" + +import ast +from pathlib import Path + +_ROOT = Path(__file__).resolve().parent.parent +_RUN_PY = _ROOT / "studio" / "backend" / "run.py" +_STUDIO_CLI_PY = _ROOT / "unsloth_cli" / "commands" / "studio.py" + + +def _parse(path: Path) -> ast.Module: + return ast.parse(path.read_text(encoding = "utf-8")) + + +def _function(tree: ast.AST, name: str) -> ast.FunctionDef: + for node in ast.walk(tree): + if isinstance(node, ast.FunctionDef) and node.name == name: + return node + raise AssertionError(f"missing function {name}") + + +def _calls_name(tree: ast.AST, name: str) -> int: + return sum( + 1 + for node in ast.walk(tree) + if isinstance(node, ast.Call) and isinstance(node.func, ast.Name) and node.func.id == name + ) + + +def _calls_shutdown_wait_getattr(tree: ast.AST) -> int: + count = 0 + for node in ast.walk(tree): + if not isinstance(node, ast.Call): + continue + func = node.func + if not isinstance(func, ast.Call): + continue + if not (isinstance(func.func, ast.Name) and func.func.id == "getattr"): + continue + if len(func.args) < 2: + continue + target, attr = func.args[:2] + if ( + isinstance(target, ast.Name) + and target.id == "run_mod" + and isinstance(attr, ast.Constant) + and attr.value == "_wait_for_server_shutdown" + ): + count += 1 + return count + + +def test_run_server_records_uvicorn_thread_for_terminal_shutdown_wait(): + tree = _parse(_RUN_PY) + run_server = _function(tree, "run_server") + + assigns_thread_global = any( + isinstance(node, ast.Assign) + and any( + isinstance(target, ast.Name) and target.id == "_server_thread" + for target in node.targets + ) + for node in ast.walk(run_server) + ) + + assert ( + assigns_thread_global + ), "run_server must retain the uvicorn thread so terminal shutdown can join it" + + +def test_wait_for_server_shutdown_joins_uvicorn_thread(): + tree = _parse(_RUN_PY) + wait_func = _function(tree, "_wait_for_server_shutdown") + + joins_thread = any( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and node.func.attr == "join" + for node in ast.walk(wait_func) + ) + + assert ( + joins_thread + ), "_wait_for_server_shutdown must join the uvicorn thread before process exit" + + +def test_wait_for_server_shutdown_join_is_bounded(): + tree = _parse(_RUN_PY) + wait_func = _function(tree, "_wait_for_server_shutdown") + + bounded_join = any( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and node.func.attr == "join" + and any(kw.arg == "timeout" for kw in node.keywords) + for node in ast.walk(wait_func) + ) + + assert ( + bounded_join + ), "the join must pass a timeout so a stalled uvicorn shutdown cannot hang the terminal" + + +def test_direct_backend_entrypoint_waits_before_returning_to_shell(): + tree = _parse(_RUN_PY) + + assert ( + _calls_name(tree, "_wait_for_server_shutdown") >= 1 + ), "run.py must wait after the main shutdown event loop before returning to the shell" + + +def test_signal_handler_restores_default_handlers_for_force_quit(): + tree = _parse(_RUN_PY) + handler = _function(tree, "_signal_handler") + + restores_default = any( + isinstance(node, ast.Attribute) and node.attr == "SIG_DFL" for node in ast.walk(handler) + ) + + assert ( + restores_default + ), "the signal handler must restore SIG_DFL so a second Ctrl+C can force-quit" + + +def test_cli_entrypoints_wait_before_returning_to_shell(): + tree = _parse(_STUDIO_CLI_PY) + + assert ( + _calls_shutdown_wait_getattr(tree) >= 3 + ), "Studio CLI terminal paths must wait for the backend thread after requesting shutdown" diff --git a/unsloth_cli/commands/studio.py b/unsloth_cli/commands/studio.py index 6ce6c5ea47..cde3d2cdbc 100644 --- a/unsloth_cli/commands/studio.py +++ b/unsloth_cli/commands/studio.py @@ -895,6 +895,8 @@ def studio_default( except KeyboardInterrupt: run_mod._graceful_shutdown(run_mod._server) typer.echo("\nShutting down...") + finally: + getattr(run_mod, "_wait_for_server_shutdown", lambda: None)() # ── unsloth studio run ─────────────────────────────────────────────── @@ -1304,6 +1306,7 @@ def run( raise typer.Exit(1) except BaseException: _graceful_shutdown(_server) + getattr(run_mod, "_wait_for_server_shutdown", lambda: None)() raise loaded_model = result.get("model", model) @@ -1410,6 +1413,8 @@ def run( except KeyboardInterrupt: run_mod._graceful_shutdown(run_mod._server) typer.echo("\nShutting down...") + finally: + getattr(run_mod, "_wait_for_server_shutdown", lambda: None)() # ── unsloth studio stop ───────────────────────────────────────────────