fix: harden MCP endpoint after code review

Critical: the MCP SDK auto-enables DNS-rebinding protection when its host
parameter is left at the 127.0.0.1 default, hard-coding a loopback-only Host
allow-list. Every remote client (the only deployment this feature targets) was
refused with HTTP 421 before auth or the tool ran. Now driven by a new
MCP_ALLOWED_HOSTS setting, with '*' as an explicit, logged opt-out.

Also bounds query/top_k to match AskRequest (top_k is amplified 4x downstream,
so an unbounded value was a resource-exhaustion vector), decodes the
Authorization header as latin-1 per the ASGI spec instead of raising a 500 on
malformed bytes, and returns WWW-Authenticate on 401 per RFC 7235.

Moves the psycopg2 import guard into backend/tests/conftest.py: duplicated
across four test modules, it only worked because of alphabetical collection
order, and any earlier-sorting package would have reintroduced a live
connection attempt against the production database.

Registers the mcp module in the authoritative backend architecture doc.

84 backend tests pass. Verified against a live server: allowed remote Host
returns a valid initialize result, unknown Host returns 421, missing token
returns 401 with WWW-Authenticate.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
wangwei
2026-07-29 17:11:54 +08:00
co-authored by Copilot
parent bd3dc38d1d
commit 49ee50c104
13 changed files with 330 additions and 51 deletions
+14
View File
@@ -198,6 +198,20 @@ class Settings(BaseSettings):
description="Comma-separated allowed CORS origins. Never use * in production.",
)
# ── MCP ───────────────────────────────────────────────────────────────────
# The MCP SDK enables DNS-rebinding protection whenever the transport is
# bound to a loopback host, which rejects any Host header not in this list
# with HTTP 421. Deployments reachable by a real hostname/IP must list it
# here or every remote MCP client is refused before the handler runs.
mcp_allowed_hosts: str = Field(
default="127.0.0.1:*,localhost:*,[::1]:*",
description=(
"Comma-separated Host header values accepted by the MCP endpoint. "
"A ':*' suffix matches any port. Set to '*' to disable DNS-rebinding "
"protection entirely (not recommended)."
),
)
@lru_cache
def get_settings() -> Settings:
"""Return settings."""
+59 -5
View File
@@ -8,13 +8,20 @@ dict. No new retrieval, ranking, or LLM orchestration logic lives here.
from __future__ import annotations
import logging
from typing import Annotated
from mcp.server import MCPServer
from mcp.server.transport_security import TransportSecuritySettings
from pydantic import Field
from starlette.responses import PlainTextResponse
from starlette.types import ASGIApp, Receive, Scope, Send
from app.config.settings import settings
from app.shared.bootstrap import get_agent_conversation_service, get_jwt_handler
logger = logging.getLogger(__name__)
# Single shared MCPServer instance — analogous to the single shared FastAPI
# `app` instance in app/api/main.py. Tools registered via @mcp.tool() below.
# Note: the installed mcp SDK (2.0.0) renamed the older "FastMCP" class to
@@ -24,12 +31,22 @@ mcp = MCPServer("ai-regulations")
@mcp.tool()
def search_regulations(query: str, top_k: int = 5) -> dict:
def search_regulations(
query: Annotated[str, Field(min_length=1, max_length=2000)],
top_k: Annotated[int, Field(ge=1, le=20)] = 5,
) -> dict:
"""Search the compliance knowledge base and return a grounded answer.
query: Natural-language search question, e.g. "国六排放标准最新要求".
top_k: Maximum number of cited sources to return (default 5).
top_k: Maximum number of cited sources to return (1-20, default 5).
"""
# Bounds mirror AskRequest in app/api/models/agent.py so the MCP path cannot
# be used to bypass the REST endpoint's limits. They matter more here than
# there: KnowledgeRetrievalService amplifies top_k (candidate_k = top_k * 4)
# when reranking, and an LLM client can easily hallucinate a huge value.
# Declaring them via Annotated puts them in the advertised JSON schema too,
# so well-behaved clients never send an out-of-range value in the first place.
#
# No session_id is passed: this keeps each call stateless (no
# ConversationStore reads/writes), matching "search" semantics rather
# than multi-turn chat semantics.
@@ -61,19 +78,53 @@ class MCPAuthMiddleware:
return
headers = dict(scope["headers"])
auth_header = headers.get(b"authorization", b"").decode()
# ASGI header values are raw bytes specified as latin-1, not UTF-8;
# decoding strictly as UTF-8 would raise on a malformed byte and turn a
# bad request into an unhandled 500.
auth_header = headers.get(b"authorization", b"").decode("latin-1")
token = auth_header.removeprefix("Bearer ").strip()
try:
get_jwt_handler().decode_token(token)
except ValueError as exc:
# Reject before the MCP session/protocol layer ever sees the request.
response = PlainTextResponse(str(exc), status_code=401)
# WWW-Authenticate matches the get_current_user dependency (auth.py)
# and is required by RFC 7235 so clients can tell "needs credentials"
# apart from a generic failure.
response = PlainTextResponse(
str(exc), status_code=401, headers={"WWW-Authenticate": "Bearer"}
)
await response(scope, receive, send)
return
await self.app(scope, receive, send)
def _build_transport_security() -> TransportSecuritySettings:
"""Translate the configured MCP host allow-list into SDK transport settings.
Without this the SDK infers its own allow-list from the bind host, which
defaults to 127.0.0.1 and therefore rejects every remote client with HTTP
421 — fatal for a remotely deployed backend.
"""
allowed = [h.strip() for h in settings.mcp_allowed_hosts.split(",") if h.strip()]
if "*" in allowed:
# Explicit, logged opt-out. Kept as an escape hatch for environments
# behind a proxy that rewrites Host unpredictably, but never the default.
logger.warning(
"MCP DNS-rebinding protection is disabled (mcp_allowed_hosts='*'). "
"Set MCP_ALLOWED_HOSTS to the real deployment host(s) instead."
)
return TransportSecuritySettings(enable_dns_rebinding_protection=False)
return TransportSecuritySettings(
enable_dns_rebinding_protection=True,
allowed_hosts=allowed,
# Browser clients send Origin; reuse the already-maintained CORS list so
# there is one place to declare trusted web origins. Non-browser MCP
# clients send no Origin at all, which the SDK treats as allowed.
allowed_origins=[o.strip() for o in settings.cors_allow_origins.split(",") if o.strip()],
)
def build_mcp_asgi_app() -> ASGIApp:
"""Return the Streamable HTTP ASGI app for the MCP server, auth-guarded.
@@ -83,6 +134,9 @@ def build_mcp_asgi_app() -> ASGIApp:
path to "/", the effective external path would be the confusing "/mcp/mcp"
instead of "/mcp".
"""
asgi_app = mcp.streamable_http_app(streamable_http_path="/")
asgi_app = mcp.streamable_http_app(
streamable_http_path="/",
transport_security=_build_transport_security(),
)
asgi_app.add_middleware(MCPAuthMiddleware)
return asgi_app
+27
View File
@@ -0,0 +1,27 @@
"""Shared pytest fixtures and import-time guards for the backend test suite.
pytest imports this file before any test module beneath backend/tests/, which
makes it the only reliable place to install import-time guards: individual test
modules cannot guarantee they run first, because collection order follows
directory names.
"""
from __future__ import annotations
import sys
from unittest.mock import MagicMock
# app/shared/bootstrap.py (the composition root) eagerly imports the Postgres
# store modules, which do `import psycopg2` at their own module scope and later
# open a real connection pool. Any test that transitively imports bootstrap
# would therefore bind the real driver and attempt a live TCP connection to the
# configured production database, surfacing as a multi-second timeout rather
# than an obvious error. Binding mocks here — before the first test module is
# imported — makes that impossible regardless of collection order.
# setdefault (not assignment) keeps a real psycopg2 in place if something has
# already imported it deliberately.
_mock_psycopg2 = MagicMock()
_mock_psycopg2.extras = MagicMock()
sys.modules.setdefault("psycopg2", _mock_psycopg2)
sys.modules.setdefault("psycopg2.extras", _mock_psycopg2.extras)
sys.modules.setdefault("psycopg2.pool", MagicMock())
+25 -17
View File
@@ -7,23 +7,7 @@ independent of FastMCP internals.
from __future__ import annotations
import sys
from unittest.mock import MagicMock, patch
# app.mcp.server imports app.shared.bootstrap at module scope (needed for
# get_agent_conversation_service/get_jwt_handler), which in turn eagerly
# imports several Postgres store modules that do `import psycopg2` at their
# own module scope. Since this is the only test file in backend/tests/mcp/
# that imports app.mcp.server at module scope (not inside a test function),
# it is the first thing to trigger that chain during pytest collection —
# guard psycopg2 here the same way backend/tests/observability/
# test_model_usage_bootstrap.py and test_model_usage_persistence.py already
# do, so a real (network-connecting) psycopg2 never gets bound first.
mock_psycopg2 = MagicMock()
mock_psycopg2.extras = MagicMock()
sys.modules.setdefault("psycopg2", mock_psycopg2)
sys.modules.setdefault("psycopg2.extras", mock_psycopg2.extras)
sys.modules.setdefault("psycopg2.pool", MagicMock())
from unittest.mock import patch
from starlette.applications import Starlette
from starlette.responses import PlainTextResponse
@@ -83,3 +67,27 @@ def test_auth_disabled_always_passes_through():
client = TestClient(_dummy_app())
response = client.get("/ping")
assert response.status_code == 200
def test_401_includes_www_authenticate_header():
"""RFC 7235 requires WWW-Authenticate on 401 so clients can tell why they failed."""
with patch("app.mcp.server.settings") as fake_settings:
fake_settings.auth_enabled = True
client = TestClient(_dummy_app())
response = client.get("/ping")
assert response.status_code == 401
assert response.headers["WWW-Authenticate"] == "Bearer"
def test_non_utf8_authorization_header_is_rejected_not_crashed():
"""A non-UTF-8 header byte must yield a clean 401, not an unhandled 500.
ASGI header values are latin-1 bytes, so any remote client could otherwise
trigger a UnicodeDecodeError inside the middleware at will.
"""
with patch("app.mcp.server.settings") as fake_settings:
fake_settings.auth_enabled = True
client = TestClient(_dummy_app(), raise_server_exceptions=False)
# Bypass the http client's own header encoding by writing raw bytes.
response = client.get("/ping", headers={"Authorization": b"Bearer \xff\xfe"})
assert response.status_code == 401
@@ -0,0 +1,107 @@
"""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"]
@@ -7,6 +7,7 @@ correct dict shape out.
from __future__ import annotations
import asyncio
from dataclasses import dataclass
from unittest.mock import MagicMock, patch
@@ -15,6 +16,9 @@ from unittest.mock import MagicMock, patch
class _FakeSource:
"""Minimal stand-in for a real Source dataclass (only __dict__ is used)."""
# A dataclass, not a MagicMock: the adapter serializes sources via
# source.__dict__, and a MagicMock's __dict__ is full of internal mock
# attributes, which would make the assertions meaningless.
doc_id: str
doc_title: str
score: float
@@ -75,3 +79,20 @@ def test_search_regulations_default_top_k():
search_regulations(query="q")
assert fake_service.ask.call_args.kwargs["top_k"] == 5
def test_advertised_schema_bounds_top_k_and_query():
"""The advertised JSON schema must carry the same bounds as AskRequest.
Bounds declared via Annotated are what the SDK validates against and what
clients see, so asserting on the generated schema is the only way to catch
a regression that silently drops them.
"""
from app.mcp.server import mcp
schema = asyncio.run(mcp.list_tools())[0].input_schema["properties"]
assert schema["top_k"]["minimum"] == 1
assert schema["top_k"]["maximum"] == 20
assert schema["query"]["minLength"] == 1
assert schema["query"]["maxLength"] == 2000
@@ -12,17 +12,10 @@ is needed anywhere in this file — asyncio.create_task itself is also mocked.
from __future__ import annotations
import sys
from unittest.mock import MagicMock, patch
# Patch psycopg2 before importing anything that transitively imports it, in
# case this file is collected before test_model_usage_persistence.py.
mock_psycopg2 = MagicMock()
mock_psycopg2.extras = MagicMock()
sys.modules.setdefault("psycopg2", mock_psycopg2)
sys.modules.setdefault("psycopg2.extras", mock_psycopg2.extras)
sys.modules.setdefault("psycopg2.pool", MagicMock())
# psycopg2 is mocked centrally in backend/tests/conftest.py, which pytest
# imports before any test module regardless of collection order.
from app.shared import bootstrap
from app.shared.model_usage_tracker import ModelUsageEntry, ModelUsageTracker
@@ -6,17 +6,11 @@ Mirrors the mocking pattern in backend/tests/perception/test_postgres_event_stor
from __future__ import annotations
import sys
from datetime import datetime, timezone
from unittest.mock import MagicMock, patch
# Patch psycopg2 before importing the module under test.
mock_psycopg2 = MagicMock()
mock_psycopg2.extras = MagicMock()
sys.modules.setdefault("psycopg2", mock_psycopg2)
sys.modules.setdefault("psycopg2.extras", mock_psycopg2.extras)
sys.modules.setdefault("psycopg2.pool", MagicMock())
# psycopg2 is mocked centrally in backend/tests/conftest.py, so importing the
# module under test here never binds the real driver.
from app.shared.model_usage_tracker import ModelUsageEntry
@@ -4,14 +4,8 @@ import json
from unittest.mock import MagicMock, patch
import pytest
# Patch psycopg2 before importing the module under test
import sys
mock_psycopg2 = MagicMock()
mock_psycopg2.extras = MagicMock()
sys.modules.setdefault("psycopg2", mock_psycopg2)
sys.modules.setdefault("psycopg2.extras", mock_psycopg2.extras)
sys.modules.setdefault("psycopg2.pool", MagicMock())
# psycopg2 is mocked centrally in backend/tests/conftest.py, so importing the
# module under test here never binds the real driver.
from app.infrastructure.perception.base_event_store import BaseEventStore