空串逗号拆分得空列表,TransportSecuritySettings(allowed_hosts=[]) 启用 DNS 重绑定防护后连 localhost 族一并 421。main() 解析层抽出 _parse_allowed_hosts,空列表以 or None 归一,保持 mcp 默认;含回归测试。 Co-Authored-By: Claude <noreply@anthropic.com>
130 lines
4.8 KiB
Python
130 lines
4.8 KiB
Python
"""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
|
||
|
||
def test_empty_string_allowed_hosts_to_none(self):
|
||
"""回归:`--allowed-hosts ""` 空串自锁——逗号拆分结果为空列表时须按
|
||
None 处理(保持 mcp 默认 localhost 族),否则空白名单启用防护后
|
||
连 localhost 请求也被 421。"""
|
||
from ppclock.mcp_server import _parse_allowed_hosts
|
||
a = parse_args(["--allowed-hosts", ""])
|
||
assert _parse_allowed_hosts(a.allowed_hosts) is None
|