diff --git a/src/ppclock/mcp_tools.py b/src/ppclock/mcp_tools.py new file mode 100644 index 0000000..d057737 --- /dev/null +++ b/src/ppclock/mcp_tools.py @@ -0,0 +1,171 @@ +"""MCP 工具层:13 个高层工具,统一输出契约(与 CLI --json 同构)。 + +成功 {"ok": true, "tool": name, "data": {...}} +失败 {"ok": false, "tool": name, "error": {"code", "message", "hint"?}} +""" +from __future__ import annotations + +import base64 +import binascii +import datetime +import io +import os + +from PIL import Image + +from .client import PPClient +from .transports.base import TransportError + +_TIMEOUT_HINT = "设备长睡眠、广播窗口 30-60s;调大 --connect-timeout 或稍后重试" +_NOTFOUND_HINT = "未发现 NRF- 前缀设备;确认设备在位且未处于深睡" + + +def decode_image_source(source: str) -> Image.Image: + """source = 服务器本地路径,或 base64(可带 data URI 前缀)。""" + if os.path.exists(source): + return Image.open(source) + data = source + if data.startswith("data:"): + if "," not in data: + raise ValueError("data URI 缺少 ',' 分隔") + data = data.split(",", 1)[1] + try: + raw = base64.b64decode(data, validate=True) + except (binascii.Error, ValueError): + raise ValueError("source 既不是存在的文件路径,也不是合法 base64") from None + return Image.open(io.BytesIO(raw)) + + +def _err(tool: str, code: str, message: str, hint: str | None = None) -> dict: + error = {"code": code, "message": message} + if hint: + error["hint"] = hint + return {"ok": False, "tool": tool, "error": error} + + +def _map_transport_error(tool: str, exc: TransportError) -> dict: + msg = str(exc) + low = msg.lower() + if "timeout" in low or "timed out" in low: + return _err(tool, "CONNECT_TIMEOUT", msg, _TIMEOUT_HINT) + if "not found" in low or "no device" in low: + return _err(tool, "DEVICE_NOT_FOUND", msg, _NOTFOUND_HINT) + return _err(tool, "BLE_ERROR", msg) + + +def register_tools(mcp, manager) -> None: + """把 13 个工具注册到 FastMCP 实例;manager 为 DeviceManager(测试可注入 fake)。""" + + async def _call(tool: str, op) -> dict: + try: + data = await manager.run(op) + return {"ok": True, "tool": tool, "data": data if data is not None else {}} + except TransportError as exc: + return _map_transport_error(tool, exc) + except ValueError as exc: + return _err(tool, "INVALID_PARAM", str(exc)) + except Exception as exc: # noqa: BLE001 - 工具层兜底,错误须回 agent 而非炸会话 + return _err(tool, "INTERNAL", f"{type(exc).__name__}: {exc}") + + @mcp.tool() + async def scan_devices(timeout: float = 5.0) -> dict: + """扫描附近 NRF- 前缀墨水屏设备,返回 [{name,address,rssi}]。""" + try: + devices = await PPClient.scan_devices("local", timeout=timeout) + return {"ok": True, "tool": "scan_devices", "data": {"devices": devices}} + except TransportError as exc: + return _map_transport_error("scan_devices", exc) + except Exception as exc: # noqa: BLE001 + return _err("scan_devices", "INTERNAL", f"{type(exc).__name__}: {exc}") + + @mcp.tool() + async def device_status() -> dict: + """查询连接状态/绑定 MAC/设备 ID/空闲秒数。""" + try: + return {"ok": True, "tool": "device_status", "data": await manager.status()} + except Exception as exc: # noqa: BLE001 + return _err("device_status", "INTERNAL", f"{type(exc).__name__}: {exc}") + + @mcp.tool() + async def set_time(tz: float = 8.0) -> dict: + """对时;tz 为时区偏移(默认 8=北京时间)。""" + return await _call("set_time", lambda dev: dev.set_time(tz=tz)) + + @mcp.tool() + async def set_mode(mode: str) -> dict: + """切换显示模式:clock1-3/calendar1-3/image0-3/tricolor/mono。""" + return await _call("set_mode", lambda dev: dev.set_mode(mode)) + + @mcp.tool() + async def toggle(name: str) -> dict: + """单字节切换:invert/font/rotate180/hour_format/clock_color。""" + return await _call("toggle", lambda dev: dev.toggle(name)) + + @mcp.tool() + async def upload_image(source: str, slot: int = 0, algo: str = "atkinson", + mono: bool = False, + threshold: float | None = None, + diffusion: float | None = None, + brightness: float | None = None, + contrast: float | None = None, + saturation: float | None = None, + rotate: float | None = None) -> dict: + """传图到指定槽位。source=服务器路径或 base64;algo∈ + none/floydsteinberg/atkinson/bayer/stucki/jarvis。""" + try: + adjust = {k: v for k, v in { + "threshold": threshold, "diffusion": diffusion, + "brightness": brightness, "contrast": contrast, + "saturation": saturation, "rotate": rotate}.items() if v is not None} + img = decode_image_source(source) + except ValueError as exc: + return _err("upload_image", "INVALID_PARAM", str(exc)) + return await _call("upload_image", lambda dev: dev.upload_image( + img, slot=slot, algo=algo, mono=mono, **adjust)) + + @mcp.tool() + async def render_template(name: str, payload: dict | None = None, + slot: int = 0, algo: str = "atkinson") -> dict: + """本地渲染模板并上传:schedule/businesscard/memo/course/qrcode/custom。""" + return await _call("render_template", lambda dev: dev.render_template( + name, payload, slot=slot, algo=algo)) + + @mcp.tool() + async def countdown(date: str, mode: str = "clock", + prefix: str | None = None) -> dict: + """倒计时;date 为 ISO YYYY-MM-DD;mode=clock/calendar。""" + try: + d = datetime.date.fromisoformat(date) + except ValueError as exc: + return _err("countdown", "INVALID_PARAM", str(exc)) + return await _call("countdown", lambda dev: dev.countdown( + d, mode=mode, prefix=prefix)) + + @mcp.tool() + async def countdown_off() -> dict: + """关闭倒计时。""" + return await _call("countdown_off", lambda dev: dev.countdown_off()) + + @mcp.tool() + async def calendar_text(text: str) -> dict: + """设置日历中文文字并切到日历模式一。""" + return await _call("calendar_text", lambda dev: dev.calendar_text(text)) + + @mcp.tool() + async def set_sleep(on: bool, start_h: int, end_h: int) -> dict: + """设置休眠时段(0-23 时)。""" + return await _call("set_sleep", lambda dev: dev.sleep(on, start_h, end_h)) + + @mcp.tool() + async def clear_screen() -> dict: + """刷屏(EPD 清屏)。""" + return await _call("clear_screen", lambda dev: dev.clear()) + + @mcp.tool() + async def raw_send(data_hex: str, channel: str = "alt") -> dict: + """逃生舱:直发十六进制字节串;channel=rxtx/epd/alt。""" + try: + data = bytes.fromhex(data_hex) + except ValueError as exc: + return _err("raw_send", "INVALID_PARAM", str(exc)) + return await _call("raw_send", lambda dev: dev.raw(data, channel)) diff --git a/tests/test_mcp_tools.py b/tests/test_mcp_tools.py new file mode 100644 index 0000000..d852e8e --- /dev/null +++ b/tests/test_mcp_tools.py @@ -0,0 +1,174 @@ +"""mcp_tools 测试:图源解码 + 工具契约(内存 MCP 会话端到端)。""" +import base64 +import io +import json + +import pytest +from mcp.server.fastmcp import FastMCP +from mcp.shared.memory import create_connected_server_and_client_session +from PIL import Image + +from ppclock.client import PPClient +from ppclock.mcp_tools import decode_image_source, register_tools +from ppclock.transports.base import TransportError + + +class FakeTransport: + def __init__(self): + self.address = "AA:BB:CC:DD:EE:FF" + self.epd_writes = [] + self.rxtx_writes = [] + + async def __aenter__(self): + return self + + async def __aexit__(self, *exc): + return False + + async def write_epd(self, data, response=True): + self.epd_writes.append(data) + + async def write_rxtx(self, data, response=True): + self.rxtx_writes.append(data) + + async def read_rxtx(self): + return b"\x00" + + async def request_device_id(self, timeout=8.0): + return "81233F3C267112" + + async def run_ota(self, image, on_progress=None): + return {} + + async def delay(self, seconds): + pass + + +class FakeManager: + """模拟 DeviceManager:直接对 fake client 执行 op,或抛预置错误。""" + + def __init__(self, error=None): + self.client = PPClient(FakeTransport()) + self.error = error + self.ops = 0 + + async def run(self, op): + self.ops += 1 + if self.error is not None: + raise self.error + return await op(self.client) + + async def status(self): + return {"connected": True, "mac": "AA:BB:CC:DD:EE:FF", + "idle_seconds": 0.0, "connect_count": 1, + "connect_timeout": 90.0, "idle_timeout": 300.0, + "device_id": "81233F3C267112"} + + +def make_session_coro(manager): + mcp = FastMCP("ppclock-test") + register_tools(mcp, manager) + return mcp + + +async def call(mcp, tool, args=None): + async with create_connected_server_and_client_session(mcp._mcp_server) as s: + result = await s.call_tool(tool, args or {}) + assert not result.isError, result.content + return json.loads(result.content[0].text) + + +def png_b64() -> str: + buf = io.BytesIO() + Image.new("RGB", (400, 300), "white").save(buf, "PNG") + return base64.b64encode(buf.getvalue()).decode() + + +class TestDecodeImageSource: + def test_path(self, tmp_path): + p = tmp_path / "a.png" + Image.new("RGB", (10, 10)).save(p) + assert decode_image_source(str(p)).size == (10, 10) + + def test_base64(self): + assert decode_image_source(png_b64()).size == (400, 300) + + def test_data_uri(self): + src = "data:image/png;base64," + png_b64() + assert decode_image_source(src).size == (400, 300) + + def test_garbage(self): + with pytest.raises(ValueError): + decode_image_source("!!!not-base64!!!") + + +class TestTools: + @pytest.mark.asyncio + async def test_set_time_ok(self): + r = await call(make_session_coro(FakeManager()), "set_time", {"tz": 8.0}) + assert r["ok"] is True and r["tool"] == "set_time" + + @pytest.mark.asyncio + async def test_upload_image_base64(self): + r = await call(make_session_coro(FakeManager()), "upload_image", + {"source": png_b64(), "slot": 0}) + assert r["ok"] is True + assert r["data"]["bytes_bw"] == 15000 + assert r["data"]["slot"] == 0 + + @pytest.mark.asyncio + async def test_set_mode_invalid_param(self): + r = await call(make_session_coro(FakeManager()), "set_mode", + {"mode": "nonsense"}) + assert r["ok"] is False + assert r["error"]["code"] == "INVALID_PARAM" + + @pytest.mark.asyncio + async def test_transport_error_mapped(self): + m = FakeManager(error=TransportError("Connection timed out")) + r = await call(make_session_coro(m), "set_time", {}) + assert r["ok"] is False + assert r["error"]["code"] == "CONNECT_TIMEOUT" + assert "hint" in r["error"] + + @pytest.mark.asyncio + async def test_device_status(self): + r = await call(make_session_coro(FakeManager()), "device_status", {}) + assert r["ok"] is True + assert r["data"]["device_id"] == "81233F3C267112" + + @pytest.mark.asyncio + async def test_raw_send(self): + r = await call(make_session_coro(FakeManager()), "raw_send", + {"data_hex": "e201", "channel": "alt"}) + assert r["ok"] is True + + @pytest.mark.asyncio + async def test_raw_send_bad_hex(self): + r = await call(make_session_coro(FakeManager()), "raw_send", + {"data_hex": "zz", "channel": "alt"}) + assert r["ok"] is False + assert r["error"]["code"] == "INVALID_PARAM" + + @pytest.mark.asyncio + async def test_render_template(self): + r = await call(make_session_coro(FakeManager()), "render_template", + {"name": "custom", "payload": {"text": "你好"}, "slot": 0}) + assert r["ok"] is True + assert r["data"]["template"] == "custom" + + @pytest.mark.asyncio + async def test_countdown(self): + r = await call(make_session_coro(FakeManager()), "countdown", + {"date": "2026-12-31", "mode": "clock", "prefix": "目标"}) + assert r["ok"] is True + + @pytest.mark.asyncio + async def test_scan_devices(self, monkeypatch): + async def fake_scan(uri="local", timeout=5.0): + return [{"name": "NRF-5DBF28", "address": "18:BC:5A:5D:BF:28", + "rssi": -60}] + monkeypatch.setattr(PPClient, "scan_devices", staticmethod(fake_scan)) + r = await call(make_session_coro(FakeManager()), "scan_devices", {}) + assert r["ok"] is True + assert r["data"]["devices"][0]["address"] == "18:BC:5A:5D:BF:28"