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")
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):
"""FileHistory subclass that sanitizes surrogate characters on write.
@@ -1106,6 +1162,16 @@ def _run_gateway(
async def run():
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:
await cron.start()
tasks = [
@@ -1122,7 +1188,20 @@ def _run_gateway(
_open_browser_when_ready(),
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:
console.print("\nShutting down...")
except Exception:
@@ -1131,20 +1210,30 @@ def _run_gateway(
console.print("\n[red]Error: Gateway crashed unexpectedly[/red]")
console.print(traceback.format_exc())
finally:
cron.stop()
agent.stop()
for task in tasks:
if not task.done():
task.cancel()
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
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)
try:
if shutdown_task and not shutdown_task.done():
shutdown_task.cancel()
with suppress(asyncio.CancelledError):
await shutdown_task
cron.stop()
agent.stop()
for task in tasks:
if not task.done():
task.cancel()
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
if runtime_tasks is not None and not runtime_tasks_drained:
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())
+169
View File
@@ -2,6 +2,8 @@ import asyncio
import json
import re
import shutil
import signal
from contextlib import suppress
from pathlib import Path
from unittest.mock import AsyncMock, MagicMock, patch
@@ -9,6 +11,7 @@ import pytest
from typer.testing import CliRunner
from nanobot.bus.events import InboundMessage, OutboundMessage
from nanobot.cli import commands as cli_commands
from nanobot.cli.commands import app
from nanobot.config.schema import Config
from nanobot.cron.service import CronJobSkippedError
@@ -58,6 +61,54 @@ class _StopGatewayError(RuntimeError):
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
def mock_paths():
"""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
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(
monkeypatch, tmp_path: Path
) -> None: