diff --git a/headroom/memory/mcp_server.py b/headroom/memory/mcp_server.py index 52830ff06..2a42dc7ab 100644 --- a/headroom/memory/mcp_server.py +++ b/headroom/memory/mcp_server.py @@ -161,34 +161,75 @@ def create_memory_server(db_path: str, user_id: str = "default") -> Server: server = Server("headroom-memory") _backend: LocalBackend | None = None - _init_task: asyncio.Task | None = None + _init_task: asyncio.Task[LocalBackend] | None = None async def _init_backend() -> LocalBackend: """Initialize backend with ONNX embedder (fast, no PyTorch).""" - nonlocal _backend + nonlocal _backend, _init_task config = LocalBackendConfig(db_path=db_path, embedder_backend="onnx") - _backend = LocalBackend(config) - await _warm_up_backend(_backend, user_id) + backend = LocalBackend(config) + init_task = asyncio.current_task() + try: + await _warm_up_backend(backend, user_id) + except (Exception, asyncio.CancelledError): + try: + await backend.close() + except Exception as cleanup_error: + logger.warning("Memory MCP: failed backend cleanup: %s", cleanup_error) + finally: + if _init_task is init_task: + _init_task = None + raise + + _backend = backend + if _init_task is init_task: + _init_task = None logger.info(f"Memory MCP: ready (db={db_path}, user={user_id})") - return _backend + return backend + + def _handle_backend_init_done(init_task: asyncio.Task[LocalBackend]) -> None: + """Clear and log failed background initialization tasks.""" + nonlocal _init_task + if init_task.cancelled(): + if _init_task is init_task: + _init_task = None + return + error = init_task.exception() + if error is not None: + if _init_task is init_task: + _init_task = None + logger.warning("Memory MCP: backend initialization failed: %s", error) + + def _start_backend_init() -> asyncio.Task[LocalBackend]: + """Start backend initialization once for all concurrent callers.""" + nonlocal _init_task + if _init_task is None: + _init_task = asyncio.create_task(_init_backend()) + _init_task.add_done_callback(_handle_backend_init_done) + return _init_task async def _get_backend() -> LocalBackend: nonlocal _backend, _init_task if _backend is not None: return _backend - # Wait for background init if it's running - if _init_task is not None: - await _init_task - return _backend # type: ignore[return-value] - # Fallback: init inline (shouldn't normally happen) - return await _init_backend() + + init_task = _start_backend_init() + try: + return await asyncio.shield(init_task) + except asyncio.CancelledError: + if init_task.done() and _init_task is init_task: + _init_task = None + raise + except Exception: + if _init_task is init_task: + _init_task = None + raise @server.list_tools() async def list_tools() -> list[Tool]: # Kick off background init on first list_tools (called at MCP handshake) - nonlocal _init_task - if _backend is None and _init_task is None: - _init_task = asyncio.create_task(_init_backend()) + if _backend is None: + _start_backend_init() return _TOOLS @server.call_tool() diff --git a/tests/test_memory/test_mcp_server.py b/tests/test_memory/test_mcp_server.py index b86e15179..cd2779627 100644 --- a/tests/test_memory/test_mcp_server.py +++ b/tests/test_memory/test_mcp_server.py @@ -10,6 +10,27 @@ from tests._mcp_stub import import_module_with_mcp_stub mcp_server_mod = import_module_with_mcp_stub("headroom.memory.mcp_server") +class _CapturingServer: + def __init__(self, name: str) -> None: + self.name = name + self.list_tools_handler = None + self.call_tool_handler = None + + def list_tools(self): + def decorator(handler): + self.list_tools_handler = handler + return handler + + return decorator + + def call_tool(self): + def decorator(handler): + self.call_tool_handler = handler + return handler + + return decorator + + def test_warm_up_backend_batches_embedding_and_indexing() -> None: """Warm-up should batch missing embeddings and vector indexing.""" warmup_embedding = np.ones(384, dtype=np.float32) @@ -62,6 +83,142 @@ def test_warm_up_backend_batches_embedding_and_indexing() -> None: assert np.array_equal(memory_without_embedding_b.embedding, batch_embeddings[1]) +def test_tool_call_waits_for_handshake_warm_up(monkeypatch) -> None: + async def scenario() -> None: + warm_up_started = asyncio.Event() + release_warm_up = asyncio.Event() + backend = SimpleNamespace() + + async def warm_up(candidate, user_id: str) -> None: + assert candidate is backend + assert user_id == "alice" + warm_up_started.set() + await release_warm_up.wait() + + handle_search = AsyncMock(return_value=["search result"]) + 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) + monkeypatch.setattr(mcp_server_mod, "_handle_search", handle_search) + + server = mcp_server_mod.create_memory_server("memory.db", user_id="alice") + await server.list_tools_handler() + await warm_up_started.wait() + + tool_call = asyncio.create_task( + server.call_tool_handler("memory_search", {"query": "preferences"}) + ) + await asyncio.sleep(0) + + handle_search.assert_not_awaited() + assert not tool_call.done() + + release_warm_up.set() + assert await tool_call == ["search result"] + handle_search.assert_awaited_once_with( + backend, + {"query": "preferences"}, + "alice", + ) + + asyncio.run(scenario()) + + +def test_failed_handshake_init_is_discarded_and_retried(monkeypatch) -> None: + async def scenario() -> None: + first_warm_up_started = asyncio.Event() + fail_first_warm_up = asyncio.Event() + failed_backend_closed = asyncio.Event() + + async def close_failed_backend() -> None: + failed_backend_closed.set() + + failed_backend = SimpleNamespace(close=AsyncMock(side_effect=close_failed_backend)) + ready_backend = SimpleNamespace(close=AsyncMock()) + backends = iter([failed_backend, ready_backend]) + + async def warm_up(candidate, user_id: str) -> None: + assert user_id == "alice" + if candidate is failed_backend: + first_warm_up_started.set() + await fail_first_warm_up.wait() + raise RuntimeError("warm-up failed") + + handle_search = AsyncMock(return_value=["search result"]) + monkeypatch.setattr(mcp_server_mod, "Server", _CapturingServer) + monkeypatch.setattr(mcp_server_mod, "LocalBackend", lambda config: next(backends)) + monkeypatch.setattr(mcp_server_mod, "_warm_up_backend", warm_up) + monkeypatch.setattr(mcp_server_mod, "_handle_search", handle_search) + + server = mcp_server_mod.create_memory_server("memory.db", user_id="alice") + await server.list_tools_handler() + await first_warm_up_started.wait() + + fail_first_warm_up.set() + await failed_backend_closed.wait() + await asyncio.sleep(0) + + failed_backend.close.assert_awaited_once() + handle_search.assert_not_awaited() + + assert await server.call_tool_handler("memory_search", {"query": "preferences"}) == [ + "search result" + ] + handle_search.assert_awaited_once_with( + ready_backend, + {"query": "preferences"}, + "alice", + ) + ready_backend.close.assert_not_awaited() + + asyncio.run(scenario()) + + +def test_concurrent_tool_calls_share_backend_initialization(monkeypatch) -> None: + async def scenario() -> None: + warm_up_started = asyncio.Event() + release_warm_up = asyncio.Event() + backend = SimpleNamespace() + created_backends = 0 + + def create_backend(config): + nonlocal created_backends + created_backends += 1 + return backend + + async def warm_up(candidate, user_id: str) -> None: + assert candidate is backend + assert user_id == "alice" + warm_up_started.set() + await release_warm_up.wait() + + handle_search = AsyncMock(return_value=["search result"]) + monkeypatch.setattr(mcp_server_mod, "Server", _CapturingServer) + monkeypatch.setattr(mcp_server_mod, "LocalBackend", create_backend) + monkeypatch.setattr(mcp_server_mod, "_warm_up_backend", warm_up) + monkeypatch.setattr(mcp_server_mod, "_handle_search", handle_search) + + server = mcp_server_mod.create_memory_server("memory.db", user_id="alice") + calls = [ + asyncio.create_task( + server.call_tool_handler("memory_search", {"query": f"query-{index}"}) + ) + for index in range(2) + ] + await warm_up_started.wait() + await asyncio.sleep(0) + + assert created_backends == 1 + handle_search.assert_not_awaited() + + release_warm_up.set() + assert await asyncio.gather(*calls) == [["search result"], ["search result"]] + assert handle_search.await_count == 2 + assert all(call.args[0] is backend for call in handle_search.await_args_list) + + 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()