Files
qianbian/tests/test_mcp_server.py
T
chenweiandClaude 69c4d209c0 fix(mcp): HTTP 模式新增 --allowed-hosts,修 LAN Host 421
Windows 实测:绑 0.0.0.0 后局域网请求被 mcp DNS 重绑定防护 421 拒。
- parse_args 加 --allowed-hosts(逗号分隔 Host 白名单,默认 None 保持
  mcp 仅 localhost 族行为)
- build_server(manager, allowed_hosts=...) 经 TransportSecuritySettings
  传入(mcp 1.29.0 API 与设计一致,无需调整)
- main() 逗号拆分(去空白、丢空段)
- 回归测试:TestClient 驱动真实 ASGI app(含 lifespan),
  Host=192.168.61.35:8972 initialize → 200;Host=evil.example.com → 421

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-30 23:55:57 +00:00

122 lines
4.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""mcp_server 测试:参数解析 + 非回环强制 token + ASGI 鉴权中间件 + HTTP Host 白名单。"""
import asyncio
import pytest
from ppclock.mcp_server import (TOKEN_ENV, TokenAuthMiddleware,
parse_args)
class TestParseArgs:
def test_defaults(self):
a = parse_args([])
assert a.transport == "stdio"
assert a.host == "127.0.0.1"
assert a.port == 8972
assert a.mac is None
assert a.connect_timeout == 90.0
assert a.idle_timeout == 300.0
def test_http_flags(self):
a = parse_args(["--transport", "http", "--host", "0.0.0.0",
"--port", "9000", "--mac", "AA:BB:CC:DD:EE:FF",
"--connect-timeout", "120", "--idle-timeout", "60"])
assert a.transport == "http" and a.host == "0.0.0.0"
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)。"""
captured = {}
async def app(scope, receive, send):
captured["called"] = True
await send({"type": "http.response.start", "status": 200, "headers": []})
await send({"type": "http.response.body", "body": b"ok"})
sent = []
async def send(msg):
sent.append(msg)
headers = []
if auth_header is not None:
headers.append((b"authorization", auth_header.encode()))
scope = {"type": "http", "headers": headers}
mw = TokenAuthMiddleware(app, token_set)
asyncio.run(mw(scope, None, send))
status = next(m["status"] for m in sent if m["type"] == "http.response.start")
return status, captured.get("called", False)
class TestTokenAuth:
def test_no_header_401(self):
status, called = run_middleware("secret", None)
assert status == 401 and called is False
def test_wrong_token_401(self):
status, called = run_middleware("secret", "Bearer nope")
assert status == 401 and called is False
def test_right_token_passes(self):
status, called = run_middleware("secret", "Bearer secret")
assert status == 200 and called is True
class TestBindGuard:
def test_non_loopback_requires_token(self, monkeypatch):
from ppclock.mcp_server import main
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