feat(mcp): ppclock-mcp 入口——stdio/HTTP 双模 + token 鉴权 + 非回环守卫
This commit is contained in:
1 parent
bb3dbcb06a
commit
9506490ce2
3 files changed
+162
No files matched your search
@@ -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"]
|
||||
|
||||
@@ -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 <token>,不符 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())
|
||||
@@ -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"])
|
||||
Reference in new issue
Block a user