diff --git a/src/ppclock/mcp_server.py b/src/ppclock/mcp_server.py index b812a47..52f6493 100644 --- a/src/ppclock/mcp_server.py +++ b/src/ppclock/mcp_server.py @@ -17,8 +17,19 @@ from .mcp_tools import register_tools TOKEN_ENV = "PPCLOCK_MCP_TOKEN" -def build_server(manager: DeviceManager) -> FastMCP: - mcp = FastMCP("ppclock", stateless_http=True) +def build_server(manager: DeviceManager, + allowed_hosts: list[str] | None = None) -> FastMCP: + """装配 FastMCP。allowed_hosts 为 HTTP Host 白名单(mcp DNS 重绑定防护): + HTTP 模式绑非回环地址时,必须把本机的局域网地址加进来 + (如 ["192.168.61.35:8972", "localhost:8972"]),否则 LAN Host 请求被 421。 + None 保持 mcp 默认(仅 localhost 族)。""" + if allowed_hosts is not None: + from mcp.server.transport_security import TransportSecuritySettings + mcp = FastMCP("ppclock", stateless_http=True, + transport_security=TransportSecuritySettings( + allowed_hosts=allowed_hosts)) + else: + mcp = FastMCP("ppclock", stateless_http=True) register_tools(mcp, manager) return mcp @@ -34,6 +45,10 @@ def parse_args(argv=None) -> argparse.Namespace: help="守候连接秒数(覆盖设备 30-60s 广播窗口)") p.add_argument("--idle-timeout", type=float, default=300.0, help="空闲断开秒数;0=每次调用独立连接") + p.add_argument("--allowed-hosts", default=None, + help="HTTP Host 白名单(逗号分隔)。HTTP 模式绑非回环地址时," + "把本机的局域网地址加进来,如 " + "192.168.61.35:8972,localhost:8972;缺省仅 localhost 族") return p.parse_args(argv) @@ -71,7 +86,9 @@ def main(argv=None) -> None: manager = DeviceManager(mac=args.mac, connect_timeout=args.connect_timeout, idle_timeout=args.idle_timeout) - mcp = build_server(manager) + allowed_hosts = (None if args.allowed_hosts is None else + [h.strip() for h in args.allowed_hosts.split(",") if h.strip()]) + mcp = build_server(manager, allowed_hosts=allowed_hosts) if args.transport == "stdio": try: mcp.run(transport="stdio") diff --git a/tests/test_mcp_server.py b/tests/test_mcp_server.py index f78095c..c672ebc 100644 --- a/tests/test_mcp_server.py +++ b/tests/test_mcp_server.py @@ -1,4 +1,4 @@ -"""mcp_server 测试:参数解析 + 非回环强制 token + ASGI 鉴权中间件。""" +"""mcp_server 测试:参数解析 + 非回环强制 token + ASGI 鉴权中间件 + HTTP Host 白名单。""" import asyncio import pytest @@ -25,6 +25,13 @@ class TestParseArgs: assert a.port == 9000 and a.mac == "AA:BB:CC:DD:EE:FF" assert a.connect_timeout == 120.0 and a.idle_timeout == 60.0 + def test_allowed_hosts_default_none(self): + assert parse_args([]).allowed_hosts is None + + def test_allowed_hosts_flag(self): + a = parse_args(["--allowed-hosts", "192.168.61.35:8972,localhost:8972"]) + assert a.allowed_hosts == "192.168.61.35:8972,localhost:8972" + def run_middleware(token_set, auth_header): """驱动 TokenAuthMiddleware,返回 (status, app_called)。""" @@ -70,3 +77,45 @@ class TestBindGuard: monkeypatch.delenv(TOKEN_ENV, raising=False) with pytest.raises(SystemExit): main(["--transport", "http", "--host", "0.0.0.0"]) + + +class TestAllowedHosts: + """HTTP 模式 Host 白名单回归(Windows 实测:LAN Host 被 421 拒)。 + + 用 starlette TestClient 驱动真实 ASGI app(含 lifespan,stateless 会话管理器 + 需要其 task group);manager 只构造不连接,initialize 不触达设备。 + """ + + LAN_HOST = "192.168.61.35:8972" + + def _make_app(self): + from ppclock.device_manager import DeviceManager + from ppclock.mcp_server import build_server + manager = DeviceManager(mac="AA:BB:CC:DD:EE:FF") + mcp = build_server(manager, allowed_hosts=[self.LAN_HOST, + "localhost:8972"]) + return mcp.streamable_http_app() + + def _initialize(self, app, host): + from starlette.testclient import TestClient + with TestClient(app) as client: + return client.post( + "/mcp", + headers={ + "Host": host, + "Accept": "application/json, text/event-stream", + "Content-Type": "application/json", + }, + json={"jsonrpc": "2.0", "id": 1, "method": "initialize", + "params": {"protocolVersion": "2025-03-26", + "capabilities": {}, + "clientInfo": {"name": "test", "version": "0"}}}, + ) + + def test_lan_host_not_421(self): + r = self._initialize(self._make_app(), self.LAN_HOST) + assert r.status_code != 421 + + def test_evil_host_421(self): + r = self._initialize(self._make_app(), "evil.example.com") + assert r.status_code == 421