diff --git a/CHANGELOG.md b/CHANGELOG.md index 9b87fefe93..cf66258667 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Changed +- `batch_fetch_server_info` now looks up independent MCP registry documents concurrently with a bounded `ThreadPoolExecutor` (max 4 workers), matching the existing install-check fan-out. Closes #2981. (#2983) - `docs/src/content/docs/specs/openapm-v0.1.md` adds proposed `req-tg-015` for Codex-native `model`/`model_reasoning_effort` preservation and bounded dropped-metadata diagnostics. (#3150) ### Fixed diff --git a/src/apm_cli/registry/operations.py b/src/apm_cli/registry/operations.py index ce4f3664fc..bd8a6d5fa7 100644 --- a/src/apm_cli/registry/operations.py +++ b/src/apm_cli/registry/operations.py @@ -289,23 +289,40 @@ def _validate_one(server_ref: str) -> tuple[str, bool, dict | None]: return valid_servers, invalid_servers - def batch_fetch_server_info(self, server_references: list[str]) -> dict[str, dict | None]: + def batch_fetch_server_info( + self, + server_references: list[str], + max_workers: int = 4, + ) -> dict[str, dict | None]: """Batch fetch server info for all servers to avoid duplicate registry calls. + Each registry lookup is independent. Lookups run in a bounded + ThreadPoolExecutor (same pattern as ``check_servers_needing_installation``) + and results are collected in submission order via ``executor.map``. + Args: server_references: List of MCP server references + max_workers: Requested parallel lookups; clamped to 4. Returns: Dictionary mapping server reference to server info (or None if not found) """ - server_info_cache = {} + from concurrent.futures import ThreadPoolExecutor - for server_ref in server_references: + def _fetch_one(server_ref: str) -> tuple[str, dict | None]: try: - server_info = self.registry_client.find_server_by_reference(server_ref) - server_info_cache[server_ref] = server_info + return (server_ref, self.registry_client.find_server_by_reference(server_ref)) except Exception: - server_info_cache[server_ref] = None + return (server_ref, None) + + server_info_cache: dict[str, dict | None] = {} + if not server_references: + return server_info_cache + + workers = min(max(1, max_workers), 4, len(server_references)) + with ThreadPoolExecutor(max_workers=workers, thread_name_prefix="mcp-fetch") as executor: + for server_ref, server_info in executor.map(_fetch_one, server_references): + server_info_cache[server_ref] = server_info return server_info_cache diff --git a/tests/unit/integration/test_mcp_registry_parallel.py b/tests/unit/integration/test_mcp_registry_parallel.py index 142e3594cb..065ffb2b68 100644 --- a/tests/unit/integration/test_mcp_registry_parallel.py +++ b/tests/unit/integration/test_mcp_registry_parallel.py @@ -1,7 +1,7 @@ """WS2b (#1116): parallel MCP registry batch lookup tests. -Verifies that ``validate_servers_exist`` and ``check_servers_needing_installation`` -run in parallel and complete within bounded wall time. +Verifies that ``validate_servers_exist``, ``check_servers_needing_installation``, +and ``batch_fetch_server_info`` run in parallel and complete within bounded wall time. No real network calls -- all registry HTTP is mocked. """ @@ -9,7 +9,7 @@ from __future__ import annotations import time -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch from apm_cli.registry.operations import MCPServerOperations @@ -116,3 +116,74 @@ def test_validate_single_server_does_not_error(self) -> None: valid, invalid = ops.validate_servers_exist(["only-one"], max_workers=4) assert valid == ["only-one"] assert invalid == [] + + def test_batch_fetch_server_info_parallel_wall_time(self) -> None: + """3 servers each sleeping 500ms: wall time < 1.0s (vs 1.5s serial).""" + ops = MCPServerOperations.__new__(MCPServerOperations) + ops.registry_client = MagicMock() + + call_count = {"n": 0} + + def slow_find(ref: str): + import time as _t + + call_count["n"] += 1 + _t.sleep(0.5) + return {"id": f"uuid-{ref}", "name": ref} + + ops.registry_client.find_server_by_reference = slow_find + + servers = ["server-a", "server-b", "server-c"] + + start = time.monotonic() + result = ops.batch_fetch_server_info(servers, max_workers=4) + elapsed = time.monotonic() - start + + assert call_count["n"] == 3 + assert list(result.keys()) == servers + assert all(result[ref]["id"] == f"uuid-{ref}" for ref in servers) + assert elapsed < 1.0, f"Wall time {elapsed:.3f}s >= 1.0s (not parallel)" + + def test_batch_fetch_preserves_submission_order_and_exceptions(self) -> None: + """Results keep input order; per-ref exceptions map to None.""" + ops = MCPServerOperations.__new__(MCPServerOperations) + ops.registry_client = MagicMock() + + import random + + def jittered_find(ref: str): + import time as _t + + _t.sleep(random.uniform(0.01, 0.05)) # noqa: S311 + if ref == "bad": + raise RuntimeError("lookup failed") + return {"id": f"uuid-{ref}", "name": ref} + + ops.registry_client.find_server_by_reference = jittered_find + + servers = ["alpha", "bad", "gamma", "delta"] + result = ops.batch_fetch_server_info(servers, max_workers=4) + + assert list(result.keys()) == servers + assert result["alpha"]["id"] == "uuid-alpha" + assert result["bad"] is None + assert result["gamma"]["id"] == "uuid-gamma" + assert result["delta"]["id"] == "uuid-delta" + + def test_batch_fetch_clamps_max_workers_to_four(self) -> None: + """Callers cannot raise the pool past the established four-worker cap.""" + ops = MCPServerOperations.__new__(MCPServerOperations) + ops.registry_client = MagicMock() + ops.registry_client.find_server_by_reference = lambda ref: {"id": ref} + + with patch("concurrent.futures.ThreadPoolExecutor") as mock_pool: + mock_pool.return_value.__enter__.return_value.map = lambda fn, refs: ( + fn(ref) for ref in refs + ) + result = ops.batch_fetch_server_info( + ["a", "b", "c", "d", "e"], + max_workers=1000, + ) + + assert mock_pool.call_args.kwargs["max_workers"] == 4 + assert list(result.keys()) == ["a", "b", "c", "d", "e"]