From 9506490ce2e877519405b1e810383644a3b3fc55 Mon Sep 17 00:00:00 2001 From: chenwei Date: Thu, 30 Jul 2026 16:39:13 +0000 Subject: [PATCH] =?UTF-8?q?feat(mcp):=20ppclock-mcp=20=E5=85=A5=E5=8F=A3?= =?UTF-8?q?=E2=80=94=E2=80=94stdio/HTTP=20=E5=8F=8C=E6=A8=A1=20+=20token?= =?UTF-8?q?=20=E9=89=B4=E6=9D=83=20+=20=E9=9D=9E=E5=9B=9E=E7=8E=AF?= =?UTF-8?q?=E5=AE=88=E5=8D=AB?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- pyproject.toml | 2 + src/ppclock/mcp_server.py | 88 +++++++++++++++++++++++++++++++++++++++ tests/test_mcp_server.py | 72 ++++++++++++++++++++++++++++++++ 3 files changed, 162 insertions(+) create mode 100644 src/ppclock/mcp_server.py create mode 100644 tests/test_mcp_server.py diff --git a/pyproject.toml b/pyproject.toml index 7e675a3..1de320a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -14,10 +14,12 @@ dependencies = [ [project.optional-dependencies] dev = ["pytest>=8"] +mcp = ["mcp>=1.10,<2"] [project.scripts] ppclock = "ppclock.cli:main" ppclock-bridge = "ppclock.bridge_server:run" +ppclock-mcp = "ppclock.mcp_server:main" [tool.setuptools.packages.find] where = ["src"] diff --git a/src/ppclock/mcp_server.py b/src/ppclock/mcp_server.py new file mode 100644 index 0000000..b812a47 --- /dev/null +++ b/src/ppclock/mcp_server.py @@ -0,0 +1,88 @@ +"""ppclock-mcp —— MCP 服务器入口(stdio / streamable HTTP 双模)。 + +stdio:本机 agent 直挂(Claude Code/Desktop 等)。 +http :局域网 agent 平台远程调用;绑非回环地址必须设 PPCLOCK_MCP_TOKEN。 +""" +from __future__ import annotations + +import argparse +import asyncio +import os + +from mcp.server.fastmcp import FastMCP + +from .device_manager import DeviceManager +from .mcp_tools import register_tools + +TOKEN_ENV = "PPCLOCK_MCP_TOKEN" + + +def build_server(manager: DeviceManager) -> FastMCP: + mcp = FastMCP("ppclock", stateless_http=True) + register_tools(mcp, manager) + return mcp + + +def parse_args(argv=None) -> argparse.Namespace: + p = argparse.ArgumentParser(prog="ppclock-mcp", + description="ppclock 墨水屏 MCP 服务器") + p.add_argument("--transport", choices=["stdio", "http"], default="stdio") + p.add_argument("--host", default="127.0.0.1") + p.add_argument("--port", type=int, default=8972) + p.add_argument("--mac", default=None, help="设备 MAC;缺省自动扫描 NRF- 前缀") + p.add_argument("--connect-timeout", type=float, default=90.0, + help="守候连接秒数(覆盖设备 30-60s 广播窗口)") + p.add_argument("--idle-timeout", type=float, default=300.0, + help="空闲断开秒数;0=每次调用独立连接") + return p.parse_args(argv) + + +class TokenAuthMiddleware: + """纯 ASGI 中间件:校验 Authorization: Bearer ,不符 401。""" + + def __init__(self, app, token: str): + self.app = app + self.token = token + + async def __call__(self, scope, receive, send): + if scope["type"] == "http": + headers = dict(scope.get("headers") or []) + auth = headers.get(b"authorization", b"").decode("latin1") + if auth != f"Bearer {self.token}": + await send({"type": "http.response.start", "status": 401, + "headers": [(b"content-type", b"text/plain")]}) + await send({"type": "http.response.body", + "body": b"unauthorized"}) + return + await self.app(scope, receive, send) + + +async def _run_http(mcp: FastMCP, host: str, port: int, token: str | None): + import uvicorn + app = mcp.streamable_http_app() + if token: + app = TokenAuthMiddleware(app, token) + config = uvicorn.Config(app, host=host, port=port, log_level="info") + await uvicorn.Server(config).serve() + + +def main(argv=None) -> None: + args = parse_args(argv) + manager = DeviceManager(mac=args.mac, + connect_timeout=args.connect_timeout, + idle_timeout=args.idle_timeout) + mcp = build_server(manager) + if args.transport == "stdio": + try: + mcp.run(transport="stdio") + finally: + asyncio.run(manager.close()) + return + token = os.environ.get(TOKEN_ENV) + loopback = args.host in ("127.0.0.1", "localhost", "::1") + if not loopback and not token: + raise SystemExit(f"绑定非回环地址 {args.host} 必须设置 {TOKEN_ENV}") + try: + asyncio.run(_run_http(mcp, args.host, args.port, token)) + finally: + asyncio.run(manager.close()) diff --git a/tests/test_mcp_server.py b/tests/test_mcp_server.py new file mode 100644 index 0000000..f78095c --- /dev/null +++ b/tests/test_mcp_server.py @@ -0,0 +1,72 @@ +"""mcp_server 测试:参数解析 + 非回环强制 token + ASGI 鉴权中间件。""" +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 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"])