fix(gateway): make foreground shutdown responsive

This commit is contained in:
Xubin Ren
2026-06-22 22:00:07 +08:00
parent a9d1fdcee8
commit 747104c9cc
2 changed files with 273 additions and 15 deletions
+104 -15
View File
@@ -74,6 +74,62 @@ def _sanitize_surrogates(text: str) -> str:
return text.encode("utf-16-le", errors="surrogatepass").decode("utf-16-le", errors="replace") return text.encode("utf-16-le", errors="surrogatepass").decode("utf-16-le", errors="replace")
def _signal_name(signum: int) -> str:
with suppress(ValueError):
return signal.Signals(signum).name
return f"signal {signum}"
def _install_gateway_shutdown_handlers(
loop: asyncio.AbstractEventLoop,
shutdown_event: asyncio.Event,
tasks: list[asyncio.Task],
print_status: Callable[[str], None],
) -> Callable[[], None]:
"""Install foreground gateway signal handlers and return a restore callback."""
loop_signals: list[int] = []
previous_handlers: list[tuple[int, Any]] = []
shutdown_requested = False
def request_shutdown(signum: int) -> None:
nonlocal shutdown_requested
sig_name = _signal_name(signum)
if shutdown_requested:
logger.warning("Forcing gateway shutdown after repeated {}", sig_name)
for task in tasks:
if not task.done():
task.cancel()
return
shutdown_requested = True
logger.info("Gateway shutdown requested by {}", sig_name)
print_status("\nShutting down... Press Ctrl+C again to force.")
shutdown_event.set()
for signum in (signal.SIGINT, signal.SIGTERM):
try:
loop.add_signal_handler(signum, request_shutdown, signum)
except (NotImplementedError, RuntimeError, ValueError):
try:
previous = signal.getsignal(signum)
signal.signal(signum, lambda sig, _frame: request_shutdown(sig))
except (RuntimeError, ValueError):
logger.debug("Could not install gateway handler for {}", _signal_name(signum))
continue
previous_handlers.append((signum, previous))
else:
loop_signals.append(signum)
def restore() -> None:
for signum in loop_signals:
with suppress(NotImplementedError, RuntimeError, ValueError):
loop.remove_signal_handler(signum)
for signum, handler in previous_handlers:
with suppress(RuntimeError, ValueError):
signal.signal(signum, handler)
return restore
class SafeFileHistory(FileHistory): class SafeFileHistory(FileHistory):
"""FileHistory subclass that sanitizes surrogate characters on write. """FileHistory subclass that sanitizes surrogate characters on write.
@@ -1106,6 +1162,16 @@ def _run_gateway(
async def run(): async def run():
tasks: list[asyncio.Task] = [] tasks: list[asyncio.Task] = []
shutdown_task: asyncio.Task | None = None
runtime_tasks: asyncio.Future | None = None
runtime_tasks_drained = False
shutdown_event = asyncio.Event()
restore_shutdown_handlers = _install_gateway_shutdown_handlers(
asyncio.get_running_loop(),
shutdown_event,
tasks,
console.print,
)
try: try:
await cron.start() await cron.start()
tasks = [ tasks = [
@@ -1122,7 +1188,20 @@ def _run_gateway(
_open_browser_when_ready(), _open_browser_when_ready(),
name="nanobot-open-browser", name="nanobot-open-browser",
)) ))
await asyncio.gather(*tasks) runtime_tasks = asyncio.gather(*tasks)
shutdown_task = asyncio.create_task(
shutdown_event.wait(),
name="nanobot-gateway-shutdown",
)
done, _pending = await asyncio.wait(
{runtime_tasks, shutdown_task},
return_when=asyncio.FIRST_COMPLETED,
)
if runtime_tasks in done:
runtime_tasks_drained = True
await runtime_tasks
elif runtime_tasks is not None:
runtime_tasks.cancel()
except KeyboardInterrupt: except KeyboardInterrupt:
console.print("\nShutting down...") console.print("\nShutting down...")
except Exception: except Exception:
@@ -1131,20 +1210,30 @@ def _run_gateway(
console.print("\n[red]Error: Gateway crashed unexpectedly[/red]") console.print("\n[red]Error: Gateway crashed unexpectedly[/red]")
console.print(traceback.format_exc()) console.print(traceback.format_exc())
finally: finally:
cron.stop() try:
agent.stop() if shutdown_task and not shutdown_task.done():
for task in tasks: shutdown_task.cancel()
if not task.done(): with suppress(asyncio.CancelledError):
task.cancel() await shutdown_task
if tasks: cron.stop()
await asyncio.gather(*tasks, return_exceptions=True) agent.stop()
await channels.stop_all() for task in tasks:
# Flush all cached sessions to durable storage before exit. if not task.done():
# This prevents data loss on filesystems with write-back task.cancel()
# caching (rclone VFS, NFS, FUSE mounts, etc.). if tasks:
flushed = agent.sessions.flush_all() await asyncio.gather(*tasks, return_exceptions=True)
if flushed: if runtime_tasks is not None and not runtime_tasks_drained:
logger.info("Shutdown: flushed {} session(s) to disk", flushed) with suppress(asyncio.CancelledError, Exception):
await runtime_tasks
await channels.stop_all()
# Flush all cached sessions to durable storage before exit.
# This prevents data loss on filesystems with write-back
# caching (rclone VFS, NFS, FUSE mounts, etc.).
flushed = agent.sessions.flush_all()
if flushed:
logger.info("Shutdown: flushed {} session(s) to disk", flushed)
finally:
restore_shutdown_handlers()
asyncio.run(run()) asyncio.run(run())
+169
View File
@@ -2,6 +2,8 @@ import asyncio
import json import json
import re import re
import shutil import shutil
import signal
from contextlib import suppress
from pathlib import Path from pathlib import Path
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
@@ -9,6 +11,7 @@ import pytest
from typer.testing import CliRunner from typer.testing import CliRunner
from nanobot.bus.events import InboundMessage, OutboundMessage from nanobot.bus.events import InboundMessage, OutboundMessage
from nanobot.cli import commands as cli_commands
from nanobot.cli.commands import app from nanobot.cli.commands import app
from nanobot.config.schema import Config from nanobot.config.schema import Config
from nanobot.cron.service import CronJobSkippedError from nanobot.cron.service import CronJobSkippedError
@@ -58,6 +61,54 @@ class _StopGatewayError(RuntimeError):
pass pass
def test_gateway_signal_handler_first_signal_stops_and_second_forces() -> None:
class _FakeLoop:
def __init__(self) -> None:
self.handlers: dict[int, tuple[object, tuple[object, ...]]] = {}
self.removed: list[int] = []
def add_signal_handler(self, signum, callback, *args) -> None:
self.handlers[int(signum)] = (callback, args)
def remove_signal_handler(self, signum) -> bool:
self.removed.append(int(signum))
self.handlers.pop(int(signum), None)
return True
async def _run() -> None:
loop = _FakeLoop()
shutdown_event = asyncio.Event()
never = asyncio.Event()
task = asyncio.create_task(never.wait())
output: list[str] = []
restore = cli_commands._install_gateway_shutdown_handlers(
loop, shutdown_event, [task], output.append,
)
try:
callback, args = loop.handlers[int(signal.SIGINT)]
assert callable(callback)
callback(*args)
assert shutdown_event.is_set()
assert output == ["\nShutting down... Press Ctrl+C again to force."]
assert not task.done()
callback(*args)
await asyncio.sleep(0)
assert task.cancelled()
finally:
restore()
task.cancel()
with suppress(asyncio.CancelledError):
await task
assert int(signal.SIGINT) in loop.removed
assert int(signal.SIGTERM) in loop.removed
asyncio.run(_run())
@pytest.fixture @pytest.fixture
def mock_paths(): def mock_paths():
"""Mock config/workspace paths for test isolation.""" """Mock config/workspace paths for test isolation."""
@@ -2057,6 +2108,124 @@ def test_gateway_shutdown_lets_agent_task_own_mcp_cleanup(
assert seen["cron_stopped"] is True assert seen["cron_stopped"] is True
def test_gateway_shutdown_event_exits_forever_runtime_tasks(
monkeypatch,
tmp_path: Path,
) -> None:
config_file = _write_instance_config(tmp_path)
config = Config()
config.gateway.port = 18791
seen: dict[str, object] = {}
class _FakeSessionManager:
def flush_all(self) -> int:
return 0
class _FakeAgentLoop:
@classmethod
def from_config(cls, config, bus=None, **extra):
return cls(**extra)
def __init__(self, **_kwargs) -> None:
self.model = "test-model"
self.provider = object()
self.sessions = _FakeSessionManager()
def llm_runtime(self) -> None:
return None
async def run(self) -> None:
try:
await asyncio.Event().wait()
finally:
seen["agent_task_cleaned_up"] = True
async def close_mcp(self) -> None:
raise AssertionError("gateway must not close MCP from the outer task")
def stop(self) -> None:
seen["agent_stopped"] = True
class _FakeChannelManager:
def __init__(self, _config, _bus, **_kwargs) -> None:
self.enabled_channels = ["websocket"]
async def start_all(self) -> None:
try:
await asyncio.Event().wait()
finally:
seen["channel_task_cleaned_up"] = True
async def stop_all(self) -> None:
seen["channels_stopped"] = True
class _FakeCronService:
def __init__(self, _store_path: Path) -> None:
self.on_job = None
async def start(self) -> None:
return None
def stop(self) -> None:
seen["cron_stopped"] = True
def status(self) -> dict[str, int]:
return {"jobs": 0}
def register_system_job(self, _job) -> None:
return None
class _FakeServer:
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, exc, tb) -> bool:
return False
async def serve_forever(self) -> None:
await asyncio.Event().wait()
async def _fake_start_server(_handler, _host: str, _port: int):
return _FakeServer()
def _fake_install_shutdown_handlers(_loop, event, _tasks, _print_status):
async def _trigger_shutdown() -> None:
await asyncio.sleep(0)
event.set()
asyncio.create_task(_trigger_shutdown())
def _restore() -> None:
seen["shutdown_handlers_restored"] = True
return _restore
_patch_cli_command_runtime(
monkeypatch,
config,
message_bus=lambda: object(),
session_manager=lambda _workspace: object(),
)
monkeypatch.setattr("nanobot.cli.commands.AgentLoop", _FakeAgentLoop)
monkeypatch.setattr("nanobot.channels.manager.ChannelManager", _FakeChannelManager)
monkeypatch.setattr("nanobot.cron.service.CronService", _FakeCronService)
monkeypatch.setattr("asyncio.start_server", _fake_start_server)
monkeypatch.setattr(
"nanobot.cli.commands._install_gateway_shutdown_handlers",
_fake_install_shutdown_handlers,
)
result = runner.invoke(app, ["gateway", "--config", str(config_file)])
assert result.exit_code == 0
assert seen["agent_stopped"] is True
assert seen["agent_task_cleaned_up"] is True
assert seen["channel_task_cleaned_up"] is True
assert seen["channels_stopped"] is True
assert seen["cron_stopped"] is True
assert seen["shutdown_handlers_restored"] is True
def test_serve_uses_api_config_defaults_and_workspace_override( def test_serve_uses_api_config_defaults_and_workspace_override(
monkeypatch, tmp_path: Path monkeypatch, tmp_path: Path
) -> None: ) -> None: