Files

108 lines
4.4 KiB
Python
Raw Permalink Normal View History

2026-07-29 17:11:54 +08:00
"""Tests for the MCP endpoint's DNS-rebinding (Host header) protection.
The MCP SDK auto-enables DNS-rebinding protection and derives its allow-list
from the bind host, which defaults to 127.0.0.1. Left alone, that rejects every
request whose Host header is the real deployment address (6.86.80.9:8000) with
HTTP 421 — before the auth middleware or the tool ever runs. These tests pin
the configured allow-list behavior so that failure mode cannot come back.
"""
from __future__ import annotations
import json
from contextlib import contextmanager
from unittest.mock import patch
from starlette.testclient import TestClient
from app.mcp.server import _build_transport_security, build_mcp_asgi_app
# A minimal JSON-RPC initialize call. Reaching the MCP handler at all is what
# matters here; transport security rejects the request long before this body is
# parsed, so its exact contents only need to be structurally valid.
_INITIALIZE = {
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-06-18",
"capabilities": {},
"clientInfo": {"name": "test", "version": "1.0"},
},
}
_HEADERS = {
"Content-Type": "application/json",
"Accept": "application/json, text/event-stream",
}
@contextmanager
def _mcp_client(allowed_hosts: str):
"""Yield a TestClient over the real MCP app with auth off and hosts configured.
The settings patch must stay active for the requests themselves, not just
for app construction, because MCPAuthMiddleware reads settings per request.
Entering the TestClient as a context manager is also required: it runs the
app's lifespan, without which the SDK's session manager task group is never
initialized and every request raises RuntimeError.
"""
with patch("app.mcp.server.settings") as fake_settings:
fake_settings.mcp_allowed_hosts = allowed_hosts
fake_settings.cors_allow_origins = "http://localhost:5173"
fake_settings.auth_enabled = False
with TestClient(build_mcp_asgi_app()) as client:
yield client
def test_remote_host_allowed_when_configured():
"""A configured non-loopback Host must reach the MCP handler, not 421."""
with _mcp_client("6.86.80.9:*,127.0.0.1:*") as client:
response = client.post(
"/", json=_INITIALIZE, headers={**_HEADERS, "Host": "6.86.80.9:8000"}
)
assert response.status_code == 200
assert "Invalid Host header" not in response.text
def test_unconfigured_host_still_rejected():
"""Protection must stay on: a Host outside the allow-list is refused with 421."""
with _mcp_client("6.86.80.9:*") as client:
response = client.post(
"/", json=_INITIALIZE, headers={**_HEADERS, "Host": "evil.example.com"}
)
assert response.status_code == 421
def test_initialize_response_is_event_stream():
"""Sanity check that a permitted request really completes the MCP handshake."""
with _mcp_client("6.86.80.9:*") as client:
response = client.post(
"/", json=_INITIALIZE, headers={**_HEADERS, "Host": "6.86.80.9:8000"}
)
assert response.status_code == 200
# The Streamable HTTP transport replies as SSE; the JSON-RPC result is
# embedded in a "data:" line rather than being the whole body.
payload = json.loads(response.text.split("data:", 1)[1].strip())
assert payload["result"]["serverInfo"]["name"] == "ai-regulations"
def test_wildcard_disables_protection_explicitly():
"""'*' is the documented opt-out; it must disable the check, not allow-list '*'."""
with patch("app.mcp.server.settings") as fake_settings:
fake_settings.mcp_allowed_hosts = "*"
fake_settings.cors_allow_origins = "http://localhost:5173"
security = _build_transport_security()
assert security.enable_dns_rebinding_protection is False
def test_allow_list_is_parsed_into_transport_settings():
"""Comma-separated config must become the SDK's allowed_hosts list verbatim."""
with patch("app.mcp.server.settings") as fake_settings:
fake_settings.mcp_allowed_hosts = "6.86.80.9:*, localhost:* ,"
fake_settings.cors_allow_origins = "http://localhost:5173"
security = _build_transport_security()
assert security.enable_dns_rebinding_protection is True
assert security.allowed_hosts == ["6.86.80.9:*", "localhost:*"]
assert security.allowed_origins == ["http://localhost:5173"]