"""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