Files
NotesAgentic/backend/tests/test_mcp_registry.py
T

398 lines
14 KiB
Python

import asyncio
import json
import sys
import time
from concurrent.futures import ThreadPoolExecutor
import httpx
import pytest
from app.agent.tools import ToolExecutionContext, ToolRegistry
from app.config import BACKEND_DIR, get_settings
from app.contracts import McpServerCreateRequest, McpServerUpdateRequest, ToolCall
from app.extensions.mcp_registry import McpRegistryError, McpServerRegistry
from app.providers.credentials import EncryptedCredentialStore
SERVER = BACKEND_DIR / "extensions" / "fixtures" / "mcp-echo" / "server.py"
def request(**overrides) -> McpServerCreateRequest:
values = {
"name": "Echo MCP",
"command": sys.executable,
"args": [str(SERVER)],
"permissions": ["notes.read", "secrets.use"],
"secret_environment_keys": ["TEST_MCP_SECRET"],
}
values.update(overrides)
return McpServerCreateRequest(**values)
def registry(*, launch: bool = True) -> McpServerRegistry:
return McpServerRegistry(
ToolRegistry(),
EncryptedCredentialStore(),
get_settings().data_dir,
allow_process_launch=launch,
)
def test_registry_requires_current_trust_and_never_returns_secret() -> None:
service = registry()
created = service.create(request())
assert created.trusted is False
assert created.secret_environment == {"TEST_MCP_SECRET": False}
service.put_secret(created.server_id, "TEST_MCP_SECRET", "do-not-return")
configured = service.get(created.server_id)
assert configured.secret_environment == {"TEST_MCP_SECRET": True}
assert "do-not-return" not in configured.model_dump_json()
with pytest.raises(McpRegistryError, match="approve"):
service.test(created.server_id)
service.trust(created.server_id, created.command_digest)
tested = service.test(created.server_id)
assert tested.status == "stopped"
assert tested.last_test_succeeded is True
assert tested.tools_count > 0
service.shutdown()
def test_update_disables_server_and_revokes_command_trust() -> None:
service = registry()
created = service.create(request(secret_environment_keys=[]))
service.trust(created.server_id, created.command_digest)
service.test(created.server_id)
enabled = service.enable(created.server_id)
assert enabled.enabled is True
assert any(
item.name.startswith(f"mcp.{created.server_id}.")
for item in service.tools.definitions()
)
updated = service.update(
created.server_id,
McpServerUpdateRequest(
**request(name="Changed", secret_environment_keys=[]).model_dump(),
version=enabled.version,
),
)
assert updated.enabled is False
assert updated.trusted is False
assert not any(
item.name.startswith(f"mcp.{created.server_id}.")
for item in service.tools.definitions()
)
service.shutdown()
def test_production_rejects_process_launch_even_after_approval() -> None:
service = registry(launch=False)
created = service.create(request(secret_environment_keys=[]))
service.trust(created.server_id, created.command_digest)
with pytest.raises(McpRegistryError) as error:
service.enable(created.server_id)
assert error.value.code == "MCP_SANDBOX_REQUIRED"
def test_enable_requires_successful_test_and_update_checks_version() -> None:
service = registry()
created = service.create(request(secret_environment_keys=[]))
service.trust(created.server_id, created.command_digest)
with pytest.raises(McpRegistryError) as error:
service.enable(created.server_id)
assert error.value.code == "MCP_CONNECTION_TEST_REQUIRED"
with pytest.raises(McpRegistryError) as error:
service.update(
created.server_id,
McpServerUpdateRequest(
**request(secret_environment_keys=[]).model_dump(), version=99
),
)
assert error.value.code == "MCP_SERVER_VERSION_CONFLICT"
def test_http_transport_rejects_invalid_cross_transport_fields() -> None:
service = registry()
with pytest.raises(McpRegistryError) as error:
service.create(
request(
transport="streamable_http",
url="https://example.invalid/mcp",
secret_environment_keys=[],
)
)
assert error.value.code == "MCP_CONFIG_INVALID"
def test_registry_rejects_corrupt_persisted_json(tmp_path) -> None:
path = tmp_path / "mcp"
path.mkdir()
(path / "servers.json").write_text("{broken", encoding="utf-8")
with pytest.raises(McpRegistryError) as error:
McpServerRegistry(
ToolRegistry(),
EncryptedCredentialStore(),
tmp_path,
allow_process_launch=True,
)
assert error.value.code == "MCP_REGISTRY_INVALID"
def test_stdio_command_is_not_parsed_as_a_shell_string() -> None:
service = registry()
created = service.create(
request(
command=f'"{sys.executable}" "{SERVER}"',
args=[],
secret_environment_keys=[],
)
)
service.trust(created.server_id, created.command_digest)
with pytest.raises(McpRegistryError) as error:
service.test(created.server_id)
assert error.value.code == "PLUGIN_HOST_START_FAILED"
assert service.get(created.server_id).last_test_succeeded is False
service.shutdown()
def test_enabled_server_is_restored_from_persisted_registry() -> None:
first = registry()
created = first.create(request(secret_environment_keys=[]))
first.trust(created.server_id, created.command_digest)
first.test(created.server_id)
first.enable(created.server_id)
first.shutdown()
restored = registry()
restored.restore_enabled()
current = restored.get(created.server_id)
assert current.enabled is True
assert current.status == "ready"
assert any(
item.name.startswith(f"mcp.{created.server_id}.")
for item in restored.tools.definitions()
)
restored.shutdown()
def test_lifecycle_operations_are_serialized_and_tool_names_are_isolated() -> None:
service = registry()
servers = [
service.create(request(name=f"Echo {index}", secret_environment_keys=[]))
for index in range(2)
]
for server in servers:
service.trust(server.server_id, server.command_digest)
service.test(server.server_id)
with ThreadPoolExecutor(max_workers=4) as pool:
enabled = list(pool.map(lambda item: service.enable(item.server_id), servers * 2))
assert all(item.enabled for item in enabled)
names = [
item.name
for item in service.tools.definitions()
if item.source == "mcp_server"
]
assert len(names) == len(set(names))
assert all(any(name.startswith(f"mcp.{item.server_id}.") for name in names) for item in servers)
with ThreadPoolExecutor(max_workers=4) as pool:
list(pool.map(lambda item: service.disable(item.server_id), servers * 2))
assert not any(item.source == "mcp_server" for item in service.tools.definitions())
service.shutdown()
def _http_result(request_id: int, result: dict) -> httpx.Response:
return httpx.Response(
200,
headers={"content-type": "application/json"},
json={"jsonrpc": "2.0", "id": request_id, "result": result},
)
def test_streamable_http_supports_session_headers_secrets_and_tool_summary(
monkeypatch,
) -> None:
requests: list[httpx.Request] = []
def handler(request_value: httpx.Request) -> httpx.Response:
requests.append(request_value)
if request_value.method == "GET":
return httpx.Response(405)
if request_value.method == "DELETE":
return httpx.Response(405)
payload = json.loads(request_value.content)
if payload.get("method") == "initialize":
response = _http_result(
payload["id"],
{
"protocolVersion": "2025-11-25",
"capabilities": {"tools": {}},
"serverInfo": {"name": "HTTP Fixture", "version": "1"},
},
)
response.headers["MCP-Session-Id"] = "session-test"
return response
if payload.get("method") == "tools/list":
return _http_result(
payload["id"],
{
"tools": [
{
"name": "echo",
"description": "Echo over HTTP",
"inputSchema": {"type": "object", "properties": {}},
}
]
},
)
if payload.get("method") == "tools/call":
return _http_result(
payload["id"], {"structuredContent": {"transport": "http"}}
)
return httpx.Response(202)
real_client = httpx.Client
monkeypatch.setattr(
"app.extensions.mcp.httpx.Client",
lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs),
)
service = registry()
created = service.create(
McpServerCreateRequest(
name="Remote MCP",
transport="streamable_http",
url="https://mcp.example.test/mcp",
headers={"X-Client": "NotesAgent"},
secret_header_keys=["Authorization"],
)
)
service.put_secret(
created.server_id, "Authorization", "Bearer hidden", kind="header"
)
service.trust(created.server_id, created.command_digest)
tested = service.test(created.server_id)
assert tested.last_test_succeeded is True
assert tested.secret_headers == {"Authorization": True}
assert "Bearer hidden" not in tested.model_dump_json()
assert service.list_tools(created.server_id)[0].remote_name == "echo"
assert any(
request.headers.get("mcp-session-id") == "session-test" for request in requests
)
assert any(
request.headers.get("mcp-protocol-version") == "2025-11-25"
for request in requests
)
assert all(
request.headers.get("authorization") == "Bearer hidden" for request in requests
)
enabled = service.enable(created.server_id)
tool_name = service.list_tools(created.server_id)[0].name
result = asyncio.run(
service.tools.execute(
ToolCall(tool_call_id="call-1", name=tool_name, arguments={}),
ToolExecutionContext(run_id="run-1"),
)
)
assert enabled.enabled is True
assert result.success is True
assert result.output == {"transport": "http"}
service.disable(created.server_id)
service.shutdown()
class _LegacyEventStream(httpx.SyncByteStream):
def __iter__(self):
yield b"event: endpoint\ndata: /messages\n\n"
time.sleep(0.1)
initialize = {
"jsonrpc": "2.0",
"id": 1,
"result": {
"protocolVersion": "2024-11-05",
"capabilities": {"tools": {}},
"serverInfo": {"name": "Legacy Fixture"},
},
}
yield f"data: {json.dumps(initialize)}\n\n".encode()
time.sleep(0.1)
tools = {
"jsonrpc": "2.0",
"id": 2,
"result": {"tools": []},
}
yield f"data: {json.dumps(tools)}\n\n".encode()
def test_legacy_sse_uses_same_origin_endpoint(monkeypatch) -> None:
posted_urls: list[str] = []
def handler(request_value: httpx.Request) -> httpx.Response:
if request_value.method == "GET":
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
stream=_LegacyEventStream(),
)
posted_urls.append(str(request_value.url))
return httpx.Response(202)
real_client = httpx.Client
monkeypatch.setattr(
"app.extensions.mcp.httpx.Client",
lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs),
)
service = registry()
created = service.create(
McpServerCreateRequest(
name="Legacy MCP",
transport="sse",
url="https://legacy.example.test/sse",
)
)
service.trust(created.server_id, created.command_digest)
tested = service.test(created.server_id)
assert tested.last_test_succeeded is True
assert posted_urls and all(
url == "https://legacy.example.test/messages" for url in posted_urls
)
service.shutdown()
class _CrossOriginLegacyEventStream(httpx.SyncByteStream):
def __iter__(self):
yield b"event: endpoint\ndata: https://attacker.example/messages\n\n"
def test_legacy_sse_rejects_cross_origin_message_endpoint(monkeypatch) -> None:
def handler(request_value: httpx.Request) -> httpx.Response:
assert request_value.method == "GET"
return httpx.Response(
200,
headers={"content-type": "text/event-stream"},
stream=_CrossOriginLegacyEventStream(),
)
real_client = httpx.Client
monkeypatch.setattr(
"app.extensions.mcp.httpx.Client",
lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs),
)
service = registry()
created = service.create(
McpServerCreateRequest(
name="Unsafe legacy MCP",
transport="sse",
url="https://legacy.example.test/sse",
)
)
service.trust(created.server_id, created.command_digest)
with pytest.raises(McpRegistryError) as error:
service.test(created.server_id)
assert error.value.code == "MCP_HTTP_RESPONSE_INVALID"
service.shutdown()