diff --git a/headroom/memory/mcp_server.py b/headroom/memory/mcp_server.py index 2a42dc7ab..3720b3e4f 100644 --- a/headroom/memory/mcp_server.py +++ b/headroom/memory/mcp_server.py @@ -162,6 +162,7 @@ def create_memory_server(db_path: str, user_id: str = "default") -> Server: server = Server("headroom-memory") _backend: LocalBackend | None = None _init_task: asyncio.Task[LocalBackend] | None = None + _close_lock = asyncio.Lock() async def _init_backend() -> LocalBackend: """Initialize backend with ONNX embedder (fast, no PyTorch).""" @@ -225,6 +226,26 @@ def create_memory_server(db_path: str, user_id: str = "default") -> Server: _init_task = None raise + async def _close_backend() -> None: + """Cancel backend initialization and close any initialized backend.""" + nonlocal _backend, _init_task + async with _close_lock: + init_task = _init_task + if init_task is not None: + if not init_task.done(): + init_task.cancel() + await asyncio.gather(init_task, return_exceptions=True) + if _init_task is init_task: + _init_task = None + + backend = _backend + _backend = None + if backend is not None: + try: + await backend.close() + except Exception as cleanup_error: + logger.warning("Memory MCP: failed backend cleanup: %s", cleanup_error) + @server.list_tools() async def list_tools() -> list[Tool]: # Kick off background init on first list_tools (called at MCP handshake) @@ -243,6 +264,7 @@ def create_memory_server(db_path: str, user_id: str = "default") -> Server: return [TextContent(type="text", text=f"Unknown tool: {name}")] + server._headroom_close = _close_backend # type: ignore[attr-defined] return server @@ -357,8 +379,13 @@ async def _handle_save( async def _run(db_path: str, user_id: str) -> None: server = create_memory_server(db_path, user_id) - async with stdio_server() as (read_stream, write_stream): - await server.run(read_stream, write_stream, server.create_initialization_options()) + try: + async with stdio_server() as (read_stream, write_stream): + await server.run(read_stream, write_stream, server.create_initialization_options()) + finally: + close_backend = getattr(server, "_headroom_close", None) + if close_backend is not None: + await close_backend() def _memory_mcp_startup_context( diff --git a/tests/test_memory/test_mcp_server.py b/tests/test_memory/test_mcp_server.py index cd2779627..b55f28218 100644 --- a/tests/test_memory/test_mcp_server.py +++ b/tests/test_memory/test_mcp_server.py @@ -219,6 +219,77 @@ def test_concurrent_tool_calls_share_backend_initialization(monkeypatch) -> None asyncio.run(scenario()) +def test_server_cleanup_closes_initialized_backend_once(monkeypatch) -> None: + async def scenario() -> None: + backend = SimpleNamespace(close=AsyncMock()) + monkeypatch.setattr(mcp_server_mod, "Server", _CapturingServer) + monkeypatch.setattr(mcp_server_mod, "LocalBackend", lambda config: backend) + monkeypatch.setattr(mcp_server_mod, "_warm_up_backend", AsyncMock()) + + server = mcp_server_mod.create_memory_server("memory.db", user_id="alice") + await server.list_tools_handler() + await asyncio.sleep(0) + + close_backend = server._headroom_close + await close_backend() + await close_backend() + + backend.close.assert_awaited_once() + + asyncio.run(scenario()) + + +def test_server_cleanup_cancels_pending_backend_initialization(monkeypatch) -> None: + async def scenario() -> None: + init_started = asyncio.Event() + backend = SimpleNamespace(close=AsyncMock()) + + async def warm_up(_backend, _user_id: str) -> None: + init_started.set() + await asyncio.Event().wait() + + monkeypatch.setattr(mcp_server_mod, "Server", _CapturingServer) + monkeypatch.setattr(mcp_server_mod, "LocalBackend", lambda config: backend) + monkeypatch.setattr(mcp_server_mod, "_warm_up_backend", warm_up) + + server = mcp_server_mod.create_memory_server("memory.db", user_id="alice") + await server.list_tools_handler() + await init_started.wait() + + await server._headroom_close() + + backend.close.assert_awaited_once() + + asyncio.run(scenario()) + + +def test_run_closes_backend_when_stdio_exits(monkeypatch) -> None: + async def scenario() -> None: + close_backend = AsyncMock() + server = SimpleNamespace( + create_initialization_options=lambda: {}, + run=AsyncMock(), + _headroom_close=close_backend, + ) + + class _StdioContext: + async def __aenter__(self): + return object(), object() + + async def __aexit__(self, exc_type, exc_value, traceback): + return False + + monkeypatch.setattr(mcp_server_mod, "create_memory_server", lambda *args: server) + monkeypatch.setattr(mcp_server_mod, "stdio_server", lambda: _StdioContext()) + + await mcp_server_mod._run("memory.db", "alice") + + server.run.assert_awaited_once() + close_backend.assert_awaited_once() + + asyncio.run(scenario()) + + def test_memory_mcp_startup_context_reports_dynamic_project_db(tmp_path) -> None: project_dir = tmp_path / "project-a" project_dir.mkdir()