From 239b2ad5fad31c37092f8fb33a89c1016a9b5f4f Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:10:21 -0700 Subject: [PATCH 01/16] PybindSwitchProController: add controller_name() controller_name() returns the active controller implementation, e.g. "Nintendo Switch: Pro Controller", or an empty string if not ready. Co-Authored-By: Claude Opus 5.5 --- .../Source/Integrations/PybindSwitchController.cpp | 5 +++++ SerialPrograms/Source/Integrations/PybindSwitchController.h | 6 ++++++ 2 files changed, 11 insertions(+) diff --git a/SerialPrograms/Source/Integrations/PybindSwitchController.cpp b/SerialPrograms/Source/Integrations/PybindSwitchController.cpp index 77aae9d4a1..da3f4dd57e 100644 --- a/SerialPrograms/Source/Integrations/PybindSwitchController.cpp +++ b/SerialPrograms/Source/Integrations/PybindSwitchController.cpp @@ -129,6 +129,11 @@ std::string PybindSwitchProController::current_status() const{ PybindSwitchProControllerInternal* internal = (PybindSwitchProControllerInternal*)m_internals; return internal->m_connection->raw_status_text(); } +std::string PybindSwitchProController::controller_name() const{ + PybindSwitchProControllerInternal* internal = (PybindSwitchProControllerInternal*)m_internals; + ProController* controller = internal->controller_if_ready(); + return controller == nullptr ? "" : controller->name(); +} void PybindSwitchProController::wait_for_all_requests(){ diff --git a/SerialPrograms/Source/Integrations/PybindSwitchController.h b/SerialPrograms/Source/Integrations/PybindSwitchController.h index decbe6a252..a59f733e09 100644 --- a/SerialPrograms/Source/Integrations/PybindSwitchController.h +++ b/SerialPrograms/Source/Integrations/PybindSwitchController.h @@ -32,6 +32,8 @@ namespace NintendoSwitch{ // Button bitfields use `NintendoSwitch::Button` values, and d-pad positions use // `NintendoSwitch::DpadPosition` values (0 = up, clockwise to 7 = up-left, 8 = none). // Joystick coordinates are in [-1.0, 1.0] with +x = right and +y = up. +// +// Thread safety: all methods can be called from any thread. class PybindSwitchProController{ PybindSwitchProController(const PybindSwitchProController&) = delete; void operator=(const PybindSwitchProController&) = delete; @@ -58,6 +60,10 @@ class PybindSwitchProController{ // or the error message if the connection failed. std::string current_status() const; + // Name of the active controller implementation, e.g. "Nintendo Switch: Pro Controller". + // Empty if not ready. + std::string controller_name() const; + // Block until every command queued so far has been executed by the device. // Returns immediately if the controller is not ready. void wait_for_all_requests(); From e8e845c367cde33c13dd97f26909203ebc2cab47 Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:10:21 -0700 Subject: [PATCH 02/16] Add AbstractController::cancel_all_commands_blocking() and use it in PybindSwitchProController cancel_all_commands() only starts the cancellation: it returns as soon as the cancel request is handed to the connection, and the host marks its command queue empty at that moment, so waiting for the queue proves nothing. AbstractController::cancel_all_commands_blocking() cancels, queues a 10 ms neutral no-op, and waits (bounded by a timeout) for the device to report that the no-op finished. Commands are delivered in order, so that report means the device has processed the cancel and is holding the neutral state. It is a non-virtual member built only on AbstractController's virtual interface, so it works for any controller. PybindSwitchProController::cancel_all_commands_blocking() exposes it with a millisecond timeout. Returns false on timeout, e.g. when the Switch is asleep and the device can't execute commands. Safe to call from another thread while one is blocked in wait_for_all_requests(). Co-Authored-By: Claude Opus 5.5 --- .../Source/Controllers/Controller.cpp | 42 +++++++++++++++++++ .../Source/Controllers/Controller.h | 18 ++++++++ .../Integrations/PybindSwitchController.cpp | 16 +++++++ .../Integrations/PybindSwitchController.h | 19 ++++++++- 4 files changed, 94 insertions(+), 1 deletion(-) diff --git a/SerialPrograms/Source/Controllers/Controller.cpp b/SerialPrograms/Source/Controllers/Controller.cpp index 7187cfb646..78df4b0cc4 100644 --- a/SerialPrograms/Source/Controllers/Controller.cpp +++ b/SerialPrograms/Source/Controllers/Controller.cpp @@ -8,6 +8,9 @@ #include "Common/Cpp/ListenerSet.h" #include "Common/Cpp/RecursiveThrottler.h" #include "Common/Cpp/Containers/Pimpl.tpp" +#include "Common/Cpp/Concurrency/Mutex.h" +#include "Common/Cpp/Concurrency/ConditionVariable.h" +#include "Common/Cpp/Concurrency/Thread.h" #include "Controller.h" namespace PokemonAutomation{ @@ -41,6 +44,45 @@ RecursiveThrottler& AbstractController::logging_throttler(){ } +bool AbstractController::cancel_all_commands_blocking(Milliseconds timeout){ + if (!is_ready()){ + return false; + } + cancel_all_commands(); + + // Bound the confirmation: this timer cancels `scope` when the timeout passes, + // which makes `issue_nop()` / `wait_for_all()` below throw + // OperationCancelledException. + CancellableHolder scope; + Mutex lock; + ConditionVariable cv; + bool done = false; + Thread timer([&]{ + std::unique_lock lg(lock); + if (!cv.wait_for(lg, timeout, [&]{ return done; })){ + scope.cancel(nullptr); + } + }); + + bool confirmed = false; + try{ + // Queued after the cancel, so the device reports this no-op finished only + // after it has dropped everything before it and held neutral for 10 ms. + issue_nop(&scope, Milliseconds(10)); + wait_for_all(&scope); + confirmed = true; + }catch (OperationCancelledException&){} + + { + std::lock_guard lg(lock); + done = true; + } + cv.notify_all(); + timer.join(); + return confirmed; +} + + void AbstractController::throw_bad_cast(const char* desired_typename){ throw UserSetupError( logger(), diff --git a/SerialPrograms/Source/Controllers/Controller.h b/SerialPrograms/Source/Controllers/Controller.h index 90df819854..9ed355a794 100644 --- a/SerialPrograms/Source/Controllers/Controller.h +++ b/SerialPrograms/Source/Controllers/Controller.h @@ -112,6 +112,24 @@ class AbstractController{ // ever releasing it during the transition. virtual void replace_on_next_command() = 0; + // Unlike the two functions above, this one blocks: + // Cancel all commands (`cancel_all_commands()`), then wait at most `timeout` for + // the device to confirm that it is in the neutral state. + // + // `cancel_all_commands()` only starts the cancellation: it returns as soon as the + // cancel request is handed to the connection, and the host marks its command + // queue empty at that moment, so waiting for the queue proves nothing. Instead this + // queues a 10 ms neutral no-op after the cancel and waits for the device to report + // that the no-op finished. Commands are delivered in order, so that report means + // the device has processed the cancel and is holding the neutral state. + // + // Returns true if confirmed; false on timeout (e.g. the Switch is asleep and the + // device can't execute commands) or if the controller is not ready. + // Thread-safe: may be called while another thread is issuing commands or waiting on + // this controller. The timeout covers waiting for the device; a command that + // another thread is in the middle of issuing is allowed to finish first. + bool cancel_all_commands_blocking(Milliseconds timeout); + public: // diff --git a/SerialPrograms/Source/Integrations/PybindSwitchController.cpp b/SerialPrograms/Source/Integrations/PybindSwitchController.cpp index da3f4dd57e..e0b034b36b 100644 --- a/SerialPrograms/Source/Integrations/PybindSwitchController.cpp +++ b/SerialPrograms/Source/Integrations/PybindSwitchController.cpp @@ -145,6 +145,22 @@ void PybindSwitchProController::wait_for_all_requests(){ } controller->wait_for_all(nullptr); } +bool PybindSwitchProController::cancel_all_commands_blocking(uint64_t timeout_millis){ + PybindSwitchProControllerInternal* internal = (PybindSwitchProControllerInternal*)m_internals; + ProController* controller = internal->controller_if_ready(); + if (controller == nullptr){ + return false; + } + bool confirmed = controller->cancel_all_commands_blocking(Milliseconds(timeout_millis)); + if (!confirmed){ + internal->m_logger.log( + "cancel_all_commands_blocking(): device did not confirm the neutral state within " + + std::to_string(timeout_millis) + " ms.", + COLOR_RED + ); + } + return confirmed; +} void PybindSwitchProController::wait(uint64_t duration){ PybindSwitchProControllerInternal* internal = (PybindSwitchProControllerInternal*)m_internals; internal->controller().issue_nop(nullptr, Milliseconds(duration)); diff --git a/SerialPrograms/Source/Integrations/PybindSwitchController.h b/SerialPrograms/Source/Integrations/PybindSwitchController.h index a59f733e09..6662eaa843 100644 --- a/SerialPrograms/Source/Integrations/PybindSwitchController.h +++ b/SerialPrograms/Source/Integrations/PybindSwitchController.h @@ -33,7 +33,9 @@ namespace NintendoSwitch{ // `NintendoSwitch::DpadPosition` values (0 = up, clockwise to 7 = up-left, 8 = none). // Joystick coordinates are in [-1.0, 1.0] with +x = right and +y = up. // -// Thread safety: all methods can be called from any thread. +// Thread safety: all methods can be called from any thread. In particular +// `cancel_all_commands_blocking()` may be called while another thread is blocked in +// `wait_for_all_requests()`, e.g. to stop everything in an emergency. class PybindSwitchProController{ PybindSwitchProController(const PybindSwitchProController&) = delete; void operator=(const PybindSwitchProController&) = delete; @@ -68,6 +70,21 @@ class PybindSwitchProController{ // Returns immediately if the controller is not ready. void wait_for_all_requests(); + // Drop every queued command that has not executed yet, return to the neutral + // state (no buttons pressed, sticks centered), and wait for the device to confirm + // it, for at most `timeout_millis`. The controller stays usable afterwards. + // + // Confirmation (see `AbstractController::cancel_all_commands_blocking()`) works by + // queueing a short neutral no-op after the cancel and waiting for the device to + // report that the no-op finished. The serial protocol delivers messages in order, + // so that report means the device has processed the cancel and is holding the + // neutral state. + // + // Returns true if confirmed, false on timeout or if the controller is not ready. + // The timeout covers waiting for the device; if another thread is in the middle + // of issuing a command on this controller, that call is allowed to finish first. + bool cancel_all_commands_blocking(uint64_t timeout_millis); + public: // Commands. These throw InvalidConnectionStateException if the controller is not // ready, and block only if the device's command queue is full. From 2cce136e8c5685e50a28175b146b7b4dc320b0d1 Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:13:17 -0700 Subject: [PATCH 03/16] Python: add the pokemon_automation package and its input vocabulary Start the pokemon_automation Python package (SerialPrograms/Source/PythonBindings/). - _core.py: finds and imports the compiled _pa_core module (added next) lazily, so the pure-Python parts work without it. - buttons.py: the input vocabulary scripts and AI agents use: button names and aliases ("A", "L+R", "start", "L3"), d-pad directions ("up", "up-right"), and joystick positions (direction names or [x, y], +y = up). Bit values match NintendoSwitch::Button and DpadPosition. Co-Authored-By: Claude Opus 5.5 --- .../Source/PythonBindings/.gitignore | 6 + .../pokemon_automation/__init__.py | 18 ++ .../pokemon_automation/_core.py | 57 +++++ .../pokemon_automation/buttons.py | 223 ++++++++++++++++++ .../Source/PythonBindings/pyproject.toml | 27 +++ .../PythonBindings/tests/test_buttons.py | 55 +++++ 6 files changed, 386 insertions(+) create mode 100644 SerialPrograms/Source/PythonBindings/.gitignore create mode 100644 SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py create mode 100644 SerialPrograms/Source/PythonBindings/pokemon_automation/_core.py create mode 100644 SerialPrograms/Source/PythonBindings/pokemon_automation/buttons.py create mode 100644 SerialPrograms/Source/PythonBindings/pyproject.toml create mode 100644 SerialPrograms/Source/PythonBindings/tests/test_buttons.py diff --git a/SerialPrograms/Source/PythonBindings/.gitignore b/SerialPrograms/Source/PythonBindings/.gitignore new file mode 100644 index 0000000000..53101df5e9 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/.gitignore @@ -0,0 +1,6 @@ +# Built by CMake and copied here after each build. +pokemon_automation/_pa_core*.so +pokemon_automation/_pa_core*.pyd +__pycache__/ +*.egg-info/ +.pytest_cache/ diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py new file mode 100644 index 0000000000..5307065205 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py @@ -0,0 +1,18 @@ +"""Pokemon Automation headless console: control a Nintendo Switch from Python. + +Layers (each built on the one below): + +- `_pa_core` (C++, pybind11): acts as a Switch controller through a PABotBase2 serial device, + built only from this codebase's CoreLib. Build it from SerialPrograms with + `-DPA_PYTHON_BINDINGS=ON`. +- `buttons`: the input vocabulary (button names, d-pad, sticks). +""" + +from .buttons import parse_buttons, parse_stick + +__all__ = [ + "parse_buttons", + "parse_stick", +] + +__version__ = "0.1.0" diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/_core.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/_core.py new file mode 100644 index 0000000000..e7115a7127 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/_core.py @@ -0,0 +1,57 @@ +"""Locate and import the compiled `_pa_core` extension module. + +The module is built by CMake (`-DPA_PYTHON_BINDINGS=ON`, target `_pa_core`) and copied +next to this file after every build. Set the environment variable `PA_CORE_PATH` to a +folder containing `_pa_core*.so` / `_pa_core*.pyd` to load it from somewhere else. + +Importing is deferred until first use so that the pure-Python parts of the package +(button parsing, fake devices, the MCP server in `--fake` mode) work without it. +""" + +from __future__ import annotations + +import importlib +import os +import sys +from types import ModuleType + +_module: ModuleType | None = None + +BUILD_HINT = ( + "The compiled module pokemon_automation._pa_core was not found. Build it with:\n" + " cmake -DPA_PYTHON_BINDINGS=ON -DPython_EXECUTABLE=" + sys.executable + "\n" + " cmake --build . --target _pa_core\n" + "or set PA_CORE_PATH to the folder containing the built module." +) + + +def core() -> ModuleType: + """Return the `_pa_core` module, importing it on first call. + + Raises ImportError with build instructions if the module can't be found. + """ + global _module + if _module is not None: + return _module + override = os.environ.get("PA_CORE_PATH") + if override: + sys.path.insert(0, override) + try: + _module = importlib.import_module("_pa_core") + finally: + sys.path.remove(override) + return _module + try: + from . import _pa_core # type: ignore[attr-defined] + except ImportError as e: + raise ImportError(BUILD_HINT) from e + _module = _pa_core + return _module + + +def core_available() -> bool: + try: + core() + return True + except ImportError: + return False diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/buttons.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/buttons.py new file mode 100644 index 0000000000..255bdae33f --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/buttons.py @@ -0,0 +1,223 @@ +"""Controller input vocabulary shared by the Python API and the MCP server. + +Everything a caller (a script or an AI agent) types to describe an input is parsed +here, so both front ends accept exactly the same names: + +- Buttons: "A", "B", "X", "Y", "L", "R", "ZL", "ZR", "PLUS" ("+", "START"), + "MINUS" ("-", "SELECT"), "HOME", "CAPTURE", "LCLICK" ("L3", "LS"), + "RCLICK" ("R3", "RS"). Case-insensitive. +- D-pad: "UP", "DOWN", "LEFT", "RIGHT" and diagonals such as "UP_RIGHT" (also + "UP-RIGHT", "UPRIGHT"). Listing "UP" and "RIGHT" together also means up-right. +- Combinations: a list (["L", "R"]) or a "+"-joined string ("L+R", "ZL+A"). +- Joystick positions: a direction name ("up", "down_left", ...) or an (x, y) pair in + [-1, 1] with +x = right and +y = up. "neutral"/"center" means (0, 0). + +The button bit values match `NintendoSwitch::Button` in +Source/NintendoSwitch/Controllers/NintendoSwitch_ControllerButtons.h. When the +compiled module is available, `check_against_core()` verifies that they still agree. +""" + +from __future__ import annotations + +import math +from collections.abc import Sequence +from dataclasses import dataclass + +# Bit positions from NintendoSwitch_ControllerButtons.h. +BUTTON_BITS: dict[str, int] = { + "Y": 1 << 0, + "B": 1 << 1, + "A": 1 << 2, + "X": 1 << 3, + "L": 1 << 4, + "R": 1 << 5, + "ZL": 1 << 6, + "ZR": 1 << 7, + "MINUS": 1 << 8, + "PLUS": 1 << 9, + "LCLICK": 1 << 10, + "RCLICK": 1 << 11, + "HOME": 1 << 12, + "CAPTURE": 1 << 13, +} + +BUTTON_ALIASES: dict[str, str] = { + "+": "PLUS", + "START": "PLUS", + "-": "MINUS", + "SELECT": "MINUS", + "L3": "LCLICK", + "LS": "LCLICK", + "LSTICK": "LCLICK", + "R3": "RCLICK", + "RS": "RCLICK", + "RSTICK": "RCLICK", + "SCREENSHOT": "CAPTURE", +} + +# D-pad positions from `NintendoSwitch::DpadPosition`: 0 = up, clockwise, 8 = none. +DPAD_POSITIONS: dict[str, int] = { + "UP": 0, + "UP_RIGHT": 1, + "RIGHT": 2, + "DOWN_RIGHT": 3, + "DOWN": 4, + "DOWN_LEFT": 5, + "LEFT": 6, + "UP_LEFT": 7, +} +DPAD_NONE = 8 + +# Unit vectors for the 8 directions, used for both the d-pad and joysticks. +_DIRECTION_VECTORS: dict[str, tuple[int, int]] = { + "UP": (0, 1), + "UP_RIGHT": (1, 1), + "RIGHT": (1, 0), + "DOWN_RIGHT": (1, -1), + "DOWN": (0, -1), + "DOWN_LEFT": (-1, -1), + "LEFT": (-1, 0), + "UP_LEFT": (-1, 1), +} + +ButtonsArg = str | Sequence[str] | None +StickArg = str | Sequence[float] | None + + +def _normalize_name(name: str) -> str: + name = name.strip().upper().replace("-", "_").replace(" ", "_") + # "UPRIGHT" -> "UP_RIGHT", "DPAD_UP" -> "UP" + if name.startswith("DPAD_"): + name = name[5:] + for vertical in ("UP", "DOWN"): + for horizontal in ("LEFT", "RIGHT"): + if name in (vertical + horizontal, horizontal + vertical, horizontal + "_" + vertical): + return vertical + "_" + horizontal + return name + + +def _split(buttons: ButtonsArg) -> list[str]: + if buttons is None: + return [] + if isinstance(buttons, str): + text = buttons.strip() + # A lone "+" or "-" is the PLUS/MINUS button, not a separator. + if text in ("+", "-"): + return [text] + return [p for p in text.replace(",", "+").split("+") if p.strip()] + ret: list[str] = [] + for item in buttons: + ret.extend(_split(item)) + return ret + + +@dataclass(frozen=True) +class ParsedButtons: + """The result of parsing a button combination.""" + + bitfield: int + dpad: int # DPAD_NONE if no d-pad direction was given + + @property + def has_buttons(self) -> bool: + return self.bitfield != 0 + + @property + def has_dpad(self) -> bool: + return self.dpad != DPAD_NONE + + +def parse_buttons(buttons: ButtonsArg) -> ParsedButtons: + """Parse a button combination into a bitfield plus a d-pad position. + + Examples: + parse_buttons("A") -> bitfield A, no d-pad + parse_buttons("L+R") -> bitfield L|R + parse_buttons(["ZL", "up"]) -> bitfield ZL, d-pad up + parse_buttons("up+right") -> d-pad up-right + + Raises ValueError for unknown names or contradictory d-pad directions + (e.g. "up+down"). + """ + bitfield = 0 + dx = dy = 0 + seen_dpad: list[str] = [] + for raw in _split(buttons): + name = raw.strip() + key = name.upper() if name in ("+", "-") else _normalize_name(name) + key = BUTTON_ALIASES.get(key, key) + if key in BUTTON_BITS: + bitfield |= BUTTON_BITS[key] + continue + if key in _DIRECTION_VECTORS: + vx, vy = _DIRECTION_VECTORS[key] + if (vx and dx and vx != dx) or (vy and dy and vy != dy): + raise ValueError(f"Contradictory d-pad directions: {seen_dpad + [name]}") + dx = vx or dx + dy = vy or dy + seen_dpad.append(name) + continue + raise ValueError( + f"Unknown button {name!r}. Valid buttons: {', '.join(BUTTON_BITS)}, " + f"d-pad: {', '.join(DPAD_POSITIONS)}." + ) + dpad = DPAD_NONE + if dx or dy: + for direction, (vx, vy) in _DIRECTION_VECTORS.items(): + if (vx, vy) == (dx, dy): + dpad = DPAD_POSITIONS[direction] + return ParsedButtons(bitfield, dpad) + + +def parse_stick(position: StickArg) -> tuple[float, float]: + """Parse a joystick position into (x, y) in [-1, 1], +y = up. + + Accepts a direction name ("up", "down_left", "neutral"), or an (x, y) pair. + Diagonal names are normalized to length 1 so they tilt the stick fully. + Raises ValueError for unknown names or out-of-range coordinates. + """ + if position is None: + return (0.0, 0.0) + if isinstance(position, str): + key = _normalize_name(position) + if key in ("NEUTRAL", "CENTER", "NONE"): + return (0.0, 0.0) + if key not in _DIRECTION_VECTORS: + raise ValueError( + f"Unknown stick direction {position!r}. Use one of " + f"{', '.join(d.lower() for d in _DIRECTION_VECTORS)}, or an [x, y] pair." + ) + vx, vy = _DIRECTION_VECTORS[key] + length = math.hypot(vx, vy) + return (vx / length, vy / length) + values = list(position) + if len(values) != 2: + raise ValueError(f"A stick position needs exactly two numbers [x, y], got {position!r}.") + x, y = float(values[0]), float(values[1]) + if not (-1.0 <= x <= 1.0 and -1.0 <= y <= 1.0): + raise ValueError(f"Stick coordinates must be within [-1, 1], got ({x}, {y}).") + return (x, y) + + +def button_names(bitfield: int) -> list[str]: + """Inverse of the bitfield part of `parse_buttons()`, for logging.""" + return [name for name, bit in BUTTON_BITS.items() if bitfield & bit] + + +def dpad_name(position: int) -> str | None: + for name, value in DPAD_POSITIONS.items(): + if value == position: + return name + return None + + +def check_against_core(core_buttons: dict[str, int]) -> None: + """Raise AssertionError if `BUTTON_BITS` disagrees with the C++ enum. + + `core_buttons` is `_pa_core.BUTTONS`, generated from `NintendoSwitch::Button`. + """ + for name, bit in BUTTON_BITS.items(): + if core_buttons.get(name) != bit: + raise AssertionError( + f"Button {name} is {bit:#x} in buttons.py but {core_buttons.get(name)!r} in _pa_core." + ) diff --git a/SerialPrograms/Source/PythonBindings/pyproject.toml b/SerialPrograms/Source/PythonBindings/pyproject.toml new file mode 100644 index 0000000000..f38b27930b --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pyproject.toml @@ -0,0 +1,27 @@ +[build-system] +requires = ["setuptools>=68"] +build-backend = "setuptools.build_meta" + +# The compiled `_pa_core` module (acts as a Switch controller) is built by CMake (see +# README.md) and copied into pokemon_automation/. This file packages the Python side +# for `pip install -e .`. + +[project] +name = "pokemon-automation" +version = "0.1.0" +description = "Control a Nintendo Switch from Python with Pokemon Automation hardware." +requires-python = ">=3.10" +dependencies = [] + +[project.optional-dependencies] +serial = ["pyserial>=3.5"] +test = ["pytest>=7"] + +[tool.setuptools] +packages = ["pokemon_automation"] + +[tool.setuptools.package-data] +pokemon_automation = ["_pa_core*.so", "_pa_core*.pyd"] + +[tool.pytest.ini_options] +testpaths = ["tests"] diff --git a/SerialPrograms/Source/PythonBindings/tests/test_buttons.py b/SerialPrograms/Source/PythonBindings/tests/test_buttons.py new file mode 100644 index 0000000000..38b54b4894 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/tests/test_buttons.py @@ -0,0 +1,55 @@ +import math + +import pytest + +from pokemon_automation import buttons as btn + + +def test_single_and_combined_buttons(): + assert btn.parse_buttons("A").bitfield == btn.BUTTON_BITS["A"] + both = btn.parse_buttons("L+R") + assert both.bitfield == btn.BUTTON_BITS["L"] | btn.BUTTON_BITS["R"] + assert not both.has_dpad + assert btn.parse_buttons(["zl", "a"]).bitfield == btn.BUTTON_BITS["ZL"] | btn.BUTTON_BITS["A"] + + +def test_aliases_and_plus_minus(): + assert btn.parse_buttons("+").bitfield == btn.BUTTON_BITS["PLUS"] + assert btn.parse_buttons("-").bitfield == btn.BUTTON_BITS["MINUS"] + assert btn.parse_buttons("start").bitfield == btn.BUTTON_BITS["PLUS"] + assert btn.parse_buttons("L3").bitfield == btn.BUTTON_BITS["LCLICK"] + assert btn.parse_buttons(["A", "+"]).bitfield == btn.BUTTON_BITS["A"] | btn.BUTTON_BITS["PLUS"] + + +def test_dpad(): + assert btn.parse_buttons("up").dpad == btn.DPAD_POSITIONS["UP"] + assert btn.parse_buttons("up+right").dpad == btn.DPAD_POSITIONS["UP_RIGHT"] + assert btn.parse_buttons("down-left").dpad == btn.DPAD_POSITIONS["DOWN_LEFT"] + assert btn.parse_buttons("UPLEFT").dpad == btn.DPAD_POSITIONS["UP_LEFT"] + mixed = btn.parse_buttons("ZL+DOWN") + assert mixed.bitfield == btn.BUTTON_BITS["ZL"] and mixed.dpad == btn.DPAD_POSITIONS["DOWN"] + assert btn.parse_buttons(None).dpad == btn.DPAD_NONE + + +def test_invalid_buttons(): + with pytest.raises(ValueError, match="Unknown button"): + btn.parse_buttons("Q") + with pytest.raises(ValueError, match="Contradictory"): + btn.parse_buttons("up+down") + + +def test_sticks(): + assert btn.parse_stick("up") == (0.0, 1.0) + assert btn.parse_stick("neutral") == (0.0, 0.0) + x, y = btn.parse_stick("down_right") + assert math.isclose(math.hypot(x, y), 1.0) and x > 0 and y < 0 + assert btn.parse_stick([0.5, -0.25]) == (0.5, -0.25) + with pytest.raises(ValueError): + btn.parse_stick([2, 0]) + with pytest.raises(ValueError): + btn.parse_stick("sideways") + + +def test_matches_cpp_enum(): + core = pytest.importorskip("pokemon_automation._pa_core") + btn.check_against_core(dict(core.BUTTONS)) From b7c40c9c171a54a9c6ffa1aacdcb4d84e7174cd7 Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:13:17 -0700 Subject: [PATCH 04/16] Add the _pa_core pybind11 module (CMake option PA_PYTHON_BINDINGS) _pa_core wraps PybindSwitchProController for Python. It links only CoreLib (no Qt, OpenCV or Tesseract), so it also builds with -DPA_CORE_ONLY=ON. - CMake option PA_PYTHON_BINDINGS (OFF by default, since it needs a Python interpreter with development headers), in both the full and the core-only build. pybind11 is taken from the environment or downloaded. - The built module is copied into the pokemon_automation package folder. - A log sink that never writes to stdout (stdout carries the protocol for MCP over stdio), with recent lines available to Python. Co-Authored-By: Claude Opus 5.5 --- SerialPrograms/CMakeLists.txt | 17 +- .../Integrations/PybindSwitchController.h | 8 +- .../PythonBindings/PythonBindings.cmake | 60 +++++ .../PythonBindings/PythonBindings_Module.cpp | 221 ++++++++++++++++++ .../PythonBindings/tests/test_core_module.py | 24 ++ SerialPrograms/cmake/CoreLib.cmake | 1 + 6 files changed, 327 insertions(+), 4 deletions(-) create mode 100644 SerialPrograms/Source/PythonBindings/PythonBindings.cmake create mode 100644 SerialPrograms/Source/PythonBindings/PythonBindings_Module.cpp create mode 100644 SerialPrograms/Source/PythonBindings/tests/test_core_module.py diff --git a/SerialPrograms/CMakeLists.txt b/SerialPrograms/CMakeLists.txt index d456953b4b..6634e2dfcc 100644 --- a/SerialPrograms/CMakeLists.txt +++ b/SerialPrograms/CMakeLists.txt @@ -63,9 +63,13 @@ find_package(Threads REQUIRED) # Define a toggle option (ON by default) option(COMMAND_LINE_EXE "Build the optional executable" ON) -# Build only the GUI-free core: CoreLib and SerialProgramsCommandLine. This needs -# no Qt, OpenCV, ONNX Runtime, Tesseract, DPP or Discord SDK, so none of them are -# searched for. +# Python bindings and the MCP server for AI agents (see Source/PythonBindings/README.md). +# OFF by default because it needs a Python interpreter with development headers. +option(PA_PYTHON_BINDINGS "Build the _pa_core Python module (acts as a Switch controller)" OFF) + +# Build only the GUI-free core: CoreLib, SerialProgramsCommandLine and (with +# PA_PYTHON_BINDINGS) the Python module. This needs no Qt, OpenCV, ONNX Runtime, +# Tesseract, DPP or Discord SDK, so none of them are searched for. option(PA_CORE_ONLY "Build only the GUI-free core (basic microcontroller functions, no Qt, OpenCV, etc.)" OFF) if(PA_CORE_ONLY) @@ -78,6 +82,9 @@ if(PA_CORE_ONLY) if(COMMAND_LINE_EXE) include(Source/CommandLine/CommandLineExecutable.cmake) endif() + if(PA_PYTHON_BINDINGS) + include(Source/PythonBindings/PythonBindings.cmake) + endif() return() endif() @@ -846,3 +853,7 @@ if(COMMAND_LINE_EXE) # Add command-line executable (GUI-free) from subdirectory include(Source/CommandLine/CommandLineExecutable.cmake) endif() + +if(PA_PYTHON_BINDINGS) + include(Source/PythonBindings/PythonBindings.cmake) +endif() diff --git a/SerialPrograms/Source/Integrations/PybindSwitchController.h b/SerialPrograms/Source/Integrations/PybindSwitchController.h index 6662eaa843..01cd3d4a96 100644 --- a/SerialPrograms/Source/Integrations/PybindSwitchController.h +++ b/SerialPrograms/Source/Integrations/PybindSwitchController.h @@ -2,6 +2,12 @@ * * From: https://github.com/PokemonAutomation/ * + * A GUI-free Nintendo Switch controller that talks to a PABotBase2 microcontroller + * over a serial port. It is part of CoreLib and uses only plain types in its + * interface, so it can be bound to other languages: the `_pa_core` Python module + * (Source/PythonBindings/) wraps it, and the Python package and MCP server are built + * on top of that. SerialProgramsCommandLine also uses it directly to test its + * functionality. */ #ifndef PokemonAutomation_Integrations_PybindSwitchController_H @@ -35,7 +41,7 @@ namespace NintendoSwitch{ // // Thread safety: all methods can be called from any thread. In particular // `cancel_all_commands_blocking()` may be called while another thread is blocked in -// `wait_for_all_requests()`, e.g. to stop everything in an emergency. +// `wait_for_all_requests()`, which is how the Python layer implements an emergency stop. class PybindSwitchProController{ PybindSwitchProController(const PybindSwitchProController&) = delete; void operator=(const PybindSwitchProController&) = delete; diff --git a/SerialPrograms/Source/PythonBindings/PythonBindings.cmake b/SerialPrograms/Source/PythonBindings/PythonBindings.cmake new file mode 100644 index 0000000000..b90613ac56 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/PythonBindings.cmake @@ -0,0 +1,60 @@ +# This cmake file is included by CMakeLists.txt when PA_PYTHON_BINDINGS is ON. +# +# It builds `_pa_core`, a pybind11 extension module that acts as a Nintendo Switch +# controller through a PABotBase2 serial device. It links only CoreLib (this +# codebase, GUI-free); it does not need Qt at runtime, OpenCV or Tesseract. Video +# capture and OCR are done by the pure-Python side of the `pokemon_automation` +# package with opencv-python / pytesseract. +# +# After each build the module is copied into Source/PythonBindings/pokemon_automation/ +# so the package (and the MCP server in it) can be used straight from the source tree +# with `pip install -e Source/PythonBindings`. +# +# Configure with, for example: +# cmake .. -DPA_PYTHON_BINDINGS=ON -DPython_EXECUTABLE=$(which python3) +# and build with: +# cmake --build . --target _pa_core -j 10 + +# The extension module is a shared library, so everything linked into it must be +# position-independent. +set_target_properties(CoreLib PROPERTIES POSITION_INDEPENDENT_CODE ON) + +# pybind11: use an installed copy if there is one, otherwise download it. +set(PYBIND11_FINDPYTHON ON) +find_package(Python 3.10 REQUIRED COMPONENTS Interpreter Development.Module) +find_package(pybind11 CONFIG QUIET) +if (NOT pybind11_FOUND) + execute_process( + COMMAND "${Python_EXECUTABLE}" -m pybind11 --cmakedir + OUTPUT_VARIABLE PYBIND11_PIP_CMAKE_DIR + OUTPUT_STRIP_TRAILING_WHITESPACE + ERROR_QUIET + ) + if (PYBIND11_PIP_CMAKE_DIR) + find_package(pybind11 CONFIG QUIET PATHS "${PYBIND11_PIP_CMAKE_DIR}" NO_DEFAULT_PATH) + endif() +endif() +if (NOT pybind11_FOUND) + message(STATUS "pybind11 not found, downloading it") + include(FetchContent) + FetchContent_Declare( + pybind11 + GIT_REPOSITORY https://github.com/pybind/pybind11.git + GIT_TAG v3.0.1 + ) + FetchContent_MakeAvailable(pybind11) +endif() +message(STATUS "Python bindings: building _pa_core for ${Python_EXECUTABLE} (${Python_VERSION})") + +# The controller class it wraps (Source/Integrations/PybindSwitchController.*) is +# part of CoreLib, so only the module file is compiled here. +pybind11_add_module(_pa_core Source/PythonBindings/PythonBindings_Module.cpp) +pa_apply_gui_free_target_properties(_pa_core) +target_link_libraries(_pa_core PRIVATE CoreLib) + +set(PA_PYTHON_PACKAGE_DIR ${CMAKE_CURRENT_SOURCE_DIR}/Source/PythonBindings/pokemon_automation) +add_custom_command( + TARGET _pa_core POST_BUILD + COMMAND ${CMAKE_COMMAND} -E copy $ ${PA_PYTHON_PACKAGE_DIR}/ + COMMENT "Copying _pa_core into the pokemon_automation Python package" +) diff --git a/SerialPrograms/Source/PythonBindings/PythonBindings_Module.cpp b/SerialPrograms/Source/PythonBindings/PythonBindings_Module.cpp new file mode 100644 index 0000000000..6723cfa6f7 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/PythonBindings_Module.cpp @@ -0,0 +1,221 @@ +/* Python Bindings: `_pa_core` Module + * + * From: https://github.com/PokemonAutomation/ + * + * pybind11 module that lets Python act as a Switch controller (serial port + + * PABotBase2 + controller scheduling from CoreLib). It is built only from + * this codebase: no Qt, OpenCV or Tesseract. Video capture and OCR live in the + * pure-Python `pokemon_automation` package next to this file, using opencv-python + * and pytesseract. + * + * The controller class is `NintendoSwitch::PybindSwitchProController` + * (Source/Integrations/PybindSwitchController.h). Methods are bound under the names + * the Python package uses (e.g. `push_button()` as `press_buttons()`). + * + * This is a thin 1:1 layer. Button-name parsing, input sequences and the MCP server + * are in the Python package, which imports this module as + * `pokemon_automation._pa_core`. + * + * Conventions: + * - All durations are milliseconds. + * - Every call that can block (waiting for the device or a full command queue) + * releases the GIL, so other Python threads (e.g. an emergency stop) keep running. + * - C++ `PokemonAutomation::Exception`s become Python `RuntimeError`s. + */ + +#include +#include +#include +#include +#include +#include "Common/Cpp/Exceptions.h" +#include "Common/Cpp/Time.h" +#include "Common/Cpp/Concurrency/Mutex.h" +#include "Common/Cpp/Logging/GlobalLogger.h" +#include "Common/Cpp/Logging/MultiOutputLogger.h" +#include "NintendoSwitch/Controllers/NintendoSwitch_ControllerButtons.h" +#include "Integrations/PybindSwitchController.h" + +namespace py = pybind11; + +namespace PokemonAutomation{ + +// Defined by every executable that links CoreLib. The Python module has no Qt UI. +bool USE_QT_UI = false; + +namespace PythonBindings{ +namespace{ + + + +// Receives every line written to the global logger and forwards it to stderr and/or a +// log file, while keeping the most recent lines in memory for `recent_logs()`. +// +// Nothing is ever written to stdout: we will use this Pybind module to run an AI agent +// MCP server. When the MCP server runs over stdio, stdout carries the JSON-RPC protocol +// and any stray line would corrupt it. +class PythonLogSink : public Logger{ +public: + static PythonLogSink& instance(){ + static PythonLogSink sink; + return sink; + } + + void set_stderr(bool enabled){ + std::lock_guard lg(m_lock); + m_stderr = enabled; + } + // Append log lines to `path`. An empty path stops file logging. + // Throws FileException if the file can't be opened. + void set_file(const std::string& path){ + std::lock_guard lg(m_lock); + m_file.close(); + if (!path.empty()){ + m_file.open(path, std::ios::app); + if (!m_file){ + throw FileException(nullptr, PA_CURRENT_FUNCTION, "Unable to open log file.", path); + } + } + } + std::vector recent(size_t count) const{ + std::lock_guard lg(m_lock); + count = std::min(count, m_recent.size()); + return std::vector(m_recent.end() - count, m_recent.end()); + } + + // Lines from `TaggedLogger`s already start with a timestamp. + virtual void log(const std::string& msg, Color color = Color()) override{ + std::string line = msg; + std::lock_guard lg(m_lock); + if (m_stderr){ + std::cerr << line << std::endl; + } + if (m_file.is_open()){ + m_file << line << std::endl; + } + m_recent.emplace_back(std::move(line)); + if (m_recent.size() > 1000){ + m_recent.pop_front(); + } + } + +private: + PythonLogSink(){ + global_multi_logger().add_listener(*this); + } + + mutable Mutex m_lock; + bool m_stderr = false; + std::ofstream m_file; + std::deque m_recent; +}; + + + +} +} +} + + +using namespace PokemonAutomation; +using namespace PokemonAutomation::PythonBindings; +using NintendoSwitch::PybindSwitchProController; + + +PYBIND11_MODULE(_pa_core, m){ + m.doc() = "Pokemon Automation: act as a Nintendo Switch controller through a PABotBase2 serial device."; + + py::register_exception_translator([](std::exception_ptr p){ + try{ + if (p){ + std::rethrow_exception(p); + } + }catch (const PokemonAutomation::Exception& e){ + PyErr_SetString(PyExc_RuntimeError, e.to_str().c_str()); + } + }); + + // Make sure the log sink is attached before anything logs. + PythonLogSink::instance(); + + + // Logging + + m.def( + "set_log_stderr", + [](bool enabled){ PythonLogSink::instance().set_stderr(enabled); }, + py::arg("enabled"), + "Echo internal log lines to stderr." + ); + m.def( + "set_log_file", + [](const std::string& path){ PythonLogSink::instance().set_file(path); }, + py::arg("path"), + "Append internal log lines to this file. Pass an empty string to stop." + ); + m.def( + "log", + [](const std::string& message){ + global_logger_raw().log(current_time_to_str() + " - " + message); + }, + py::arg("message"), + "Write a line to the internal log, so Python events appear alongside C++ ones." + ); + m.def( + "recent_logs", + [](size_t count){ return PythonLogSink::instance().recent(count); }, + py::arg("count") = 50, + "Return up to `count` of the most recent internal log lines." + ); + + + // Button constants, so the Python side never hard-codes bit positions. + // Keys are the C++ enum names without the "BUTTON_" prefix, e.g. "A", "ZL", + // "PLUS", "LCLICK", "HOME", "UP". + + py::dict buttons; + for (size_t bit = 0; bit < NintendoSwitch::TOTAL_BUTTONS; bit++){ + NintendoSwitch::Button button = (NintendoSwitch::Button)((uint32_t)1 << bit); + std::string name = NintendoSwitch::button_to_code_string(button); + if (name.starts_with("BUTTON_")){ + name = name.substr(7); + } + buttons[py::str(name)] = (uint32_t)button; + } + m.attr("BUTTONS") = buttons; + + + // Controller + + py::class_(m, "Controller", + "A Switch controller on a PABotBase2 device. Commands are queued and return " + "immediately; call wait_for_all() to block until they have executed.") + .def(py::init(), py::arg("port_name")) + .def("wait_for_ready", &PybindSwitchProController::wait_for_ready, + py::arg("timeout_ms"), py::call_guard()) + .def("is_ready", &PybindSwitchProController::is_ready) + .def("status_text", &PybindSwitchProController::current_status) + .def("controller_name", &PybindSwitchProController::controller_name) + .def("wait_for_all", &PybindSwitchProController::wait_for_all_requests, + py::call_guard()) + .def("cancel_all_commands_blocking", &PybindSwitchProController::cancel_all_commands_blocking, + py::arg("timeout_ms"), py::call_guard()) + .def("wait", &PybindSwitchProController::wait, + py::arg("duration_ms"), py::call_guard()) + .def("press_buttons", &PybindSwitchProController::push_button, + py::arg("delay_ms"), py::arg("hold_ms"), py::arg("release_ms"), py::arg("buttons"), + py::call_guard()) + .def("press_dpad", &PybindSwitchProController::push_dpad, + py::arg("delay_ms"), py::arg("hold_ms"), py::arg("release_ms"), py::arg("position"), + py::call_guard()) + .def("move_left_joystick", &PybindSwitchProController::push_left_joystick, + py::arg("delay_ms"), py::arg("hold_ms"), py::arg("release_ms"), py::arg("x"), py::arg("y"), + py::call_guard()) + .def("move_right_joystick", &PybindSwitchProController::push_right_joystick, + py::arg("delay_ms"), py::arg("hold_ms"), py::arg("release_ms"), py::arg("x"), py::arg("y"), + py::call_guard()) + .def("set_state", &PybindSwitchProController::controller_state, + py::arg("duration_ms"), py::arg("buttons"), py::arg("dpad"), + py::arg("left_x"), py::arg("left_y"), py::arg("right_x"), py::arg("right_y"), + py::call_guard()); +} diff --git a/SerialPrograms/Source/PythonBindings/tests/test_core_module.py b/SerialPrograms/Source/PythonBindings/tests/test_core_module.py new file mode 100644 index 0000000000..371ec50749 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/tests/test_core_module.py @@ -0,0 +1,24 @@ +"""Tests of the compiled `_pa_core` module that don't need hardware.""" + +import pytest + +core = pytest.importorskip("pokemon_automation._pa_core") + +from pokemon_automation import buttons as btn # noqa: E402 + + +def test_button_bits_match_cpp_enum(): + btn.check_against_core(dict(core.BUTTONS)) + + +def test_log_roundtrip(): + core.log("hello from the test") + assert any("hello from the test" in line for line in core.recent_logs(10)) + + +def test_missing_port_is_not_ready(): + controller = core.Controller("/dev/does-not-exist") + assert controller.wait_for_ready(3000) is False + assert not controller.is_ready() + with pytest.raises(RuntimeError, match="not ready"): + controller.press_buttons(80, 80, 0, btn.BUTTON_BITS["A"]) diff --git a/SerialPrograms/cmake/CoreLib.cmake b/SerialPrograms/cmake/CoreLib.cmake index 0494b954f9..e9c9ae01fe 100644 --- a/SerialPrograms/cmake/CoreLib.cmake +++ b/SerialPrograms/cmake/CoreLib.cmake @@ -9,6 +9,7 @@ # # Users of CoreLib: # - SerialProgramsCommandLine (Source/CommandLine/) +# - the `_pa_core` Python module (Source/PythonBindings/) # The GUI program does not link CoreLib; SerialProgramsLib compiles the same # sources itself with Qt enabled. # From 2136630b883c818692fdea6d711fe67c91ac1679 Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:13:17 -0700 Subject: [PATCH 05/16] Python: add SwitchController and input steps - SwitchController: press buttons/d-pad, tilt sticks, hold combinations, run sequences of InputSteps, flush and cancel_all_commands_blocking. - InputStep: one step of a sequence (buttons, sticks, hold, release, repeat, wait). The same step dicts are used by the MCP run_inputs tool. - FakeController: records commands, for tests without hardware. Co-Authored-By: Claude Opus 5.5 --- .../pokemon_automation/__init__.py | 5 +- .../pokemon_automation/controller.py | 278 ++++++++++++++++++ .../PythonBindings/pokemon_automation/fake.py | 71 +++++ .../PythonBindings/tests/test_controller.py | 79 +++++ 4 files changed, 432 insertions(+), 1 deletion(-) create mode 100644 SerialPrograms/Source/PythonBindings/pokemon_automation/controller.py create mode 100644 SerialPrograms/Source/PythonBindings/pokemon_automation/fake.py create mode 100644 SerialPrograms/Source/PythonBindings/tests/test_controller.py diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py index 5307065205..177e5db034 100644 --- a/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py @@ -5,12 +5,15 @@ - `_pa_core` (C++, pybind11): acts as a Switch controller through a PABotBase2 serial device, built only from this codebase's CoreLib. Build it from SerialPrograms with `-DPA_PYTHON_BINDINGS=ON`. -- `buttons`: the input vocabulary (button names, d-pad, sticks). +- `SwitchController`: the Python API for the controller. """ from .buttons import parse_buttons, parse_stick +from .controller import InputStep, SwitchController __all__ = [ + "InputStep", + "SwitchController", "parse_buttons", "parse_stick", ] diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/controller.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/controller.py new file mode 100644 index 0000000000..e64ffa6c53 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/controller.py @@ -0,0 +1,278 @@ +"""High-level Nintendo Switch controller API on top of `_pa_core.Controller`. + +Commands are queued on the microcontroller and return immediately, like `pbf_*()` +functions in the main C++ program. Call `flush()` to wait until everything queued so +far has executed. `cancel_all_commands_blocking()` cancels anything still queued, +releases all inputs and waits for the device to confirm; it may be called from +another thread. + +Example: + with SwitchController("/dev/cu.usbserial-0001") as sw: + sw.press("A") # tap A (80 ms hold, 80 ms release) + sw.press("L+R", hold_ms=200) # press L and R together + sw.press("down", repeat=3) # d-pad down three times + sw.stick("left", "up", duration_ms=1500) # walk forward + sw.hold(["B"], left="up", duration_ms=2000) # run forward + sw.flush() +""" + +from __future__ import annotations + +import threading +from collections.abc import Iterable, Mapping, Sequence +from dataclasses import dataclass, field +from typing import Any, Protocol + +from . import buttons as btn +from ._core import core + +DEFAULT_HOLD_MS = 80 +DEFAULT_RELEASE_MS = 80 + + +class ControllerBackend(Protocol): + """The subset of `_pa_core.Controller` used here. `fake.FakeController` implements it too.""" + + def wait_for_ready(self, timeout_ms: int) -> bool: ... + def is_ready(self) -> bool: ... + def status_text(self) -> str: ... + def controller_name(self) -> str: ... + def wait_for_all(self) -> None: ... + def cancel_all_commands_blocking(self, timeout_ms: int) -> bool: ... + def wait(self, duration_ms: int) -> None: ... + def press_buttons(self, delay_ms: int, hold_ms: int, release_ms: int, buttons: int) -> None: ... + def press_dpad(self, delay_ms: int, hold_ms: int, release_ms: int, position: int) -> None: ... + def move_left_joystick(self, delay_ms: int, hold_ms: int, release_ms: int, x: float, y: float) -> None: ... + def move_right_joystick(self, delay_ms: int, hold_ms: int, release_ms: int, x: float, y: float) -> None: ... + def set_state( + self, duration_ms: int, buttons: int, dpad: int, + left_x: float, left_y: float, right_x: float, right_y: float, + ) -> None: ... + + +@dataclass +class InputStep: + """One step of an input sequence. This is also the schema the MCP server exposes. + + A step either waits (`wait_ms` > 0 and nothing else set) or sets an input state: + the given `buttons` (buttons and/or d-pad directions) and stick positions are held + together for `hold_ms`, then everything is released for `release_ms`. The whole + step is repeated `repeat` times. + + Examples: + InputStep(wait_ms=1000) + InputStep(buttons="A") + InputStep(buttons="ZL+A", hold_ms=100) + InputStep(left_stick="up", hold_ms=2000, release_ms=0) + InputStep(buttons="B", left_stick=[0.5, 1.0], hold_ms=1500) + """ + + buttons: btn.ButtonsArg = None + left_stick: btn.StickArg = None + right_stick: btn.StickArg = None + hold_ms: int = DEFAULT_HOLD_MS + release_ms: int = DEFAULT_RELEASE_MS + repeat: int = 1 + wait_ms: int = 0 + + @classmethod + def from_dict(cls, data: Mapping[str, Any]) -> InputStep: + unknown = set(data) - {f for f in cls.__dataclass_fields__} + if unknown: + raise ValueError(f"Unknown input step field(s): {sorted(unknown)}") + return cls(**data) + + def is_wait(self) -> bool: + return self.buttons in (None, "", []) and self.left_stick is None and self.right_stick is None + + def duration_ms(self) -> int: + """Total time this step occupies on the controller.""" + if self.is_wait(): + return max(0, self.wait_ms) + return max(1, self.repeat) * (self.hold_ms + self.release_ms) + max(0, self.wait_ms) + + +@dataclass +class _ParsedStep: + step: InputStep + pressed: btn.ParsedButtons = field(default_factory=lambda: btn.ParsedButtons(0, btn.DPAD_NONE)) + left: tuple[float, float] = (0.0, 0.0) + right: tuple[float, float] = (0.0, 0.0) + + +def _parse_step(step: InputStep) -> _ParsedStep: + if step.hold_ms < 0 or step.release_ms < 0 or step.wait_ms < 0: + raise ValueError("Durations must not be negative.") + if step.repeat < 1: + raise ValueError("repeat must be at least 1.") + if step.is_wait(): + return _ParsedStep(step) + if step.hold_ms == 0: + raise ValueError("hold_ms must be positive for a step that presses something.") + return _ParsedStep( + step, + btn.parse_buttons(step.buttons), + btn.parse_stick(step.left_stick) if step.left_stick is not None else (0.0, 0.0), + btn.parse_stick(step.right_stick) if step.right_stick is not None else (0.0, 0.0), + ) + + +class SwitchController: + """A Switch controller on a PABotBase2 device (ESP32/Pico running PA firmware). + + `port` is the serial port, e.g. "/dev/cu.usbserial-0001" on macOS or "COM3" on + Windows. See `devices.list_serial_ports()`. + + Pass `backend` to wrap an existing `_pa_core.Controller` or a + `fake.FakeController` instead of opening a port. + + Raises RuntimeError if the device isn't ready within `timeout_s` seconds. + """ + + def __init__(self, port: str | None = None, *, timeout_s: float = 10.0, + backend: ControllerBackend | None = None): + if backend is None: + if port is None: + raise ValueError("Either port or backend is required.") + backend = core().Controller(port) + self._backend = backend + self._lock = threading.Lock() + self.port = port + if not self._backend.wait_for_ready(int(timeout_s * 1000)): + status = self._backend.status_text() + raise RuntimeError(f"Controller on {port} is not ready: {status}") + + # ---- status ------------------------------------------------------------- + + @property + def backend(self) -> ControllerBackend: + return self._backend + + def is_ready(self) -> bool: + return self._backend.is_ready() + + def status(self) -> str: + return self._backend.status_text() + + def name(self) -> str: + return self._backend.controller_name() + + # ---- basic commands ----------------------------------------------------- + + def press(self, buttons: btn.ButtonsArg, hold_ms: int = DEFAULT_HOLD_MS, + release_ms: int = DEFAULT_RELEASE_MS, repeat: int = 1) -> None: + """Tap a button combination `repeat` times. D-pad directions may be mixed in.""" + self.run([InputStep(buttons=buttons, hold_ms=hold_ms, release_ms=release_ms, repeat=repeat)]) + + def stick(self, side: str, position: btn.StickArg, duration_ms: int = 500, + release_ms: int = 0) -> None: + """Tilt the "left" or "right" stick to `position` for `duration_ms`.""" + side = side.lower() + if side not in ("left", "right"): + raise ValueError('side must be "left" or "right".') + step = InputStep(hold_ms=duration_ms, release_ms=release_ms) + setattr(step, side + "_stick", position) + self.run([step]) + + def hold(self, buttons: btn.ButtonsArg = None, duration_ms: int = 500, *, + left: btn.StickArg = None, right: btn.StickArg = None, release_ms: int = 0) -> None: + """Hold any combination of buttons, d-pad and sticks for `duration_ms`.""" + self.run([InputStep(buttons=buttons, left_stick=left, right_stick=right, + hold_ms=duration_ms, release_ms=release_ms)]) + + def wait(self, duration_ms: int) -> None: + """Queue a pause with nothing pressed.""" + if duration_ms > 0: + self._backend.wait(int(duration_ms)) + + def run(self, steps: Iterable[InputStep | Mapping[str, Any]]) -> int: + """Queue a sequence of steps. Returns the total duration in milliseconds. + + All steps are validated before any is sent, so a typo in step 5 doesn't leave + the first four half-executed. + """ + parsed = [ + _parse_step(s if isinstance(s, InputStep) else InputStep.from_dict(s)) + for s in steps + ] + total = 0 + with self._lock: + for p in parsed: + self._issue(p) + total += p.step.duration_ms() + return total + + def flush(self) -> None: + """Block until every queued command has executed.""" + self._backend.wait_for_all() + + def cancel_all_commands_blocking(self, timeout_ms: int = 500) -> bool: + """Cancel all queued commands, release every input and wait up to `timeout_ms` + for the device to confirm it is in the neutral state. Returns True if + confirmed. Thread-safe: may be called while another thread is sending inputs.""" + return self._backend.cancel_all_commands_blocking(int(timeout_ms)) + + # ---- context manager ---------------------------------------------------- + + def __enter__(self) -> SwitchController: + return self + + def __exit__(self, *exc: object) -> None: + self.close() + + def close(self) -> None: + """Release all inputs (waiting briefly for the device to confirm) and drop + the connection.""" + try: + if self._backend.is_ready(): + self._backend.cancel_all_commands_blocking(500) + finally: + self._backend = _ClosedBackend() # type: ignore[assignment] + + # ---- internals ---------------------------------------------------------- + + def _issue(self, p: _ParsedStep) -> None: + s = p.step + b = self._backend + if s.is_wait(): + self.wait(s.wait_ms) + return + only_buttons = p.pressed.has_buttons and not p.pressed.has_dpad \ + and s.left_stick is None and s.right_stick is None + only_dpad = p.pressed.has_dpad and not p.pressed.has_buttons \ + and s.left_stick is None and s.right_stick is None + only_left = not p.pressed.has_buttons and not p.pressed.has_dpad \ + and s.left_stick is not None and s.right_stick is None + only_right = not p.pressed.has_buttons and not p.pressed.has_dpad \ + and s.left_stick is None and s.right_stick is not None + cycle = s.hold_ms + s.release_ms + for _ in range(s.repeat): + # Single-kind inputs use the dedicated commands, which let the scheduler + # enforce per-button cooldowns exactly like the main program. + if only_buttons: + b.press_buttons(cycle, s.hold_ms, s.release_ms, p.pressed.bitfield) + elif only_dpad: + b.press_dpad(cycle, s.hold_ms, s.release_ms, p.pressed.dpad) + elif only_left: + b.move_left_joystick(cycle, s.hold_ms, s.release_ms, *p.left) + elif only_right: + b.move_right_joystick(cycle, s.hold_ms, s.release_ms, *p.right) + else: + b.set_state(s.hold_ms, p.pressed.bitfield, p.pressed.dpad, *p.left, *p.right) + self.wait(s.release_ms) + self.wait(s.wait_ms) + + +class _ClosedBackend: + def __getattr__(self, name: str) -> Any: + raise RuntimeError("This controller has been closed.") + + +def validate_steps(steps: Sequence[InputStep | Mapping[str, Any]]) -> list[InputStep]: + """Parse and validate steps without sending them. Raises ValueError on bad input.""" + ret = [] + for s in steps: + step = s if isinstance(s, InputStep) else InputStep.from_dict(s) + _parse_step(step) + ret.append(step) + return ret diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/fake.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/fake.py new file mode 100644 index 0000000000..f0c95e6473 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/fake.py @@ -0,0 +1,71 @@ +"""Fake controller backend for testing without hardware. + +`FakeController` records every command it receives. It doesn't need the compiled +`_pa_core` module. +""" + +from __future__ import annotations + +import threading +from typing import Any + + + +class FakeController: + """Implements the `_pa_core.Controller` interface and records calls in `log`.""" + + def __init__(self, port_name: str = "fake"): + self.port_name = port_name + self.log: list[tuple[str, tuple[Any, ...]]] = [] + self.cancel_count = 0 + self.confirm_release = True + self._lock = threading.Lock() + + def wait_for_ready(self, timeout_ms: int) -> bool: + return True + + def is_ready(self) -> bool: + return True + + def status_text(self) -> str: + return "Fake controller (no hardware)" + + def controller_name(self) -> str: + return "Fake Controller" + + def _record(self, name: str, *args: Any) -> None: + with self._lock: + self.log.append((name, args)) + + def wait_for_all(self) -> None: + pass + + def cancel_all_commands_blocking(self, timeout_ms: int) -> bool: + """Counts the call in `cancel_count`. Returns `confirm_release`, which tests can + set to False to simulate a device that doesn't confirm in time.""" + with self._lock: + self.cancel_count += 1 + return self.confirm_release + + def wait(self, duration_ms: int) -> None: + self._record("wait", duration_ms) + + def press_buttons(self, delay_ms, hold_ms, release_ms, buttons) -> None: + self._record("press_buttons", delay_ms, hold_ms, release_ms, buttons) + + def press_dpad(self, delay_ms, hold_ms, release_ms, position) -> None: + self._record("press_dpad", delay_ms, hold_ms, release_ms, position) + + def move_left_joystick(self, delay_ms, hold_ms, release_ms, x, y) -> None: + self._record("move_left_joystick", delay_ms, hold_ms, release_ms, x, y) + + def move_right_joystick(self, delay_ms, hold_ms, release_ms, x, y) -> None: + self._record("move_right_joystick", delay_ms, hold_ms, release_ms, x, y) + + def set_state(self, duration_ms, buttons, dpad, left_x, left_y, right_x, right_y) -> None: + self._record("set_state", duration_ms, buttons, dpad, left_x, left_y, right_x, right_y) + + def commands(self) -> list[str]: + """Names of recorded commands, excluding waits.""" + with self._lock: + return [name for name, _ in self.log if name != "wait"] diff --git a/SerialPrograms/Source/PythonBindings/tests/test_controller.py b/SerialPrograms/Source/PythonBindings/tests/test_controller.py new file mode 100644 index 0000000000..a20eb500cc --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/tests/test_controller.py @@ -0,0 +1,79 @@ +import pytest + +from pokemon_automation import InputStep, SwitchController +from pokemon_automation import buttons as btn +from pokemon_automation.fake import FakeController + + +@pytest.fixture +def fake(): + return FakeController() + + +def test_press_uses_button_command(fake): + sw = SwitchController(backend=fake) + sw.press("A", hold_ms=50, release_ms=30, repeat=2) + # delay = hold + release, so presses run back to back like pbf_press_button(). + assert fake.log == [("press_buttons", (80, 50, 30, btn.BUTTON_BITS["A"]))] * 2 + + +def test_dpad_and_sticks_use_dedicated_commands(fake): + sw = SwitchController(backend=fake) + sw.press("down") + sw.stick("left", "up", duration_ms=1000) + sw.stick("right", [0.5, 0.0], duration_ms=200) + assert fake.commands() == ["press_dpad", "move_left_joystick", "move_right_joystick"] + name, args = [c for c in fake.log if c[0] == "move_left_joystick"][0] + assert args == (1000, 1000, 0, 0.0, 1.0) + + +def test_mixed_inputs_use_full_state(fake): + sw = SwitchController(backend=fake) + sw.hold("B", 1500, left="up") + name, args = fake.log[0] + assert name == "set_state" + assert args == (1500, btn.BUTTON_BITS["B"], btn.DPAD_NONE, 0.0, 1.0, 0.0, 0.0) + + +def test_sequence_is_validated_before_sending(fake): + sw = SwitchController(backend=fake) + with pytest.raises(ValueError): + sw.run([{"buttons": "A"}, {"buttons": "NOT_A_BUTTON"}]) + assert fake.log == [] + with pytest.raises(ValueError, match="Unknown input step field"): + sw.run([{"button": "A"}]) + + +def test_sequence_duration(fake): + sw = SwitchController(backend=fake) + total = sw.run([ + {"buttons": "A", "hold_ms": 100, "release_ms": 50, "repeat": 3}, + {"wait_ms": 1000}, + InputStep(left_stick="up", hold_ms=500, release_ms=0), + ]) + assert total == 3 * 150 + 1000 + 500 + + +def test_closed_controller_raises(fake): + sw = SwitchController(backend=fake) + sw.close() + with pytest.raises(RuntimeError, match="closed"): + sw.press("A") + + +def test_step_dicts_match_mcp_schema(fake): + """The same step dicts work in Python and in the MCP `run_inputs` tool.""" + sw = SwitchController(backend=fake) + sw.run([{"left_stick": "up", "buttons": "B", "hold_ms": 2000, "release_ms": 0}, + {"right_stick": [0.5, 0], "hold_ms": 100}]) + assert fake.commands() == ["set_state", "move_right_joystick"] + + +def test_cancel_all_commands_blocking_reports_confirmation(fake): + sw = SwitchController(backend=fake) + assert sw.cancel_all_commands_blocking(100) is True + fake.confirm_release = False + assert sw.cancel_all_commands_blocking(100) is False + assert fake.cancel_count == 2 + sw.close() # close() also releases + assert fake.cancel_count == 3 From 87759d2ed9751add97f448c905b19bc9fc4c039e Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:13:17 -0700 Subject: [PATCH 06/16] Python: add video capture and OCR Pure Python, outside the C++ codebase: - video.py: continuous capture with opencv-python (only the latest frame is kept, so a frame taken after an input shows its effect), JPEG/PNG encoding, cropping with ImageFloatBox-style boxes, OCR with pytesseract (optional). - devices.py: list serial ports and video devices (device names on macOS via pyobjc AVFoundation, in OpenCV's index order). - FakeVideoCapture for tests. Co-Authored-By: Claude Opus 5.5 --- .../pokemon_automation/__init__.py | 12 +- .../pokemon_automation/devices.py | 81 ++++ .../PythonBindings/pokemon_automation/fake.py | 70 ++- .../pokemon_automation/video.py | 399 ++++++++++++++++++ .../Source/PythonBindings/pyproject.toml | 12 +- .../PythonBindings/tests/test_vision.py | 67 +++ 6 files changed, 634 insertions(+), 7 deletions(-) create mode 100644 SerialPrograms/Source/PythonBindings/pokemon_automation/devices.py create mode 100644 SerialPrograms/Source/PythonBindings/pokemon_automation/video.py create mode 100644 SerialPrograms/Source/PythonBindings/tests/test_vision.py diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py index 177e5db034..20989e0298 100644 --- a/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py @@ -5,15 +5,25 @@ - `_pa_core` (C++, pybind11): acts as a Switch controller through a PABotBase2 serial device, built only from this codebase's CoreLib. Build it from SerialPrograms with `-DPA_PYTHON_BINDINGS=ON`. -- `SwitchController`: the Python API for the controller. +- Video and OCR are pure Python: opencv-python for capture, pytesseract (optional) + for OCR. +- `SwitchController`, `VideoSource`: the Python API. """ from .buttons import parse_buttons, parse_stick from .controller import InputStep, SwitchController +from .devices import list_serial_ports, list_video_devices +from .video import Frame, VideoSource, encode_image, ocr_image __all__ = [ + "Frame", "InputStep", "SwitchController", + "VideoSource", + "encode_image", + "list_serial_ports", + "list_video_devices", + "ocr_image", "parse_buttons", "parse_stick", ] diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/devices.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/devices.py new file mode 100644 index 0000000000..77983ea18c --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/devices.py @@ -0,0 +1,81 @@ +"""Discover serial ports (controller devices) and video capture devices.""" + +from __future__ import annotations + +import glob +import json +import subprocess +import sys +from pathlib import Path + +# Serial devices that are never a PABotBase2 controller. +_IGNORED_PORT_WORDS = ("bluetooth", "debug-console", "wlan", "headphones", "airpods") + + +def list_serial_ports() -> list[str]: + """Return candidate serial ports for the controller device, most likely first. + + Uses pyserial if it is installed (required on Windows); otherwise scans /dev. + On macOS only the "cu." (call-out) nodes are listed, which is what should be + opened; the matching "tty." nodes block on open. + """ + ports: list[str] = [] + try: + from serial.tools import list_ports # type: ignore[import-not-found] + ports = [p.device for p in list_ports.comports()] + except ImportError: + if sys.platform == "darwin": + ports = glob.glob("/dev/cu.*") + elif sys.platform.startswith("linux"): + ports = glob.glob("/dev/ttyUSB*") + glob.glob("/dev/ttyACM*") + if sys.platform == "darwin": + ports = [p for p in ports if not p.startswith("/dev/tty.")] + ports = [p for p in ports if not any(w in p.lower() for w in _IGNORED_PORT_WORDS)] + # USB-serial adapters first. + ports.sort(key=lambda p: (not any(w in p.lower() for w in ("usb", "uart", "acm", "com")), p)) + return ports + + +def list_video_devices() -> list[dict[str, int | str]]: + """Return [{"index": int, "name": str}] for every video capture device, where + `index` is the OpenCV device index to open. + + - macOS: AVFoundation via pyobjc (`pip install pyobjc-framework-AVFoundation`), + in the same order OpenCV uses. Without pyobjc, names come from + `system_profiler`, whose order usually but not always matches OpenCV's. + - Linux: /sys/class/video4linux. + - Windows: device names aren't available without extra libraries; returns [] + (open devices by index). + """ + if sys.platform == "darwin": + return _list_video_devices_macos() + if sys.platform.startswith("linux"): + ret = [] + for node in sorted(Path("/sys/class/video4linux").glob("video*")): + try: + index = int(node.name[5:]) + name = (node / "name").read_text().strip() + except (ValueError, OSError): + continue + ret.append({"index": index, "name": name or f"Camera {index}"}) + return sorted(ret, key=lambda d: d["index"]) + return [] + + +def _list_video_devices_macos() -> list[dict[str, int | str]]: + try: + import AVFoundation # type: ignore[import-not-found] + + # OpenCV's AVFoundation backend indexes video devices followed by muxed ones. + devices = list(AVFoundation.AVCaptureDevice.devicesWithMediaType_(AVFoundation.AVMediaTypeVideo)) + devices += list(AVFoundation.AVCaptureDevice.devicesWithMediaType_(AVFoundation.AVMediaTypeMuxed)) + return [{"index": i, "name": str(d.localizedName())} for i, d in enumerate(devices)] + except ImportError: + pass + try: + out = subprocess.run(["system_profiler", "SPCameraDataType", "-json"], + capture_output=True, text=True, timeout=10).stdout + cameras = json.loads(out).get("SPCameraDataType", []) + except (OSError, ValueError, subprocess.TimeoutExpired): + return [] + return [{"index": i, "name": c.get("_name", f"Camera {i}")} for i, c in enumerate(cameras)] diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/fake.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/fake.py index f0c95e6473..c811d2951b 100644 --- a/SerialPrograms/Source/PythonBindings/pokemon_automation/fake.py +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/fake.py @@ -1,14 +1,23 @@ -"""Fake controller backend for testing without hardware. +"""Fake controller and video backends for testing without hardware. -`FakeController` records every command it receives. It doesn't need the compiled -`_pa_core` module. +`FakeController` records every command it receives. `FakeVideoCapture` produces +synthetic frames whose color changes with each controller command, so tests (and an +agent trying the MCP server with `--fake`) can see that inputs "did something". + +Neither needs the compiled `_pa_core` module. Frames are encoded with opencv-python +when it is installed; without it they are always encoded as PNG. """ from __future__ import annotations +import struct import threading +import time +import zlib from typing import Any +import numpy as np + class FakeController: @@ -69,3 +78,58 @@ def commands(self) -> list[str]: """Names of recorded commands, excluding waits.""" with self._lock: return [name for name, _ in self.log if name != "wait"] + + +def encode_png(image: np.ndarray) -> bytes: + """Minimal pure-Python PNG encoder for RGB uint8 images.""" + height, width, _ = image.shape + raw = b"".join(b"\x00" + image[row].tobytes() for row in range(height)) + + def chunk(kind: bytes, data: bytes) -> bytes: + return struct.pack(">I", len(data)) + kind + data + struct.pack(">I", zlib.crc32(kind + data)) + + return (b"\x89PNG\r\n\x1a\n" + + chunk(b"IHDR", struct.pack(">IIBBBBB", width, height, 8, 2, 0, 0, 0)) + + chunk(b"IDAT", zlib.compress(raw, 3)) + + chunk(b"IEND", b"")) + + +class FakeVideoCapture: + """Implements the `_pa_core.VideoCapture` interface with synthetic frames.""" + + def __init__(self, controller: FakeController | None = None, width: int = 320, height: int = 180): + self.device_index = -1 + self.width = width + self.height = height + self._controller = controller + self._sequence = 0 + + def measured_fps(self) -> float: + return 30.0 + + def is_streaming(self) -> bool: + return True + + def _render(self) -> np.ndarray: + count = len(self._controller.commands()) if self._controller else 0 + image = np.zeros((self.height, self.width, 3), dtype=np.uint8) + image[:, :, 0] = np.linspace(0, 255, self.width, dtype=np.uint8)[None, :] + image[:, :, 1] = (count * 40) % 256 + image[:, :, 2] = 128 + return image + + def snapshot(self, min_sequence: int = 0, min_timestamp_ms: int = 0, timeout_ms: int = 2000): + self._sequence += 1 + now_ms = max(int(time.time() * 1000), min_timestamp_ms) + return self._render(), now_ms, self._sequence + + def encode_latest(self, format: str = "jpg", box=None, max_width: int = 0, quality: int = 85) -> bytes: + image = self._render() + try: + from .video import encode_image + return encode_image(image, format, box, max_width, quality) + except ImportError: + return encode_png(image) + + def close(self) -> None: + pass diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/video.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/video.py new file mode 100644 index 0000000000..4a2dea90ff --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/video.py @@ -0,0 +1,399 @@ +"""Video capture, image encoding and OCR, in pure Python. + +This part of the package deliberately does not use the C++ codebase: capture and +encoding use opencv-python (`cv2`), and OCR uses pytesseract (optional; it needs the +`tesseract` program installed). Only controller input goes through `_pa_core`. + +Frames are numpy arrays of shape (height, width, 3), dtype uint8, in RGB order. +Boxes are (x, y, width, height) as fractions of the frame, the same convention as +`ImageFloatBox` in the main C++ program, so boxes can be copied between the two. + +Example: + video = VideoSource("MiraBox") # by name substring, or by index + frame = video.frame() + frame.save("screen.png") + print(frame.ocr(box=(0.05, 0.75, 0.9, 0.2))) +""" + +from __future__ import annotations + +import sys +import threading +import time +from collections import deque +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Protocol + +import numpy as np + +Box = tuple[float, float, float, float] + +# Tesseract page segmentation modes by friendly name. +OCR_MODES = {"block": 6, "line": 7, "word": 8, "sparse": 11} + + +def _cv2() -> Any: + try: + import cv2 + except ImportError as e: + raise ImportError("Video features need opencv-python: pip install opencv-python") from e + return cv2 + + +class VideoBackend(Protocol): + """What `VideoSource` needs from a capture backend. Implemented by `CaptureThread` + and by `fake.FakeVideoCapture`.""" + + device_index: int + width: int + height: int + + def measured_fps(self) -> float: ... + def is_streaming(self) -> bool: ... + def snapshot(self, min_sequence: int = 0, min_timestamp_ms: int = 0, + timeout_ms: int = 2000) -> tuple[np.ndarray | None, int, int]: ... + def encode_latest(self, format: str = "jpg", box: Box | None = None, + max_width: int = 0, quality: int = 85) -> bytes: ... + def close(self) -> None: ... + + +def check_box(box: Box | None) -> Box | None: + """Validate a normalized box. Raises ValueError if it is malformed.""" + if box is None: + return None + if len(box) != 4: + raise ValueError("A box is [x, y, width, height] as fractions of the frame.") + x, y, w, h = (float(v) for v in box) + if not (0 <= x < 1 and 0 <= y < 1 and 0 < w <= 1 and 0 < h <= 1): + raise ValueError(f"Box values must be fractions of the frame in [0, 1], got {box!r}.") + return (x, y, w, h) + + +def crop(image: np.ndarray, box: Box | None) -> np.ndarray: + """Crop `image` to a normalized box (a view, no copy). None means the whole image.""" + box = check_box(box) + if box is None: + return image + height, width = image.shape[:2] + x, y, w, h = box + x0, y0 = min(int(x * width + 0.5), width - 1), min(int(y * height + 0.5), height - 1) + x1 = min(max(x0 + 1, int((x + w) * width + 0.5)), width) + y1 = min(max(y0 + 1, int((y + h) * height + 0.5)), height) + return image[y0:y1, x0:x1] + + +def _encode_bgr(bgr: np.ndarray, format: str, box: Box | None, max_width: int, quality: int) -> bytes: + cv2 = _cv2() + image = crop(bgr, box) + if max_width > 0 and image.shape[1] > max_width: + height = max(1, image.shape[0] * max_width // image.shape[1]) + image = cv2.resize(image, (max_width, height), interpolation=cv2.INTER_AREA) + if format == "png": + ok, data = cv2.imencode(".png", image, [cv2.IMWRITE_PNG_COMPRESSION, 1]) + elif format in ("jpg", "jpeg"): + ok, data = cv2.imencode(".jpg", image, [cv2.IMWRITE_JPEG_QUALITY, max(0, min(100, quality))]) + else: + raise ValueError(f'Unknown image format {format!r}; use "png" or "jpg".') + if not ok: + raise RuntimeError("Image encoding failed.") + return data.tobytes() + + +def encode_image(image: np.ndarray, format: str = "png", box: Box | None = None, + max_width: int = 0, quality: int = 85) -> bytes: + """Encode an RGB image as PNG or JPEG bytes, optionally cropped and downscaled.""" + if image.ndim != 3 or image.shape[2] != 3: + raise ValueError("Expected an RGB image: a uint8 array of shape (height, width, 3).") + return _encode_bgr(np.ascontiguousarray(image[:, :, ::-1]), format, box, max_width, quality) + + +def ocr_available() -> bool: + """True if pytesseract and the tesseract program are installed.""" + try: + import pytesseract + pytesseract.get_tesseract_version() + return True + except Exception: + return False + + +def ocr_image(image: np.ndarray, box: Box | None = None, language: str = "eng", + mode: str | int = "block", whitelist: str = "") -> str: + """Read text in `box` of an RGB image with Tesseract (via pytesseract). + + Steps: crop to `box`, convert to grayscale, upscale small crops so lines of text + are ~40+ px tall (the size Tesseract's LSTM model works best at), then run + Tesseract. `mode` is "block" (a paragraph), "line" (one line), "word", "sparse" + (scattered text, e.g. a whole menu screen) or a raw page segmentation mode number. + `language` is a Tesseract code such as "eng", "jpn" or "eng+jpn". + + Raises ImportError if pytesseract isn't installed, or pytesseract's + TesseractNotFoundError if the tesseract program isn't. + """ + cv2 = _cv2() + try: + import pytesseract + except ImportError as e: + raise ImportError("OCR needs pytesseract and Tesseract: pip install pytesseract, " + "and install tesseract (e.g. brew install tesseract).") from e + if isinstance(mode, str): + if mode not in OCR_MODES: + raise ValueError(f"mode must be one of {', '.join(OCR_MODES)} or an integer.") + psm = OCR_MODES[mode] + else: + psm = int(mode) + gray = cv2.cvtColor(np.ascontiguousarray(crop(image, box)), cv2.COLOR_RGB2GRAY) + if gray.shape[0] < 80: + scale = min(4.0, 80.0 / max(1, gray.shape[0])) + gray = cv2.resize(gray, None, fx=scale, fy=scale, interpolation=cv2.INTER_CUBIC) + config = f"--psm {psm}" + if whitelist: + config += f" -c tessedit_char_whitelist={whitelist}" + return pytesseract.image_to_string(gray, lang=language, config=config).strip() + + +class CaptureThread: + """Continuously reads frames from one OpenCV video device and keeps only the latest. + + Grabbing continuously (instead of reading on demand) matters because OpenCV and the + OS buffer several frames: a frame read on demand after a pause can be a second or + more old, so "press a button, then look" would show the screen from before the press. + + Raises RuntimeError if the device can't be opened or delivers no frames. On macOS + that is also what happens when the app running Python has no camera permission + (System Settings > Privacy & Security > Camera). + """ + + def __init__(self, device_index: int, width: int = 1920, height: int = 1080): + cv2 = _cv2() + if sys.platform == "darwin": + api = cv2.CAP_AVFOUNDATION + elif sys.platform == "win32": + api = cv2.CAP_MSMF + elif sys.platform.startswith("linux"): + api = cv2.CAP_V4L2 + else: + api = cv2.CAP_ANY + self.device_index = device_index + self._capture = cv2.VideoCapture(device_index, api) + if not self._capture.isOpened(): + raise RuntimeError( + f"Unable to open video device {device_index}. Check that it exists and, on " + "macOS, that this app has camera permission " + "(System Settings > Privacy & Security > Camera).") + if sys.platform != "darwin": + # Most USB capture cards only reach 1080p at full frame rate with MJPG. + self._capture.set(cv2.CAP_PROP_FOURCC, cv2.VideoWriter_fourcc(*"MJPG")) + self._capture.set(cv2.CAP_PROP_FRAME_WIDTH, width) + self._capture.set(cv2.CAP_PROP_FRAME_HEIGHT, height) + self._capture.set(cv2.CAP_PROP_BUFFERSIZE, 1) + + self._cond = threading.Condition() + self._frame: np.ndarray | None = None # BGR + self._timestamp_ms = 0 + self._sequence = 0 + self._recent: deque[float] = deque(maxlen=120) + + # Read one frame synchronously so the resolution is known and a device that + # opens but never delivers frames is reported here. + first = None + for _ in range(50): + ok, first = self._capture.read() + if ok and first is not None: + break + time.sleep(0.02) + else: + self._capture.release() + raise RuntimeError(f"Video device {device_index} opened but delivered no frames.") + self.height, self.width = first.shape[:2] + self._store(first) + + self._stopping = threading.Event() + self._thread = threading.Thread(target=self._loop, name=f"video-{device_index}", daemon=True) + self._thread.start() + + def _store(self, bgr: np.ndarray) -> None: + now = time.time() + with self._cond: + self._frame = bgr + self._timestamp_ms = int(now * 1000) + self._sequence += 1 + self._recent.append(now) + self._cond.notify_all() + + def _loop(self) -> None: + # If the device stops delivering frames (e.g. unplugged), keep retrying; + # `is_streaming()` reports the stall. + while not self._stopping.is_set(): + ok, frame = self._capture.read() + if not ok or frame is None: + time.sleep(0.02) + continue + self._store(frame) + + def _latest(self, min_sequence: int, min_timestamp_ms: int, timeout_ms: int): + with self._cond: + self._cond.wait_for( + lambda: self._sequence >= min_sequence and self._timestamp_ms >= min_timestamp_ms, + timeout=timeout_ms / 1000) + return self._frame, self._timestamp_ms, self._sequence + + def measured_fps(self) -> float: + with self._cond: + recent = [t for t in self._recent if t > time.time() - 3] + if len(recent) < 2 or recent[-1] <= recent[0]: + return 0.0 + return (len(recent) - 1) / (recent[-1] - recent[0]) + + def is_streaming(self) -> bool: + with self._cond: + return self._frame is not None and time.time() * 1000 - self._timestamp_ms < 2000 + + def snapshot(self, min_sequence: int = 0, min_timestamp_ms: int = 0, timeout_ms: int = 2000): + frame, timestamp, sequence = self._latest(min_sequence, min_timestamp_ms, timeout_ms) + if frame is None: + return None, 0, 0 + return np.ascontiguousarray(frame[:, :, ::-1]), timestamp, sequence + + def encode_latest(self, format: str = "jpg", box: Box | None = None, + max_width: int = 0, quality: int = 85) -> bytes: + frame, _, _ = self._latest(0, 0, 2000) + if frame is None: + return b"" + return _encode_bgr(frame, format, box, max_width, quality) + + def close(self) -> None: + self._stopping.set() + self._thread.join(timeout=2) + self._capture.release() + + +@dataclass +class Frame: + """One captured frame.""" + + image: np.ndarray # (height, width, 3) uint8 RGB + timestamp_ms: int # capture time, milliseconds since the Unix epoch + sequence: int # frame counter since the device was opened + + @property + def width(self) -> int: + return int(self.image.shape[1]) + + @property + def height(self) -> int: + return int(self.image.shape[0]) + + def crop(self, box: Box) -> np.ndarray: + return crop(self.image, box) + + def encode(self, format: str = "png", box: Box | None = None, + max_width: int = 0, quality: int = 85) -> bytes: + return encode_image(self.image, format, box, max_width, quality) + + def save(self, path: str | Path, box: Box | None = None, max_width: int = 0) -> Path: + """Save as PNG, or JPEG if `path` ends in .jpg/.jpeg.""" + path = Path(path) + fmt = "jpg" if path.suffix.lower() in (".jpg", ".jpeg") else "png" + path.write_bytes(self.encode(fmt, box, max_width, 92)) + return path + + def ocr(self, box: Box | None = None, **kwargs: Any) -> str: + """Read text in `box`. See `ocr_image()` for the options.""" + return ocr_image(self.image, box, **kwargs) + + +def find_video_device(device: int | str) -> int: + """Resolve a device index or a case-insensitive name substring to an index. + + Raises ValueError if a name matches no device or more than one. + """ + from .devices import list_video_devices + + if isinstance(device, int): + return device + text = str(device).strip() + if text.lstrip("-").isdigit(): + return int(text) + devices = list_video_devices() + matches = [d for d in devices if text.lower() in str(d["name"]).lower()] + if len(matches) == 1: + return int(matches[0]["index"]) + names = ", ".join(f'{d["index"]}: {d["name"]}' for d in devices) or "none found" + if not matches: + raise ValueError(f"No video device matches {text!r}. Devices: {names}") + raise ValueError(f"{text!r} matches several video devices; use an index. Devices: {names}") + + +class VideoSource: + """The latest frames from a capture card. + + `device` is an index or a name substring (see `devices.list_video_devices()`). + Pass `backend` to use an existing backend, e.g. `fake.FakeVideoCapture`. + + Raises RuntimeError if the device can't be opened (see `CaptureThread`). + """ + + def __init__(self, device: int | str | None = None, width: int = 1920, height: int = 1080, + *, backend: VideoBackend | None = None): + if backend is None: + if device is None: + raise ValueError("Either device or backend is required.") + backend = CaptureThread(find_video_device(device), width, height) + self._backend: VideoBackend | None = backend + + @property + def backend(self) -> VideoBackend: + if self._backend is None: + raise RuntimeError("This video source has been closed.") + return self._backend + + @property + def resolution(self) -> tuple[int, int]: + return (self.backend.width, self.backend.height) + + def fps(self) -> float: + return self.backend.measured_fps() + + def is_streaming(self) -> bool: + return self.backend.is_streaming() + + def frame(self, *, after_ms: int | None = None, timeout_ms: int = 2000) -> Frame: + """Return the newest frame. If `after_ms` (epoch milliseconds) is given, wait + up to `timeout_ms` for a frame captured at or after that time. + + Raises RuntimeError if no frame has ever been captured. + """ + image, timestamp, sequence = self.backend.snapshot(0, after_ms or 0, timeout_ms) + if image is None: + raise RuntimeError("No video frame has been captured yet.") + return Frame(image, timestamp, sequence) + + def fresh_frame(self, settle_ms: int = 0, timeout_ms: int = 2000) -> Frame: + """Wait `settle_ms`, then return a frame captured after the wait. + + Use this after sending inputs, so the frame shows their effect rather than a + frame that was already buffered. + """ + if settle_ms > 0: + time.sleep(settle_ms / 1000) + return self.frame(after_ms=int(time.time() * 1000), timeout_ms=timeout_ms) + + def jpeg(self, box: Box | None = None, max_width: int = 1280, quality: int = 80) -> bytes: + """The newest frame as JPEG bytes, optionally cropped and downscaled.""" + data = self.backend.encode_latest("jpg", check_box(box), max_width, quality) + if not data: + raise RuntimeError("No video frame has been captured yet.") + return data + + def close(self) -> None: + if self._backend is not None: + self._backend.close() + self._backend = None + + def __enter__(self) -> VideoSource: + return self + + def __exit__(self, *exc: object) -> None: + self.close() diff --git a/SerialPrograms/Source/PythonBindings/pyproject.toml b/SerialPrograms/Source/PythonBindings/pyproject.toml index f38b27930b..5b326c8d16 100644 --- a/SerialPrograms/Source/PythonBindings/pyproject.toml +++ b/SerialPrograms/Source/PythonBindings/pyproject.toml @@ -4,18 +4,24 @@ build-backend = "setuptools.build_meta" # The compiled `_pa_core` module (acts as a Switch controller) is built by CMake (see # README.md) and copied into pokemon_automation/. This file packages the Python side -# for `pip install -e .`. +# for `pip install -e .`. Video and OCR use the third-party packages below. [project] name = "pokemon-automation" version = "0.1.0" description = "Control a Nintendo Switch from Python with Pokemon Automation hardware." requires-python = ">=3.10" -dependencies = [] +dependencies = [ + "numpy>=1.24", + "opencv-python>=4.8", + # Video device names on macOS, in OpenCV's index order. + "pyobjc-framework-AVFoundation>=10; sys_platform == 'darwin'", +] [project.optional-dependencies] +ocr = ["pytesseract>=0.3.10"] # also needs the tesseract program installed serial = ["pyserial>=3.5"] -test = ["pytest>=7"] +test = ["pytest>=7", "pytesseract>=0.3.10", "pillow"] [tool.setuptools] packages = ["pokemon_automation"] diff --git a/SerialPrograms/Source/PythonBindings/tests/test_vision.py b/SerialPrograms/Source/PythonBindings/tests/test_vision.py new file mode 100644 index 0000000000..3b5d302185 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/tests/test_vision.py @@ -0,0 +1,67 @@ +"""Tests of the pure-Python vision helpers (opencv-python / pytesseract), no hardware.""" + +import io + +import numpy as np +import pytest + +pytest.importorskip("cv2") + +from pokemon_automation import Frame, encode_image, ocr_image # noqa: E402 +from pokemon_automation.video import ocr_available # noqa: E402 + + +def make_image(width=64, height=48): + image = np.zeros((height, width, 3), dtype=np.uint8) + image[:, : width // 2] = (255, 0, 0) # left half red + image[:, width // 2:] = (0, 0, 255) # right half blue + return image + + +def test_encode_png_and_jpeg(): + image = make_image() + assert encode_image(image, "png").startswith(b"\x89PNG") + assert encode_image(image, "jpg").startswith(b"\xff\xd8") + + +def test_encode_rejects_bad_input(): + with pytest.raises(ValueError): + encode_image(np.zeros((10, 10), dtype=np.uint8)) + with pytest.raises(ValueError, match="Unknown image format"): + encode_image(make_image(), "gif") + with pytest.raises(ValueError, match="fractions"): + encode_image(make_image(), box=(0, 0, 2, 1)) + + +def test_crop_and_scale(): + image = make_image() + frame = Frame(image, 0, 1) + assert (frame.crop((0.5, 0.0, 0.5, 1.0)) == (0, 0, 255)).all() + # The PNG header (IHDR) holds the output width and height. + small = encode_image(image, "png", box=(0.5, 0.0, 0.5, 1.0), max_width=16) + assert (int.from_bytes(small[16:20], "big"), int.from_bytes(small[20:24], "big")) == (16, 24) + + +def test_encoded_colors_are_rgb(): + """Round-trip through the encoder must not swap red and blue.""" + pil = pytest.importorskip("PIL.Image") + decoded = np.asarray(pil.open(io.BytesIO(encode_image(make_image(), "png"))).convert("RGB")) + assert tuple(decoded[0, 0]) == (255, 0, 0) + assert tuple(decoded[0, -1]) == (0, 0, 255) + + +def test_ocr_reads_rendered_text(): + if not ocr_available(): + pytest.skip("pytesseract or tesseract not installed") + pil = pytest.importorskip("PIL.Image") + from PIL import ImageDraw, ImageFont + + canvas = pil.new("RGB", (640, 120), (20, 20, 20)) + draw = ImageDraw.Draw(canvas) + try: + font = ImageFont.truetype("Arial.ttf", 48) + except OSError: + font = ImageFont.load_default(size=48) + draw.text((20, 30), "Hello Switch 123", fill=(250, 250, 250), font=font) + text = ocr_image(np.asarray(canvas), mode="line") + assert "Switch" in text and "123" in text From 55cf2e2eef1f5b7ed7c08a1101edc01b602c8191 Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:13:17 -0700 Subject: [PATCH 07/16] Python: add Console and the hardware self-test - Console: a controller plus a video source. act() sends inputs, waits until the device has executed them, waits for the game to react, then returns a frame captured after all of that. - selftest.py: `python -m pokemon_automation.selftest --serial ... --video ...` checks real hardware with harmless inputs and saves screenshots. Co-Authored-By: Claude Opus 5.5 --- .../pokemon_automation/__init__.py | 11 +- .../pokemon_automation/console.py | 127 ++++++++++++++++++ .../pokemon_automation/selftest.py | 66 +++++++++ .../PythonBindings/tests/test_controller.py | 17 ++- 4 files changed, 218 insertions(+), 3 deletions(-) create mode 100644 SerialPrograms/Source/PythonBindings/pokemon_automation/console.py create mode 100644 SerialPrograms/Source/PythonBindings/pokemon_automation/selftest.py diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py index 20989e0298..4a84b32954 100644 --- a/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py @@ -7,15 +7,24 @@ `-DPA_PYTHON_BINDINGS=ON`. - Video and OCR are pure Python: opencv-python for capture, pytesseract (optional) for OCR. -- `SwitchController`, `VideoSource`: the Python API. +- `SwitchController`, `VideoSource`, `Console`: the Python API. + +Quick start: + from pokemon_automation import Console, list_serial_ports, list_video_devices + + print(list_serial_ports(), list_video_devices()) + with Console(serial_port="/dev/cu.usbserial-0001", video="MiraBox") as console: + console.act([{"buttons": "A"}]).save("after_A.png") """ from .buttons import parse_buttons, parse_stick +from .console import Console from .controller import InputStep, SwitchController from .devices import list_serial_ports, list_video_devices from .video import Frame, VideoSource, encode_image, ocr_image __all__ = [ + "Console", "Frame", "InputStep", "SwitchController", diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/console.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/console.py new file mode 100644 index 0000000000..ed6533da3c --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/console.py @@ -0,0 +1,127 @@ +"""`Console`: a controller and a video feed for one Switch, used together. + +This is the object both the Python API and the MCP server are built around. The core +operation is `act()`: send inputs, wait until the device has executed them, let the +game react, then capture a frame that is guaranteed to be newer than the inputs. + +Example: + from pokemon_automation import Console + + with Console(serial_port="/dev/cu.usbserial-0001", video="MiraBox") as console: + frame = console.act([{"buttons": "HOME"}], settle_ms=1000) + frame.save("home.png") + print(console.read_text(box=(0.0, 0.9, 1.0, 0.1))) +""" + +from __future__ import annotations + +import time +from collections.abc import Iterable, Mapping +from typing import Any + +from .controller import InputStep, SwitchController +from .video import Box, Frame, VideoSource + + +class Console: + """A Switch driven by a controller device and observed through a capture card. + + Either part is optional: without `serial_port`/`controller` the console is + read-only, and without `video`/`video_source` it can't observe. + """ + + def __init__(self, serial_port: str | None = None, video: int | str | None = None, *, + width: int = 1920, height: int = 1080, + controller: SwitchController | None = None, + video_source: VideoSource | None = None, + controller_timeout_s: float = 10.0): + self.controller = controller + self.video = video_source + try: + if self.video is None and video is not None: + self.video = VideoSource(video, width, height) + if self.controller is None and serial_port is not None: + self.controller = SwitchController(serial_port, timeout_s=controller_timeout_s) + except BaseException: + self.close() + raise + + # ---- requirements ------------------------------------------------------- + + def require_controller(self) -> SwitchController: + if self.controller is None: + raise RuntimeError("No controller is connected.") + return self.controller + + def require_video(self) -> VideoSource: + if self.video is None: + raise RuntimeError("No video device is connected.") + return self.video + + # ---- actions -------------------------------------------------------------- + + def send(self, steps: Iterable[InputStep | Mapping[str, Any]], *, wait: bool = True) -> int: + """Queue input steps; if `wait`, block until they've executed. + Returns the total input duration in milliseconds.""" + controller = self.require_controller() + total = controller.run(steps) + if wait: + controller.flush() + return total + + def act(self, steps: Iterable[InputStep | Mapping[str, Any]], *, settle_ms: int = 500) -> Frame: + """Send inputs, wait for them to finish, wait `settle_ms` more for the game to + react, then return a frame captured after all of that.""" + self.send(steps, wait=True) + return self.observe(settle_ms=settle_ms) + + def observe(self, *, settle_ms: int = 0) -> Frame: + """Return a frame captured at least `settle_ms` from now.""" + return self.require_video().fresh_frame(settle_ms) + + def screenshot_jpeg(self, box: Box | None = None, max_width: int = 1280, quality: int = 80, + *, settle_ms: int = 0) -> bytes: + """A fresh frame (captured after `settle_ms`) as JPEG bytes.""" + video = self.require_video() + if settle_ms > 0: + video.fresh_frame(settle_ms) + return video.jpeg(box, max_width, quality) + + def read_text(self, box: Box | None = None, **kwargs: Any) -> str: + """OCR the newest frame. See `video.ocr_image()` for the options.""" + return self.require_video().frame().ocr(box, **kwargs) + + def wait_until(self, predicate, timeout_s: float = 10.0, interval_ms: int = 200) -> Frame | None: + """Poll frames until `predicate(frame)` is truthy. Returns the matching frame, + or None on timeout. Useful for "wait for this menu to appear" in scripts.""" + deadline = time.monotonic() + timeout_s + while True: + frame = self.require_video().frame() + if predicate(frame): + return frame + if time.monotonic() >= deadline: + return None + time.sleep(interval_ms / 1000) + + def cancel_all_commands_blocking(self, timeout_ms: int = 500) -> bool: + """Emergency stop that waits up to `timeout_ms` for the device to confirm the + neutral state. Returns True if confirmed. Thread-safe.""" + if self.controller is None: + return False + return self.controller.cancel_all_commands_blocking(timeout_ms) + + # ---- lifetime ------------------------------------------------------------- + + def close(self) -> None: + if self.controller is not None: + self.controller.close() + self.controller = None + if self.video is not None: + self.video.close() + self.video = None + + def __enter__(self) -> Console: + return self + + def __exit__(self, *exc: object) -> None: + self.close() diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/selftest.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/selftest.py new file mode 100644 index 0000000000..e5e0952277 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/selftest.py @@ -0,0 +1,66 @@ +"""Hardware self-test: check the controller, the capture card, and that inputs reach +the Switch, saving screenshots along the way. + + python -m pokemon_automation.selftest --serial /dev/cu.usbserial-0001 --video MiraBox + +Put the Switch on the Home menu first. The test only presses HOME, the d-pad and B. +On macOS, run this from Terminal/iTerm the first time so macOS can ask for camera +permission for that app. +""" + +from __future__ import annotations + +import argparse +import sys +import time +from pathlib import Path + +from .console import Console +from .devices import list_serial_ports, list_video_devices +from .video import ocr_available + + +def main(argv: list[str] | None = None) -> int: + p = argparse.ArgumentParser(prog="python -m pokemon_automation.selftest", description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter) + p.add_argument("--serial", help="Controller serial port. Omit to skip controller tests.") + p.add_argument("--video", help="Video device index or name. Omit to skip video tests.") + p.add_argument("--out", default="selftest_output", help="Folder for screenshots.") + args = p.parse_args(argv) + + print("Serial ports: ", list_serial_ports()) + print("Video devices:", list_video_devices()) + print("OCR available:", ocr_available()) + if not args.serial and not args.video: + print("\nPass --serial and/or --video to test devices.") + return 0 + + out = Path(args.out) + out.mkdir(parents=True, exist_ok=True) + with Console(serial_port=args.serial, video=args.video) as console: + if console.video is not None: + frame = console.observe(settle_ms=500) + print(f"Video: {frame.width}x{frame.height} at {console.video.fps():.1f} fps, " + f"mean color {frame.image.mean(axis=(0, 1)).round(1).tolist()}") + print("Saved", frame.save(out / "0_start.png")) + if console.controller is not None: + print("Controller:", console.controller.name(), "|", console.controller.status()) + for i, (label, steps, settle) in enumerate([ + ("HOME", [{"buttons": "HOME", "hold_ms": 100}], 1500), + ("RIGHT x2", [{"buttons": "RIGHT", "repeat": 2, "release_ms": 300}], 500), + ("LEFT x2", [{"buttons": "LEFT", "repeat": 2, "release_ms": 300}], 500), + ], start=1): + started = time.time() + console.send(steps) + print(f"Sent {label} ({time.time() - started:.2f} s)") + if console.video is not None: + frame = console.observe(settle_ms=settle) + print("Saved", frame.save(out / f"{i}_{label.split()[0].lower()}.png")) + if console.video is not None and ocr_available(): + print("OCR of the whole screen (sparse):", + repr(console.read_text(mode="sparse")[:200])) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/SerialPrograms/Source/PythonBindings/tests/test_controller.py b/SerialPrograms/Source/PythonBindings/tests/test_controller.py index a20eb500cc..fb76f0a32f 100644 --- a/SerialPrograms/Source/PythonBindings/tests/test_controller.py +++ b/SerialPrograms/Source/PythonBindings/tests/test_controller.py @@ -1,8 +1,8 @@ import pytest -from pokemon_automation import InputStep, SwitchController +from pokemon_automation import Console, InputStep, SwitchController, VideoSource from pokemon_automation import buttons as btn -from pokemon_automation.fake import FakeController +from pokemon_automation.fake import FakeController, FakeVideoCapture @pytest.fixture @@ -54,6 +54,19 @@ def test_sequence_duration(fake): assert total == 3 * 150 + 1000 + 500 +def test_console_act_returns_new_frame(fake): + console = Console(controller=SwitchController(backend=fake), + video_source=VideoSource(backend=FakeVideoCapture(fake))) + before = console.observe() + after = console.act([{"buttons": "A"}], settle_ms=0) + assert after.sequence > before.sequence + assert after.image.shape == (180, 320, 3) + # The fake video changes color with every command, so the input "did something". + assert (after.image != before.image).any() + console.close() + assert console.controller is None and console.video is None + + def test_closed_controller_raises(fake): sw = SwitchController(backend=fake) sw.close() From b20facbe734d5e878d48cd12d02e8e75302d1022 Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:13:17 -0700 Subject: [PATCH 08/16] Add the AI agent tool definitions (AgentTools.json) The MCP interface for AI agents that control the Switch, as data, so the Python MCP server and the upcoming SerialPrograms app server expose exactly the same tools: names, descriptions, JSON Schemas of the arguments, which host implements each tool, and the instructions sent to agents. AgentInputTestCases.json lists inputs (buttons, sticks, steps) and their expected parse results or errors; both the Python and C++ parsers are tested against it. Also: agent_tools.py loads the file (inlining $refs), tests that the Python input vocabulary passes the shared cases, and the build copies both files into the Python package. Co-Authored-By: Claude Opus 5.5 --- .../AgentServer/AgentInputTestCases.json | 61 +++++ .../Integrations/AgentServer/AgentTools.json | 220 ++++++++++++++++++ .../Source/PythonBindings/.gitignore | 3 + .../PythonBindings/PythonBindings.cmake | 9 +- .../pokemon_automation/agent_tools.py | 89 +++++++ .../Source/PythonBindings/pyproject.toml | 2 +- .../tests/test_shared_interface.py | 59 +++++ 7 files changed, 441 insertions(+), 2 deletions(-) create mode 100644 SerialPrograms/Source/Integrations/AgentServer/AgentInputTestCases.json create mode 100644 SerialPrograms/Source/Integrations/AgentServer/AgentTools.json create mode 100644 SerialPrograms/Source/PythonBindings/pokemon_automation/agent_tools.py create mode 100644 SerialPrograms/Source/PythonBindings/tests/test_shared_interface.py diff --git a/SerialPrograms/Source/Integrations/AgentServer/AgentInputTestCases.json b/SerialPrograms/Source/Integrations/AgentServer/AgentInputTestCases.json new file mode 100644 index 0000000000..37069c1c31 --- /dev/null +++ b/SerialPrograms/Source/Integrations/AgentServer/AgentInputTestCases.json @@ -0,0 +1,61 @@ +{ + "_comment": [ + "Test cases for the input vocabulary of AgentTools.json, shared by the C++", + "(AgentServer_InputSteps.cpp) and Python (pokemon_automation/buttons.py)", + "implementations so both parse inputs identically.", + "Button bits are NintendoSwitch::Button values; d-pad positions are", + "NintendoSwitch::DpadPosition values (0 = up, clockwise, 8 = none).", + "Stick values are [x, y] with +y = up; diagonals are normalized to length 1." + ], + "buttons": [ + {"input": "A", "bitfield": 4, "dpad": 8}, + {"input": "a", "bitfield": 4, "dpad": 8}, + {"input": "L+R", "bitfield": 48, "dpad": 8}, + {"input": ["ZL", "a"], "bitfield": 68, "dpad": 8}, + {"input": "+", "bitfield": 512, "dpad": 8}, + {"input": "-", "bitfield": 256, "dpad": 8}, + {"input": "start", "bitfield": 512, "dpad": 8}, + {"input": "select", "bitfield": 256, "dpad": 8}, + {"input": "L3", "bitfield": 1024, "dpad": 8}, + {"input": "rs", "bitfield": 2048, "dpad": 8}, + {"input": "HOME", "bitfield": 4096, "dpad": 8}, + {"input": "capture", "bitfield": 8192, "dpad": 8}, + {"input": ["A", "+"], "bitfield": 516, "dpad": 8}, + {"input": "up", "bitfield": 0, "dpad": 0}, + {"input": "up+right", "bitfield": 0, "dpad": 1}, + {"input": "UP_RIGHT", "bitfield": 0, "dpad": 1}, + {"input": "down-left", "bitfield": 0, "dpad": 5}, + {"input": "UPLEFT", "bitfield": 0, "dpad": 7}, + {"input": "dpad_left", "bitfield": 0, "dpad": 6}, + {"input": "ZL+DOWN", "bitfield": 64, "dpad": 4}, + {"input": " b , x ", "bitfield": 10, "dpad": 8}, + {"input": "", "bitfield": 0, "dpad": 8}, + {"input": "Q", "error": "Unknown button"}, + {"input": "up+down", "error": "Contradictory"}, + {"input": "left+right", "error": "Contradictory"} + ], + "sticks": [ + {"input": "up", "x": 0.0, "y": 1.0}, + {"input": "DOWN", "x": 0.0, "y": -1.0}, + {"input": "left", "x": -1.0, "y": 0.0}, + {"input": "neutral", "x": 0.0, "y": 0.0}, + {"input": "center", "x": 0.0, "y": 0.0}, + {"input": "down_right", "x": 0.7071067811865475, "y": -0.7071067811865475}, + {"input": "up-left", "x": -0.7071067811865475, "y": 0.7071067811865475}, + {"input": [0.5, -0.25], "x": 0.5, "y": -0.25}, + {"input": [2, 0], "error": "within [-1, 1]"}, + {"input": [0.5], "error": "two numbers"}, + {"input": "sideways", "error": "Unknown stick direction"} + ], + "steps": [ + {"input": {"buttons": "A"}, "duration_ms": 160}, + {"input": {"buttons": "A", "hold_ms": 100, "release_ms": 50, "repeat": 3}, "duration_ms": 450}, + {"input": {"wait_ms": 1000}, "duration_ms": 1000}, + {"input": {"left_stick": "up", "hold_ms": 500, "release_ms": 0}, "duration_ms": 500}, + {"input": {"buttons": "B", "left_stick": [0.5, 1.0], "hold_ms": 1500, "wait_ms": 100}, "duration_ms": 1680}, + {"input": {"button": "A"}, "error": "Unknown input step field"}, + {"input": {"buttons": "A", "hold_ms": 0}, "error": "hold_ms"}, + {"input": {"buttons": "A", "repeat": 0}, "error": "repeat"}, + {"input": {"buttons": "A", "release_ms": -1}, "error": "negative"} + ] +} diff --git a/SerialPrograms/Source/Integrations/AgentServer/AgentTools.json b/SerialPrograms/Source/Integrations/AgentServer/AgentTools.json new file mode 100644 index 0000000000..d38bc54d53 --- /dev/null +++ b/SerialPrograms/Source/Integrations/AgentServer/AgentTools.json @@ -0,0 +1,220 @@ +{ + "_comment": [ + "The MCP interface for AI agents that control a Nintendo Switch.", + "Shared by both MCP servers so they expose exactly the same tools:", + " - the SerialPrograms app ('AI Agent Server' program in the ML tab): host 'app'", + " - the Python package (pokemon_automation.mcp_server): host 'python'", + "Each tool lists the hosts that implement it. 'inputSchema' is the JSON Schema", + "sent to agents in tools/list; both servers validate arguments against it.", + "Input steps (buttons, sticks, durations) follow the vocabulary in", + "AgentInputTestCases.json, which both implementations are tested against." + ], + "server_name": "pokemon-automation-switch", + "instructions": [ + "You control a real Nintendo Switch through a USB device that acts as a controller, and see its screen through a capture card.", + "", + "Workflow:", + "- Call `screenshot` first to see where you are.", + "- Input tools (`press_buttons`, `move_stick`, `run_inputs`) return a screenshot taken `settle_ms` after the inputs finished, unless observe=false. Prefer `run_inputs` to send several steps in one call when you are confident about them (e.g. navigating a known menu).", + "- Use `read_text` to read on-screen text reliably instead of reading it off a screenshot.", + "- If something goes wrong, call `cancel_all_commands_blocking`.", + "- If a tool says the user has taken control, stop sending inputs. The user is steering the console by hand; check `switch_status` later to see when control is returned.", + "", + "Conventions:", + "- Buttons: A B X Y L R ZL ZR PLUS MINUS HOME CAPTURE LCLICK RCLICK. D-pad: UP DOWN LEFT RIGHT (and UP_RIGHT etc.). Combine with \"+\", e.g. \"L+R\" or \"ZL+A\".", + "- Sticks: direction names (up, down_left, ...) or [x, y] in [-1, 1], +y = up.", + "- Boxes: [x, y, width, height] as fractions of the screen, (0, 0) = top-left.", + "- Menus usually need a short press (hold 80 ms) and ~300-800 ms to animate. Walking is done by holding a stick for a duration.", + "- Switch UI: A = confirm, B = back, HOME = Home menu, X on Home menu = close game.", + "- A black screen may mean the Switch is asleep; ask the user to wake it." + ], + "definitions": { + "box": { + "type": "array", + "items": {"type": "number", "minimum": 0, "maximum": 1}, + "minItems": 4, + "maxItems": 4, + "description": "Crop box [x, y, width, height] as fractions of the screen; omit for the full screen." + }, + "stick": { + "anyOf": [ + {"type": "string"}, + {"type": "array", "items": {"type": "number", "minimum": -1, "maximum": 1}, "minItems": 2, "maxItems": 2} + ] + }, + "settle_ms": { + "type": "integer", + "minimum": 0, + "maximum": 10000, + "description": "Wait this long after the inputs before the screenshot (default from server config)." + }, + "observe": { + "type": "boolean", + "default": true, + "description": "Return a screenshot afterwards." + } + }, + "tools": [ + { + "name": "switch_status", + "hosts": ["app", "python"], + "description": "Report controller and video connection status, resolution, who is in control, and limits.", + "inputSchema": {"type": "object", "properties": {}, "additionalProperties": false} + }, + { + "name": "list_devices", + "hosts": ["python"], + "description": "List serial ports (controller devices) and video capture devices.", + "inputSchema": {"type": "object", "properties": {}, "additionalProperties": false} + }, + { + "name": "connect", + "hosts": ["python"], + "description": "(Re)connect to the controller and/or video device. Omitted arguments keep the current setting. Use after unplugging a device or to switch devices.", + "inputSchema": { + "type": "object", + "properties": { + "serial_port": {"type": "string", "description": "Serial port of the controller device."}, + "video_device": {"type": "string", "description": "Video device index or name substring."} + }, + "additionalProperties": false + } + }, + { + "name": "screenshot", + "hosts": ["app", "python"], + "description": "Capture the current screen as a JPEG. Crop with `box` to zoom into details.", + "inputSchema": { + "type": "object", + "properties": { + "box": {"$ref": "#/definitions/box"}, + "max_width": {"type": "integer", "minimum": 64, "maximum": 3840, "description": "Downscale to this width."} + }, + "additionalProperties": false + } + }, + { + "name": "wait_and_observe", + "hosts": ["app", "python"], + "description": "Wait without pressing anything (e.g. for a loading screen), then screenshot.", + "inputSchema": { + "type": "object", + "properties": { + "duration_ms": {"type": "integer", "minimum": 0, "maximum": 60000, "description": "How long to wait."}, + "box": {"$ref": "#/definitions/box"} + }, + "required": ["duration_ms"], + "additionalProperties": false + } + }, + { + "name": "read_text", + "hosts": ["app", "python"], + "description": "Read on-screen text with OCR. Crop tightly around the text with `box` for best results.", + "inputSchema": { + "type": "object", + "properties": { + "box": {"$ref": "#/definitions/box"}, + "mode": { + "type": "string", + "enum": ["block", "line", "word", "sparse"], + "default": "block", + "description": "\"line\" for one line, \"block\" for a paragraph, \"sparse\" for scattered text." + }, + "language": {"type": "string", "default": "eng", "description": "Tesseract language code, e.g. \"eng\", \"jpn\"."} + }, + "additionalProperties": false + } + }, + { + "name": "press_buttons", + "hosts": ["app", "python"], + "description": "Press a button combination, optionally several times.", + "inputSchema": { + "type": "object", + "properties": { + "buttons": {"type": "string", "description": "Buttons/d-pad to press together, e.g. \"A\", \"L+R\", \"DOWN\"."}, + "hold_ms": {"type": "integer", "minimum": 1, "default": 80}, + "release_ms": {"type": "integer", "minimum": 0, "default": 120}, + "repeat": {"type": "integer", "minimum": 1, "maximum": 100, "default": 1, "description": "Press this many times."}, + "observe": {"$ref": "#/definitions/observe"}, + "settle_ms": {"$ref": "#/definitions/settle_ms"} + }, + "required": ["buttons"], + "additionalProperties": false + } + }, + { + "name": "move_stick", + "hosts": ["app", "python"], + "description": "Tilt a stick for a duration (walk, move a cursor, turn the camera), optionally while holding buttons.", + "inputSchema": { + "type": "object", + "properties": { + "direction": { + "$ref": "#/definitions/stick", + "description": "Direction name (\"up\", \"down_left\", ...) or [x, y] in [-1, 1], +y = up." + }, + "duration_ms": {"type": "integer", "minimum": 1, "default": 500, "description": "How long to hold the stick."}, + "stick": {"type": "string", "enum": ["left", "right"], "default": "left"}, + "buttons": {"type": "string", "description": "Buttons to hold at the same time, e.g. \"B\" to run."}, + "observe": {"$ref": "#/definitions/observe"}, + "settle_ms": {"$ref": "#/definitions/settle_ms"} + }, + "required": ["direction"], + "additionalProperties": false + } + }, + { + "name": "run_inputs", + "hosts": ["app", "python"], + "description": "Run a sequence of input steps back to back, then optionally screenshot. Example: [{\"buttons\": \"DOWN\", \"repeat\": 3}, {\"buttons\": \"A\"}, {\"wait_ms\": 1000}, {\"left_stick\": \"up\", \"hold_ms\": 2000, \"release_ms\": 0}]", + "inputSchema": { + "type": "object", + "properties": { + "steps": { + "type": "array", + "minItems": 1, + "maxItems": 200, + "items": { + "type": "object", + "description": "One input step. Set `buttons` and/or sticks to press something, or only `wait_ms` to pause.", + "properties": { + "buttons": {"type": "string", "description": "Buttons and/or d-pad to hold together, e.g. \"A\", \"L+R\", \"ZL+UP\"."}, + "left_stick": {"$ref": "#/definitions/stick", "description": "Left stick: direction name (\"up\", \"down_left\") or [x, y] in [-1, 1]."}, + "right_stick": {"$ref": "#/definitions/stick", "description": "Right stick (camera in most games), same format as left_stick."}, + "hold_ms": {"type": "integer", "minimum": 1, "default": 80, "description": "How long to hold the inputs."}, + "release_ms": {"type": "integer", "minimum": 0, "default": 80, "description": "Neutral time after releasing, before the next step."}, + "repeat": {"type": "integer", "minimum": 1, "maximum": 100, "default": 1, "description": "Repeat this press/release cycle this many times."}, + "wait_ms": {"type": "integer", "minimum": 0, "default": 0, "description": "Extra pause after the step (or the whole step if nothing is pressed)."} + }, + "additionalProperties": false + } + }, + "observe": {"$ref": "#/definitions/observe"}, + "settle_ms": {"$ref": "#/definitions/settle_ms"} + }, + "required": ["steps"], + "additionalProperties": false + } + }, + { + "name": "cancel_all_commands_blocking", + "hosts": ["app", "python"], + "description": "Emergency stop: cancel queued inputs and release every button and stick, then wait briefly for the device to confirm.", + "inputSchema": {"type": "object", "properties": {}, "additionalProperties": false} + }, + { + "name": "get_logs", + "hosts": ["app", "python"], + "description": "Recent log lines (connection status, inputs, errors).", + "inputSchema": { + "type": "object", + "properties": { + "count": {"type": "integer", "minimum": 1, "maximum": 500, "default": 30} + }, + "additionalProperties": false + } + } + ] +} diff --git a/SerialPrograms/Source/PythonBindings/.gitignore b/SerialPrograms/Source/PythonBindings/.gitignore index 53101df5e9..a6ed7df9c4 100644 --- a/SerialPrograms/Source/PythonBindings/.gitignore +++ b/SerialPrograms/Source/PythonBindings/.gitignore @@ -1,6 +1,9 @@ # Built by CMake and copied here after each build. pokemon_automation/_pa_core*.so pokemon_automation/_pa_core*.pyd +# Copied from Source/Integrations/AgentServer/ by the build (see agent_tools.py). +pokemon_automation/AgentTools.json +pokemon_automation/AgentInputTestCases.json __pycache__/ *.egg-info/ .pytest_cache/ diff --git a/SerialPrograms/Source/PythonBindings/PythonBindings.cmake b/SerialPrograms/Source/PythonBindings/PythonBindings.cmake index b90613ac56..c3f4657167 100644 --- a/SerialPrograms/Source/PythonBindings/PythonBindings.cmake +++ b/SerialPrograms/Source/PythonBindings/PythonBindings.cmake @@ -53,8 +53,15 @@ pa_apply_gui_free_target_properties(_pa_core) target_link_libraries(_pa_core PRIVATE CoreLib) set(PA_PYTHON_PACKAGE_DIR ${CMAKE_CURRENT_SOURCE_DIR}/Source/PythonBindings/pokemon_automation) +# Also copy the shared MCP interface (AgentTools.json, also compiled into the app) +# and its test cases, so an installed package doesn't need the source tree. +set(PA_AGENT_SERVER_DIR ${CMAKE_CURRENT_SOURCE_DIR}/Source/Integrations/AgentServer) add_custom_command( TARGET _pa_core POST_BUILD COMMAND ${CMAKE_COMMAND} -E copy $ ${PA_PYTHON_PACKAGE_DIR}/ - COMMENT "Copying _pa_core into the pokemon_automation Python package" + COMMAND ${CMAKE_COMMAND} -E copy + ${PA_AGENT_SERVER_DIR}/AgentTools.json + ${PA_AGENT_SERVER_DIR}/AgentInputTestCases.json + ${PA_PYTHON_PACKAGE_DIR}/ + COMMENT "Copying _pa_core and the shared agent tool definitions into the pokemon_automation Python package" ) diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/agent_tools.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/agent_tools.py new file mode 100644 index 0000000000..c4e44615fc --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/agent_tools.py @@ -0,0 +1,89 @@ +"""The shared MCP interface definition (AgentTools.json). + +The SerialPrograms app and this package both serve MCP from the same definition file, +so agents see identical tools whichever server they connect to. The file lives in the +C++ source tree at SerialPrograms/Source/Integrations/AgentServer/AgentTools.json. + +Search order: +1. `PA_AGENT_TOOLS` environment variable (path to the file). +2. The C++ source tree, when running from a source checkout, so edits to the file + take effect without rebuilding. +3. A copy inside this package (made by the CMake build, and included in wheels). +""" + +from __future__ import annotations + +import copy +import json +import os +from functools import lru_cache +from pathlib import Path +from typing import Any + +FILE_NAME = "AgentTools.json" +TEST_CASES_FILE_NAME = "AgentInputTestCases.json" + +_PACKAGE_DIR = Path(__file__).resolve().parent +_SOURCE_TREE_DIR = _PACKAGE_DIR.parent.parent / "Integrations" / "AgentServer" + + +def find_file(name: str = FILE_NAME) -> Path: + """Locate a shared definition file. Raises FileNotFoundError if it's nowhere.""" + candidates = [] + if name == FILE_NAME and os.environ.get("PA_AGENT_TOOLS"): + candidates.append(Path(os.environ["PA_AGENT_TOOLS"])) + candidates += [_SOURCE_TREE_DIR / name, _PACKAGE_DIR / name] + for path in candidates: + if path.is_file(): + return path + raise FileNotFoundError( + f"{name} not found. Looked in: " + ", ".join(str(p) for p in candidates)) + + +def resolve_refs(schema: Any, definitions: dict[str, Any]) -> Any: + """Return `schema` with every {"$ref": "#/definitions/", ...} replaced by a + copy of that definition, merged with the sibling keys (siblings win). + + Agents receive each tool's inputSchema on its own, without the file's shared + `definitions`, so references must be inlined. Raises KeyError on unknown names. + """ + if isinstance(schema, list): + return [resolve_refs(item, definitions) for item in schema] + if not isinstance(schema, dict): + return schema + if "$ref" in schema: + ref = schema["$ref"] + prefix = "#/definitions/" + if not ref.startswith(prefix): + raise KeyError(f"Unsupported $ref {ref!r}") + merged = copy.deepcopy(definitions[ref[len(prefix):]]) + merged.update({k: v for k, v in schema.items() if k != "$ref"}) + return resolve_refs(merged, definitions) + return {k: resolve_refs(v, definitions) for k, v in schema.items()} + + +@lru_cache(maxsize=None) +def load() -> dict[str, Any]: + """Load AgentTools.json with every tool's inputSchema made self-contained.""" + data = json.loads(find_file().read_text(encoding="utf-8")) + definitions = data.get("definitions", {}) + for tool in data["tools"]: + tool["inputSchema"] = resolve_refs(tool["inputSchema"], definitions) + return data + + +def instructions() -> str: + return "\n".join(load()["instructions"]) + + +def server_name() -> str: + return load()["server_name"] + + +def tools_for(host: str) -> dict[str, dict[str, Any]]: + """The tools a host ("app" or "python") implements, by name.""" + return {t["name"]: t for t in load()["tools"] if host in t["hosts"]} + + +def load_test_cases() -> dict[str, Any]: + return json.loads(find_file(TEST_CASES_FILE_NAME).read_text(encoding="utf-8")) diff --git a/SerialPrograms/Source/PythonBindings/pyproject.toml b/SerialPrograms/Source/PythonBindings/pyproject.toml index 5b326c8d16..c6a6d9057f 100644 --- a/SerialPrograms/Source/PythonBindings/pyproject.toml +++ b/SerialPrograms/Source/PythonBindings/pyproject.toml @@ -27,7 +27,7 @@ test = ["pytest>=7", "pytesseract>=0.3.10", "pillow"] packages = ["pokemon_automation"] [tool.setuptools.package-data] -pokemon_automation = ["_pa_core*.so", "_pa_core*.pyd"] +pokemon_automation = ["_pa_core*.so", "_pa_core*.pyd", "AgentTools.json", "AgentInputTestCases.json"] [tool.pytest.ini_options] testpaths = ["tests"] diff --git a/SerialPrograms/Source/PythonBindings/tests/test_shared_interface.py b/SerialPrograms/Source/PythonBindings/tests/test_shared_interface.py new file mode 100644 index 0000000000..649e9f47b1 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/tests/test_shared_interface.py @@ -0,0 +1,59 @@ +"""AgentTools.json must be usable as-is, and the Python input vocabulary must pass +the shared AgentInputTestCases.json (also used by C++).""" + +import json +import math + +import pytest + +from pokemon_automation import agent_tools +from pokemon_automation import buttons as btn +from pokemon_automation.controller import InputStep, validate_steps + + +def test_schemas_are_self_contained(): + text = json.dumps([t["inputSchema"] for t in agent_tools.load()["tools"]]) + assert "$ref" not in text + + +def test_every_tool_names_known_hosts(): + for tool in agent_tools.load()["tools"]: + assert tool["hosts"] and set(tool["hosts"]) <= {"app", "python"}, tool["name"] + + +CASES = agent_tools.load_test_cases() + + +def expect_error(fn, arg, substring): + """Shared cases name an expected substring of the error message (not a regex).""" + with pytest.raises(ValueError) as e: + fn(arg) + assert substring in str(e.value) + + +@pytest.mark.parametrize("case", CASES["buttons"], ids=lambda c: repr(c["input"])) +def test_shared_button_cases(case): + if "error" in case: + expect_error(btn.parse_buttons, case["input"], case["error"]) + return + parsed = btn.parse_buttons(case["input"]) + assert (parsed.bitfield, parsed.dpad) == (case["bitfield"], case["dpad"]) + + +@pytest.mark.parametrize("case", CASES["sticks"], ids=lambda c: repr(c["input"])) +def test_shared_stick_cases(case): + if "error" in case: + expect_error(btn.parse_stick, case["input"], case["error"]) + return + x, y = btn.parse_stick(case["input"]) + assert math.isclose(x, case["x"], abs_tol=1e-9) and math.isclose(y, case["y"], abs_tol=1e-9) + + +@pytest.mark.parametrize("case", CASES["steps"], ids=lambda c: json.dumps(c["input"])) +def test_shared_step_cases(case): + if "error" in case: + expect_error(validate_steps, [case["input"]], case["error"]) + return + (step,) = validate_steps([case["input"]]) + assert isinstance(step, InputStep) + assert step.duration_ms() == case["duration_ms"] From cedfb1864df0acde1dd59a1d49d54e6a27a49bb1 Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:13:18 -0700 Subject: [PATCH 09/16] Python: add the MCP server pokemon_automation.mcp_server serves a Console to AI agents over MCP (stdio or streamable HTTP). Tools, descriptions and schemas come from AgentTools.json; tests check that the Python functions accept exactly the shared arguments. Input tools return a screenshot taken after the game reacts; cancel_all_commands_blocking works while another call is running; per-call input limits; read-only mode; --fake runs without hardware. Co-Authored-By: Claude Opus 5.5 --- .../pokemon_automation/__init__.py | 2 + .../pokemon_automation/mcp_server.py | 475 ++++++++++++++++++ .../Source/PythonBindings/pyproject.toml | 6 +- .../PythonBindings/tests/test_mcp_server.py | 124 +++++ .../tests/test_shared_interface.py | 39 +- 5 files changed, 643 insertions(+), 3 deletions(-) create mode 100644 SerialPrograms/Source/PythonBindings/pokemon_automation/mcp_server.py create mode 100644 SerialPrograms/Source/PythonBindings/tests/test_mcp_server.py diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py index 4a84b32954..9e9712cd28 100644 --- a/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py @@ -8,6 +8,8 @@ - Video and OCR are pure Python: opencv-python for capture, pytesseract (optional) for OCR. - `SwitchController`, `VideoSource`, `Console`: the Python API. +- `pokemon_automation.mcp_server`: an MCP server exposing a `Console` to AI agents. + Run it with `python -m pokemon_automation.mcp_server --help`. Quick start: from pokemon_automation import Console, list_serial_ports, list_video_devices diff --git a/SerialPrograms/Source/PythonBindings/pokemon_automation/mcp_server.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/mcp_server.py new file mode 100644 index 0000000000..1aee3e845b --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/mcp_server.py @@ -0,0 +1,475 @@ +"""MCP server that lets AI agents see and control a Nintendo Switch. + +The server wraps one `Console` (controller + capture card). Tools are deliberately +coarse: a single call can send a whole input sequence and return a screenshot taken +after the game has reacted, because an agent needs seconds per turn and every round +trip counts. + +Run over stdio (for Claude Code, Claude Desktop and most MCP clients): + python -m pokemon_automation.mcp_server --serial /dev/cu.usbserial-0001 --video MiraBox + +Run over HTTP (for remote agents): + python -m pokemon_automation.mcp_server --serial ... --video ... \\ + --transport streamable-http --host 127.0.0.1 --port 8765 + +Try it without hardware: + python -m pokemon_automation.mcp_server --fake + +Register with Claude Code: + claude mcp add switch -- /path/to/python -m pokemon_automation.mcp_server \\ + --serial /dev/cu.usbserial-0001 --video MiraBox + +Safety: every input tool is capped (`--max-hold-ms`, `--max-sequence-ms`), +`cancel_all_commands_blocking` works even while another tool is mid-sequence, and +`--read-only` disables inputs entirely. All inputs are logged (see `get_logs`). +""" + +from __future__ import annotations + +import argparse +import functools +import logging +import os +import threading +import time +from dataclasses import dataclass +from typing import Annotated, Any, Literal + +import anyio +from pydantic import BaseModel, Field + +try: # mcp >= 2 + from mcp.server.mcpserver import Image, MCPServer as _Server + from mcp.server.mcpserver.exceptions import ToolError +except ImportError: # mcp 1.x + from mcp.server.fastmcp import FastMCP as _Server, Image # type: ignore[no-redef] + from mcp.server.fastmcp.exceptions import ToolError # type: ignore[no-redef] + +from . import agent_tools +from . import buttons as btn +from ._core import core, core_available +from .console import Console +from .controller import InputStep, SwitchController, validate_steps +from .devices import list_serial_ports, list_video_devices +from .fake import FakeController, FakeVideoCapture +from .video import VideoSource, find_video_device, ocr_available + +log = logging.getLogger("pokemon_automation.mcp") + +# Tool names, descriptions, input schemas and agent instructions come from the +# shared AgentTools.json (see agent_tools.py), which the SerialPrograms app serves +# too. The functions below implement the tools; their signatures do the argument +# validation and must stay in line with the shared schemas (tests check this). + +BoxArg = list[float] | None + + +class Step(BaseModel): + """One `run_inputs` step (schema and descriptions: AgentTools.json).""" + + model_config = {"extra": "forbid"} + + buttons: str | None = None + left_stick: str | list[float] | None = None + right_stick: str | list[float] | None = None + hold_ms: int = Field(80, ge=1) + release_ms: int = Field(80, ge=0) + repeat: int = Field(1, ge=1, le=100) + wait_ms: int = Field(0, ge=0) + + def to_input_step(self) -> InputStep: + return InputStep( + buttons=self.buttons, left_stick=self.left_stick, right_stick=self.right_stick, + hold_ms=self.hold_ms, release_ms=self.release_ms, repeat=self.repeat, + wait_ms=self.wait_ms, + ) + + +@dataclass +class ServerConfig: + serial_port: str | None = None + video_device: str | None = None + width: int = 1920 + height: int = 1080 + fake: bool = False + read_only: bool = False + max_hold_ms: int = 10_000 + max_sequence_ms: int = 60_000 + screenshot_width: int = 1280 + jpeg_quality: int = 75 + default_settle_ms: int = 500 + + +class ConsoleManager: + """Owns the `Console` and connects to devices lazily, so the server starts (and + can report useful errors) even when a device is missing or unplugged.""" + + def __init__(self, config: ServerConfig): + self.config = config + self._lock = threading.Lock() + self._console: Console | None = None + self._errors: dict[str, str] = {} + + def console(self) -> Console: + with self._lock: + if self._console is None: + self._console = self._open(self.config.serial_port, self.config.video_device) + return self._console + + def reconnect(self, serial_port: str | None, video_device: str | None) -> Console: + with self._lock: + if self._console is not None: + self._console.close() + self._console = None + if serial_port is not None: + self.config.serial_port = serial_port + if video_device is not None: + self.config.video_device = video_device + self._console = self._open(self.config.serial_port, self.config.video_device) + return self._console + + def current(self) -> Console | None: + return self._console + + def errors(self) -> dict[str, str]: + return dict(self._errors) + + def _open(self, serial_port: str | None, video_device: str | None) -> Console: + self._errors = {} + if self.config.fake: + fake_controller = FakeController() + return Console( + controller=SwitchController(backend=fake_controller), + video_source=VideoSource(backend=FakeVideoCapture(fake_controller)), + ) + controller = video = None + if serial_port and not self.config.read_only: + try: + controller = SwitchController(serial_port) + except Exception as e: # keep going so video still works + self._errors["controller"] = str(e) + log.warning("Controller: %s", e) + if video_device is not None: + try: + video = VideoSource(video_device, self.config.width, self.config.height) + except Exception as e: + self._errors["video"] = str(e) + log.warning("Video: %s", e) + return Console(controller=controller, video_source=video) + + def close(self) -> None: + with self._lock: + if self._console is not None: + self._console.close() + self._console = None + + +def log_event(message: str) -> None: + """Log to the core log (which `get_logs` returns and echoes to stderr), or to + Python logging when the core module isn't available.""" + if core_available(): + core().log("[MCP] " + message) + else: + log.info(message) + + +def _image_from_bytes(data: bytes) -> Image: + fmt = "png" if data.startswith(b"\x89PNG") else "jpeg" + return Image(data=data, format=fmt) + + +def agent_errors(fn): + """Report expected failures (bad arguments, missing devices, limits) to the agent + as a readable tool error. The MCP SDK replaces the message of any other exception + with a generic "Error executing tool", which leaves the agent guessing.""" + @functools.wraps(fn) + async def wrapper(*args: Any, **kwargs: Any) -> Any: + try: + return await fn(*args, **kwargs) + except (ValueError, RuntimeError, ImportError, OSError) as e: + raise ToolError(str(e)) from e + return wrapper + + +def create_server(config: ServerConfig) -> tuple[Any, ConsoleManager]: + """Build the MCP server and its console manager. Separate from `main()` for tests.""" + manager = ConsoleManager(config) + server = _Server(agent_tools.server_name(), instructions=agent_tools.instructions()) + shared_tools = agent_tools.tools_for("python") + + def tool(structured_output: bool | None = None): + """Register a tool under its function name, with the description and input + schema from AgentTools.json. Raises KeyError for a tool not in the file.""" + def decorator(fn): + definition = shared_tools[fn.__name__] + server.tool(description=definition["description"], + structured_output=structured_output)(fn) + # Advertise the shared schema (the SDK would otherwise derive one from + # the signature, which validates the same arguments but reads differently). + server._tool_manager.get_tool(fn.__name__).parameters = definition["inputSchema"] + return fn + return decorator + # Serializes input tools so two calls can't interleave their button presses. + input_lock = anyio.Lock() + + def check_limits(steps: list[InputStep]) -> int: + total = 0 + for s in steps: + if not s.is_wait() and s.hold_ms > config.max_hold_ms: + raise ValueError(f"hold_ms {s.hold_ms} exceeds the limit of {config.max_hold_ms} ms.") + total += s.duration_ms() + if total > config.max_sequence_ms: + raise ValueError( + f"The sequence lasts {total} ms, over the limit of {config.max_sequence_ms} ms. " + "Split it into several calls.") + return total + + def screenshot_content(console: Console, box: list[float] | None, max_width: int | None, + settle_ms: int) -> list[Any]: + video = console.require_video() + started = time.monotonic() + frame = video.fresh_frame(settle_ms) + data = video.jpeg(tuple(box) if box else None, + max_width or config.screenshot_width, config.jpeg_quality) + waited = int((time.monotonic() - started) * 1000) + return [ + f"Frame #{frame.sequence} ({frame.width}x{frame.height}), captured {waited} ms after the call.", + _image_from_bytes(data), + ] + + async def send_and_observe(steps: list[InputStep], observe: bool, settle_ms: int | None, + description: str) -> list[Any]: + if config.read_only: + raise ValueError("The server is in read-only mode; inputs are disabled.") + validate_steps(steps) + total = check_limits(steps) + settle = config.default_settle_ms if settle_ms is None else settle_ms + async with input_lock: + def work() -> list[Any]: + console = manager.console() + console.require_controller() + log_event(f"Input: {description}") + console.send(steps, wait=True) + result: list[Any] = [f"Done: {description} ({total} ms of input)."] + if observe and console.video is not None: + result += screenshot_content(console, None, None, settle) + return result + return await anyio.to_thread.run_sync(work) + + # ---- status & setup ------------------------------------------------------ + + @tool() + @agent_errors + async def switch_status() -> dict[str, Any]: + def work() -> dict[str, Any]: + console = manager.console() + status: dict[str, Any] = {"control": "agent", "read_only": config.read_only, + "fake": config.fake} + if console.controller is not None: + status["controller"] = { + "ready": console.controller.is_ready(), + "name": console.controller.name(), + "status": console.controller.status(), + "port": config.serial_port, + } + else: + status["controller"] = {"ready": False, "error": manager.errors().get( + "controller", "No serial port configured.")} + if console.video is not None: + width, height = console.video.resolution + status["video"] = { + "streaming": console.video.is_streaming(), + "resolution": [width, height], + "fps": round(console.video.fps(), 1), + "device": config.video_device, + } + else: + status["video"] = {"streaming": False, "error": manager.errors().get( + "video", "No video device configured.")} + status["limits"] = {"max_hold_ms": config.max_hold_ms, + "max_sequence_ms": config.max_sequence_ms} + status["ocr_available"] = ocr_available() + return status + return await anyio.to_thread.run_sync(work) + + @tool() + @agent_errors + async def list_devices() -> dict[str, Any]: + return {"serial_ports": list_serial_ports(), "video_devices": list_video_devices()} + + @tool() + @agent_errors + async def connect( + serial_port: str | None = None, + video_device: str | None = None, + ) -> dict[str, Any]: + def work() -> dict[str, Any]: + if video_device is not None and not config.fake: + find_video_device(video_device) # fail fast on a bad name + manager.reconnect(serial_port, video_device) + return {"errors": manager.errors() or None} + async with input_lock: + return await anyio.to_thread.run_sync(work) + + # ---- observation ------------------------------------------------------------ + + @tool(structured_output=False) + @agent_errors + async def screenshot( + box: BoxArg = None, + max_width: Annotated[int | None, Field(ge=64, le=3840)] = None, + ) -> list[Any]: + def work() -> list[Any]: + return screenshot_content(manager.console(), box, max_width, 0) + return await anyio.to_thread.run_sync(work) + + @tool(structured_output=False) + @agent_errors + async def wait_and_observe( + duration_ms: Annotated[int, Field(ge=0, le=60_000)], + box: BoxArg = None, + ) -> list[Any]: + def work() -> list[Any]: + return screenshot_content(manager.console(), box, None, duration_ms) + return await anyio.to_thread.run_sync(work) + + @tool() + @agent_errors + async def read_text( + box: BoxArg = None, + mode: Literal["block", "line", "word", "sparse"] = "block", + language: str = "eng", + ) -> str: + def work() -> str: + return manager.console().read_text( + tuple(box) if box else None, language=language, mode=mode) + return await anyio.to_thread.run_sync(work) + + # ---- input -------------------------------------------------------------------- + + @tool(structured_output=False) + @agent_errors + async def press_buttons( + buttons: str, + hold_ms: Annotated[int, Field(ge=1)] = 80, + release_ms: Annotated[int, Field(ge=0)] = 120, + repeat: Annotated[int, Field(ge=1, le=100)] = 1, + observe: bool = True, + settle_ms: Annotated[int | None, Field(ge=0, le=10_000)] = None, + ) -> list[Any]: + step = InputStep(buttons=buttons, hold_ms=hold_ms, release_ms=release_ms, repeat=repeat) + return await send_and_observe([step], observe, settle_ms, + f"press {buttons}" + (f" x{repeat}" if repeat > 1 else "")) + + @tool(structured_output=False) + @agent_errors + async def move_stick( + direction: str | list[float], + duration_ms: Annotated[int, Field(ge=1)] = 500, + stick: Literal["left", "right"] = "left", + buttons: str | None = None, + observe: bool = True, + settle_ms: Annotated[int | None, Field(ge=0, le=10_000)] = None, + ) -> list[Any]: + btn.parse_stick(direction) # validate early for a clear error + step = InputStep(buttons=buttons, hold_ms=duration_ms, release_ms=0) + setattr(step, stick + "_stick", direction) + return await send_and_observe([step], observe, settle_ms, + f"{stick} stick {direction} for {duration_ms} ms" + + (f" holding {buttons}" if buttons else "")) + + @tool(structured_output=False) + @agent_errors + async def run_inputs( + steps: Annotated[list[Step], Field(min_length=1, max_length=200)], + observe: bool = True, + settle_ms: Annotated[int | None, Field(ge=0, le=10_000)] = None, + ) -> list[Any]: + input_steps = [s.to_input_step() for s in steps] + return await send_and_observe(input_steps, observe, settle_ms, f"{len(steps)} steps") + + @tool() + @agent_errors + async def cancel_all_commands_blocking() -> str: + console = manager.current() + if console is None or console.controller is None: + return "No controller connected." + timeout_ms = 1000 + confirmed = await anyio.to_thread.run_sync(console.cancel_all_commands_blocking, timeout_ms) + log_event(f"cancel_all_commands_blocking called (confirmed: {confirmed})") + if confirmed: + return "All inputs released; the device confirmed the neutral state." + return (f"Release requested, but the device did not confirm within {timeout_ms} ms. " + "Check switch_status; the controller may be disconnected.") + + @tool() + @agent_errors + async def get_logs(count: Annotated[int, Field(ge=1, le=500)] = 30) -> list[str]: + if not core_available(): + return [] + return list(core().recent_logs(count)) + + registered = {t.name for t in server._tool_manager.list_tools()} + if registered != set(shared_tools): + raise RuntimeError( + f"Tools out of sync with AgentTools.json: missing {sorted(set(shared_tools) - registered)}, " + f"extra {sorted(registered - set(shared_tools))}") + return server, manager + + +def parse_args(argv: list[str] | None = None) -> argparse.Namespace: + p = argparse.ArgumentParser( + prog="python -m pokemon_automation.mcp_server", + description="MCP server exposing a Nintendo Switch (controller + capture card) to AI agents.") + p.add_argument("--serial", default=os.environ.get("PA_SERIAL_PORT"), + help="Controller serial port, e.g. /dev/cu.usbserial-0001 (env PA_SERIAL_PORT).") + p.add_argument("--video", default=os.environ.get("PA_VIDEO_DEVICE"), + help="Video device index or name substring (env PA_VIDEO_DEVICE).") + p.add_argument("--width", type=int, default=1920) + p.add_argument("--height", type=int, default=1080) + p.add_argument("--fake", action="store_true", help="Use fake devices (no hardware).") + p.add_argument("--read-only", action="store_true", help="Disable all controller input.") + p.add_argument("--max-hold-ms", type=int, default=10_000) + p.add_argument("--max-sequence-ms", type=int, default=60_000) + p.add_argument("--screenshot-width", type=int, default=1280) + p.add_argument("--jpeg-quality", type=int, default=75) + p.add_argument("--settle-ms", type=int, default=500, + help="Default wait between inputs finishing and the screenshot.") + p.add_argument("--transport", choices=["stdio", "streamable-http"], default="stdio") + p.add_argument("--host", default="127.0.0.1") + p.add_argument("--port", type=int, default=8765) + p.add_argument("--log-file", default=os.environ.get("PA_LOG_FILE"), + help="Append internal logs to this file.") + return p.parse_args(argv) + + +def main(argv: list[str] | None = None) -> None: + args = parse_args(argv) + # Logs go to stderr: over stdio, stdout is the protocol channel. + logging.basicConfig(level=logging.INFO, format="%(asctime)s %(name)s: %(message)s") + if core_available(): + core().set_log_stderr(True) + if args.log_file: + core().set_log_file(args.log_file) + config = ServerConfig( + serial_port=args.serial, video_device=args.video, width=args.width, height=args.height, + fake=args.fake, read_only=args.read_only, max_hold_ms=args.max_hold_ms, + max_sequence_ms=args.max_sequence_ms, screenshot_width=args.screenshot_width, + jpeg_quality=args.jpeg_quality, default_settle_ms=args.settle_ms, + ) + server, manager = create_server(config) + try: + if args.transport == "stdio": + server.run("stdio") + else: + server.run("streamable-http", host=args.host, port=args.port) + except KeyboardInterrupt: + # Ctrl+C is the normal way to stop the server; don't print a traceback. + log.info("Stopping the MCP server.") + finally: + # Releases all inputs and closes the devices. + manager.close() + + +if __name__ == "__main__": + main() diff --git a/SerialPrograms/Source/PythonBindings/pyproject.toml b/SerialPrograms/Source/PythonBindings/pyproject.toml index c6a6d9057f..a419cb7527 100644 --- a/SerialPrograms/Source/PythonBindings/pyproject.toml +++ b/SerialPrograms/Source/PythonBindings/pyproject.toml @@ -19,9 +19,13 @@ dependencies = [ ] [project.optional-dependencies] +mcp = ["mcp>=1.2"] ocr = ["pytesseract>=0.3.10"] # also needs the tesseract program installed serial = ["pyserial>=3.5"] -test = ["pytest>=7", "pytesseract>=0.3.10", "pillow"] +test = ["pytest>=7", "mcp>=1.2", "pytesseract>=0.3.10", "pillow"] + +[project.scripts] +pa-switch-mcp = "pokemon_automation.mcp_server:main" [tool.setuptools] packages = ["pokemon_automation"] diff --git a/SerialPrograms/Source/PythonBindings/tests/test_mcp_server.py b/SerialPrograms/Source/PythonBindings/tests/test_mcp_server.py new file mode 100644 index 0000000000..655e8143dd --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/tests/test_mcp_server.py @@ -0,0 +1,124 @@ +"""End-to-end tests of the MCP server with fake devices, through a real MCP client.""" + +import anyio +import pytest + +pytest.importorskip("mcp") +from mcp import Client # noqa: E402 + +from pokemon_automation import buttons as btn # noqa: E402 +from pokemon_automation.mcp_server import ServerConfig, create_server # noqa: E402 + + +def run(coro_fn): + return anyio.run(coro_fn) + + +def make(**overrides): + config = ServerConfig(fake=True, default_settle_ms=0, **overrides) + server, manager = create_server(config) + return server, manager + + +def texts(result): + return [c.text for c in result.content if c.type == "text"] + + +def images(result): + return [c for c in result.content if c.type == "image"] + + +def test_lists_expected_tools(): + server, _ = make() + + async def go(): + async with Client(server) as client: + return {t.name for t in (await client.list_tools()).tools} + + names = run(go) + assert {"screenshot", "press_buttons", "move_stick", "run_inputs", "read_text", + "cancel_all_commands_blocking", "switch_status", "wait_and_observe", "connect"} <= names + + +def test_press_returns_screenshot_and_sends_input(): + server, manager = make() + + async def go(): + async with Client(server) as client: + return await client.call_tool("press_buttons", {"buttons": "A", "repeat": 2}) + + result = run(go) + assert not result.is_error, texts(result) + assert "press A x2" in texts(result)[0] + assert len(images(result)) == 1 + fake = manager.current().controller.backend + assert [c for c in fake.log if c[0] == "press_buttons"][0][1][3] == btn.BUTTON_BITS["A"] + + +def test_run_inputs_and_limits(): + server, manager = make(max_sequence_ms=5000) + + async def go(): + async with Client(server) as client: + ok = await client.call_tool("run_inputs", { + "steps": [ + {"buttons": "DOWN", "repeat": 3}, + {"buttons": "A"}, + {"left_stick": "up", "buttons": "B", "hold_ms": 1000, "release_ms": 0}, + ], + "observe": False, + }) + too_long = await client.call_tool("run_inputs", { + "steps": [{"left_stick": "up", "hold_ms": 6000}], "observe": False}) + bad_button = await client.call_tool("press_buttons", {"buttons": "Q"}) + return ok, too_long, bad_button + + ok, too_long, bad_button = run(go) + assert not ok.is_error, texts(ok) + assert images(ok) == [] + fake = manager.current().controller.backend + assert fake.commands() == ["press_dpad"] * 3 + ["press_buttons", "set_state"] + assert too_long.is_error and "limit" in texts(too_long)[0] + assert bad_button.is_error and "Unknown button" in texts(bad_button)[0] + + +def test_read_only_mode_blocks_inputs(): + server, _ = make(read_only=True) + + async def go(): + async with Client(server) as client: + return await client.call_tool("press_buttons", {"buttons": "A"}) + + result = run(go) + assert result.is_error and "read-only" in texts(result)[0] + + +def test_status_screenshot_and_release(): + server, manager = make() + + async def go(): + async with Client(server) as client: + status = await client.call_tool("switch_status", {}) + shot = await client.call_tool("screenshot", {"box": [0, 0, 0.5, 0.5]}) + released = await client.call_tool("cancel_all_commands_blocking", {}) + return status, shot, released + + status, shot, released = run(go) + assert '"ready": true' in texts(status)[0] + assert len(images(shot)) == 1 + assert "released" in texts(released)[0] + assert manager.current().controller.backend.cancel_count == 1 + + +def test_cancel_all_commands_blocking_unconfirmed_is_reported(): + server, manager = make() + + async def go(): + async with Client(server) as client: + await client.call_tool("switch_status", {}) # connects the fake devices + manager.current().controller.backend.confirm_release = False + return await client.call_tool("cancel_all_commands_blocking", {}) + + result = run(go) + assert not result.is_error + assert "did not confirm" in texts(result)[0] diff --git a/SerialPrograms/Source/PythonBindings/tests/test_shared_interface.py b/SerialPrograms/Source/PythonBindings/tests/test_shared_interface.py index 649e9f47b1..3d03b991b2 100644 --- a/SerialPrograms/Source/PythonBindings/tests/test_shared_interface.py +++ b/SerialPrograms/Source/PythonBindings/tests/test_shared_interface.py @@ -1,9 +1,10 @@ -"""AgentTools.json must be usable as-is, and the Python input vocabulary must pass -the shared AgentInputTestCases.json (also used by C++).""" +"""The Python server must match the shared interface in AgentTools.json, and the +input vocabulary must pass the shared AgentInputTestCases.json (also used by C++).""" import json import math +import anyio import pytest from pokemon_automation import agent_tools @@ -21,6 +22,40 @@ def test_every_tool_names_known_hosts(): assert tool["hosts"] and set(tool["hosts"]) <= {"app", "python"}, tool["name"] +def test_python_server_serves_the_shared_definitions(): + pytest.importorskip("mcp") + from mcp import Client + from pokemon_automation.mcp_server import ServerConfig, create_server + + server, _ = create_server(ServerConfig(fake=True)) + + async def go(): + async with Client(server) as client: + return (await client.list_tools()).tools, client.instructions + + tools, instructions = anyio.run(go) + shared = agent_tools.tools_for("python") + assert {t.name for t in tools} == set(shared) + for t in tools: + assert t.description == shared[t.name]["description"] + assert t.input_schema == shared[t.name]["inputSchema"] + assert instructions == agent_tools.instructions() + + +def test_python_signatures_accept_the_shared_arguments(): + """Each tool function takes exactly the shared schema's properties, with the same + required ones. (The function signature is what validates arguments in Python.)""" + pytest.importorskip("mcp") + from pokemon_automation.mcp_server import ServerConfig, create_server + + server, _ = create_server(ServerConfig(fake=True)) + for name, definition in agent_tools.tools_for("python").items(): + derived = server._tool_manager.get_tool(name).fn_metadata.arg_model.model_json_schema() + schema = definition["inputSchema"] + assert set(derived.get("properties", {})) == set(schema.get("properties", {})), name + assert set(derived.get("required", [])) == set(schema.get("required", [])), name + + CASES = agent_tools.load_test_cases() From 3dd30e45629071330dd30873601fd64c63dc80ab Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:13:18 -0700 Subject: [PATCH 10/16] Python: add a README How to build _pa_core, install the package, run the self-test, use the Python API and the MCP server, and run the tests. Co-Authored-By: Claude Opus 5.5 --- .../Source/PythonBindings/README.md | 172 ++++++++++++++++++ .../Source/PythonBindings/pyproject.toml | 1 + 2 files changed, 173 insertions(+) create mode 100644 SerialPrograms/Source/PythonBindings/README.md diff --git a/SerialPrograms/Source/PythonBindings/README.md b/SerialPrograms/Source/PythonBindings/README.md new file mode 100644 index 0000000000..1f926fc954 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/README.md @@ -0,0 +1,172 @@ +# Python bindings and MCP server + +Control a Nintendo Switch from Python, or let an AI agent control it over MCP, using +the same controller hardware as the main program (an ESP32/Pico running PABotBase2 +firmware) plus a capture card. + +``` + MCP server (pokemon_automation.mcp_server) your Python scripts + \ / + pokemon_automation (Python API: Console, SwitchController, VideoSource) + / \ + _pa_core (C++/pybind11, CoreLib only) opencv-python / pytesseract + serial + PABotBase2, acts as a controller video capture, encoding, OCR +``` + +- **`_pa_core`** is built only from this codebase (`CoreLib`: serial port, PABotBase2 + protocol, controller scheduling). It links no Qt, OpenCV or Tesseract. +- **Vision is pure Python**: capture and encoding use `opencv-python`; OCR uses + `pytesseract` if installed. +- **The MCP server** is a thin layer over the Python API, so every capability is + available to scripts and to agents alike. + +## Build + +The module only needs CoreLib, the GUI-free core, so use the core-only build. It +needs only a C++23 compiler, CMake and Python (with pybind11, or CMake downloads it): +no Qt, OpenCV, ONNX Runtime, Tesseract or DPP. + +```bash +mkdir -p build-core && cd build-core +cmake ../SerialPrograms -DPA_CORE_ONLY=ON -DPA_PYTHON_BINDINGS=ON \ + -DPython_EXECUTABLE=$(which python3) -DCMAKE_BUILD_TYPE=Release +cmake --build . -j 10 # CoreLib, SerialProgramsCommandLine, _pa_core +``` + +`-DPA_PYTHON_BINDINGS=ON` also works in a full build (next to the GUI), with +`--target _pa_core`. + +The built module is copied into `pokemon_automation/` automatically. Then install the +package (editable) with the extras you want: + +```bash +pip install -e "SerialPrograms/Source/PythonBindings[mcp,ocr]" +``` + +pybind11 is taken from your environment if installed (`pip install pybind11` or +Homebrew), otherwise downloaded by CMake. OCR also needs the `tesseract` program +(`brew install tesseract`). + +## Self-test + +Put the Switch on the Home menu, then: + +```bash +python -m pokemon_automation.selftest --serial /dev/cu.usbserial-0001 --video MiraBox +``` + +It lists devices, presses HOME and a few d-pad inputs, and saves screenshots to +`selftest_output/`. **macOS:** run it from Terminal or iTerm the first time. macOS only +asks for camera permission on behalf of apps that declare camera usage; processes +started from other apps may be denied silently. + +## Python API + +```python +from pokemon_automation import Console, list_serial_ports, list_video_devices + +print(list_serial_ports(), list_video_devices()) + +with Console(serial_port="/dev/cu.usbserial-0001", video="MiraBox") as console: + frame = console.act([{"buttons": "HOME"}], settle_ms=1000) # press, wait, capture + frame.save("home.png") + + console.send([ + {"buttons": "DOWN", "repeat": 3}, # d-pad down x3 + {"buttons": "A"}, # tap A + {"wait_ms": 1000}, + {"left_stick": "up", "buttons": "B", "hold_ms": 2000, "release_ms": 0}, # run + ]) + print(console.read_text(box=(0.05, 0.8, 0.9, 0.15), mode="line")) +``` + +The controller alone, without video: + +```python +from pokemon_automation import SwitchController + +with SwitchController("/dev/cu.usbserial-0001") as sw: + sw.press("A") # 80 ms hold, 80 ms release + sw.press("L+R", hold_ms=200) + sw.stick("left", "up", duration_ms=1500) + sw.hold("ZL", 1000, right=[0.5, 0]) # anything held together + sw.flush() # block until executed +``` + +Conventions: +- Durations are milliseconds. Inputs are queued on the device and return + immediately; `flush()` / `Console.send(wait=True)` blocks until they've executed. + `cancel_all_commands_blocking(timeout_ms)` cancels queued inputs, releases + everything and waits for the device to confirm the neutral state, from any thread. + It returns False on timeout, e.g. when the Switch is asleep and the device can't + execute commands. +- Buttons: `A B X Y L R ZL ZR PLUS MINUS HOME CAPTURE LCLICK RCLICK`, d-pad `UP DOWN LEFT + RIGHT UP_RIGHT ...`; combine with `+` or a list. Aliases such as `+`, `START`, `L3` + work too (see `buttons.py`). +- Sticks: direction names or `[x, y]` in [-1, 1], +y = up. +- Images: numpy `(height, width, 3)` uint8 RGB. Boxes: `(x, y, width, height)` as + fractions of the frame, the same as `ImageFloatBox` in C++. +- The controller type (Pro Controller, wired, Switch 1/2) is whatever the device is set + to; choose it once in the main program. + +## MCP server + +The tools are defined in the shared `Source/Integrations/AgentServer/AgentTools.json`. + +```bash +python -m pokemon_automation.mcp_server --serial /dev/cu.usbserial-0001 --video MiraBox +python -m pokemon_automation.mcp_server --fake # no hardware, for trying it out +``` + +Register with Claude Code: + +```bash +claude mcp add switch -- /path/to/python -m pokemon_automation.mcp_server \ + --serial /dev/cu.usbserial-0001 --video MiraBox +``` + +Use `--transport streamable-http --host 127.0.0.1 --port 8765` for HTTP clients. + +**macOS camera permission:** the camera is granted per *app*. If an app without camera +permission (e.g. the Claude desktop app) launches the server, video fails to open. In +that case run the server from Terminal over HTTP and connect the agent to it: + +```bash +# in Terminal.app +python -m pokemon_automation.mcp_server --serial /dev/cu.usbserial-0001 --video MiraBox \ + --transport streamable-http --port 8765 +# then +claude mcp add --transport http switch http://127.0.0.1:8765/mcp +``` + +| Tool | Purpose | +|---|---| +| `switch_status`, `list_devices`, `connect` | Connection state, device discovery, (re)connect | +| `screenshot`, `wait_and_observe` | Current screen as JPEG (optionally cropped) | +| `press_buttons`, `move_stick`, `run_inputs` | Inputs, each returning a screenshot taken `settle_ms` after they finish | +| `read_text` | OCR a region of the screen | +| `cancel_all_commands_blocking` | Emergency stop; works while another input call is running, and reports whether the device confirmed the neutral state | +| `get_logs` | Recent connection/input log lines | + +Safety options: `--max-hold-ms` (default 10 s per step), `--max-sequence-ms` (default +60 s per call), `--read-only` (no inputs at all). Logs go to stderr (never stdout, +which carries the protocol) and optionally `--log-file`. + +## Tests + +```bash +cd SerialPrograms/Source/PythonBindings +pip install -e ".[test]" +pytest +``` + +The tests use fake devices (`pokemon_automation.fake`) and need no hardware. + +## Files + +- `PythonBindings.cmake`: build rules, included by `CMakeLists.txt` when + `PA_PYTHON_BINDINGS=ON`. +- `PythonBindings_Module.cpp`: the pybind11 module `_pa_core`. It wraps + `Source/Integrations/PybindSwitchController.*`, the GUI-free controller class in + CoreLib (also used by `SerialProgramsCommandLine`). +- `pokemon_automation/`: the Python package (API, vision, MCP server, fakes, self-test). diff --git a/SerialPrograms/Source/PythonBindings/pyproject.toml b/SerialPrograms/Source/PythonBindings/pyproject.toml index a419cb7527..8f3b8c38f1 100644 --- a/SerialPrograms/Source/PythonBindings/pyproject.toml +++ b/SerialPrograms/Source/PythonBindings/pyproject.toml @@ -10,6 +10,7 @@ build-backend = "setuptools.build_meta" name = "pokemon-automation" version = "0.1.0" description = "Control a Nintendo Switch from Python with Pokemon Automation hardware." +readme = "README.md" requires-python = ">=3.10" dependencies = [ "numpy>=1.24", From 5ef1a49df60b0bef3dc945223c9e2eb45a1eddab Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:14:29 -0700 Subject: [PATCH 11/16] AgentServer: add a minimal HTTP server on QTcpServer The app needs a local HTTP endpoint for AI agents (MCP's Streamable HTTP transport). Qt's QHttpServer module is GPL-3.0-only, which this project doesn't take, so this is a small HTTP/1.1 server on QTcpServer (Qt Network, LGPL): Content-Length bodies, keep-alive and "Expect: 100-continue"; no chunked requests or pipelining. The listener and sockets live on their own thread with its own event loop; each request is handled on a worker thread, so a slow request doesn't block others. stop() closes connections and waits for running handlers. Qt Network is now listed explicitly. The app already used it (FileDownloader, DiscordWebhook) but only got it through other Qt modules. Co-Authored-By: Claude Opus 5.5 --- SerialPrograms/CMakeLists.txt | 4 + .../AgentServer/AgentServer_HttpServer.cpp | 501 ++++++++++++++++++ .../AgentServer/AgentServer_HttpServer.h | 119 +++++ SerialPrograms/cmake/SourceFiles.cmake | 2 + 4 files changed, 626 insertions(+) create mode 100644 SerialPrograms/Source/Integrations/AgentServer/AgentServer_HttpServer.cpp create mode 100644 SerialPrograms/Source/Integrations/AgentServer/AgentServer_HttpServer.h diff --git a/SerialPrograms/CMakeLists.txt b/SerialPrograms/CMakeLists.txt index 6634e2dfcc..8ace70c71f 100644 --- a/SerialPrograms/CMakeLists.txt +++ b/SerialPrograms/CMakeLists.txt @@ -115,6 +115,7 @@ if(WIN32 AND QT_DEPLOY_FILES) Qml Quick QuickWidgets + Network ) else() # Find all subdirectories in the Qt base directory @@ -127,6 +128,7 @@ if(WIN32 AND QT_DEPLOY_FILES) Qml Quick QuickWidgets + Network ) file(GLOB QT_SUBDIRS LIST_DIRECTORIES true "${QT_BASE_DIR}/${QT_MAJOR}*") @@ -164,6 +166,7 @@ else() Qml Quick QuickWidgets + Network ) endif() @@ -272,6 +275,7 @@ function(apply_common_target_properties target_name) Qt${QT_MAJOR}::Qml Qt${QT_MAJOR}::Quick Qt${QT_MAJOR}::QuickWidgets + Qt${QT_MAJOR}::Network ) # more include directories diff --git a/SerialPrograms/Source/Integrations/AgentServer/AgentServer_HttpServer.cpp b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_HttpServer.cpp new file mode 100644 index 0000000000..6363bec603 --- /dev/null +++ b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_HttpServer.cpp @@ -0,0 +1,501 @@ +/* Agent Server: HTTP Server + * + * From: https://github.com/PokemonAutomation/ + * + */ + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include "Common/Cpp/Color.h" +#include "Common/Cpp/Logging/AbstractLogger.h" +#include "AgentServer_HttpServer.h" + +namespace PokemonAutomation{ +namespace AgentServer{ + + + +const std::string& HttpRequest::header(const std::string& name) const{ + static const std::string EMPTY; + auto iter = headers.find(name); + return iter == headers.end() ? EMPTY : iter->second; +} + +HttpResponse HttpResponse::text(int status, std::string body){ + HttpResponse ret; + ret.status = status; + ret.content_type = "text/plain; charset=utf-8"; + ret.body = std::move(body); + return ret; +} +HttpResponse HttpResponse::json(int status, std::string body){ + HttpResponse ret; + ret.status = status; + ret.content_type = "application/json"; + ret.body = std::move(body); + return ret; +} + + + +namespace{ + +const size_t MAX_HEADER_BYTES = 32 * 1024; + +std::string to_lower(std::string text){ + for (char& ch : text){ + ch = (char)std::tolower((unsigned char)ch); + } + return text; +} +std::string trim(const std::string& text){ + size_t start = text.find_first_not_of(" \t"); + if (start == std::string::npos){ + return ""; + } + size_t end = text.find_last_not_of(" \t"); + return text.substr(start, end - start + 1); +} + +HttpParseResult parse_error(int status, std::string message){ + HttpParseResult ret; + ret.status = HttpParseStatus::ERROR; + ret.error_status = status; + ret.error_message = std::move(message); + return ret; +} + +const char* reason_phrase(int status){ + switch (status){ + case 100: return "Continue"; + case 200: return "OK"; + case 202: return "Accepted"; + case 204: return "No Content"; + case 400: return "Bad Request"; + case 401: return "Unauthorized"; + case 403: return "Forbidden"; + case 404: return "Not Found"; + case 405: return "Method Not Allowed"; + case 406: return "Not Acceptable"; + case 411: return "Length Required"; + case 413: return "Content Too Large"; + case 415: return "Unsupported Media Type"; + case 431: return "Request Header Fields Too Large"; + case 500: return "Internal Server Error"; + case 501: return "Not Implemented"; + case 503: return "Service Unavailable"; + case 505: return "HTTP Version Not Supported"; + default: return "Unknown"; + } +} + +} + + + +HttpParseResult parse_http_request(std::string& buffer, HttpRequest& request, size_t max_body_bytes){ + size_t header_end = buffer.find("\r\n\r\n"); + if (header_end == std::string::npos){ + if (buffer.size() > MAX_HEADER_BYTES){ + return parse_error(431, "Request headers are too large."); + } + return HttpParseResult(); + } + if (header_end > MAX_HEADER_BYTES){ + return parse_error(431, "Request headers are too large."); + } + + HttpRequest parsed; + + // Request line: METHOD SP TARGET SP VERSION + size_t line_end = buffer.find("\r\n"); + std::string request_line = buffer.substr(0, line_end); + size_t space0 = request_line.find(' '); + size_t space1 = space0 == std::string::npos ? std::string::npos : request_line.find(' ', space0 + 1); + if (space0 == std::string::npos || space1 == std::string::npos){ + return parse_error(400, "Malformed request line."); + } + parsed.method = request_line.substr(0, space0); + std::string target = request_line.substr(space0 + 1, space1 - space0 - 1); + std::string version = request_line.substr(space1 + 1); + if (!version.starts_with("HTTP/1.")){ + return parse_error(505, "Only HTTP/1.x is supported."); + } + size_t question = target.find('?'); + parsed.path = target.substr(0, question); + if (question != std::string::npos){ + parsed.query = target.substr(question + 1); + } + + // Headers + size_t pos = line_end + 2; + while (pos < header_end){ + size_t end = buffer.find("\r\n", pos); + if (end == std::string::npos || end > header_end){ + end = header_end; + } + std::string line = buffer.substr(pos, end - pos); + pos = end + 2; + if (line.empty()){ + continue; + } + if (line[0] == ' ' || line[0] == '\t'){ + return parse_error(400, "Folded header lines are not supported."); + } + size_t colon = line.find(':'); + if (colon == std::string::npos || colon == 0){ + return parse_error(400, "Malformed header line."); + } + std::string name = to_lower(trim(line.substr(0, colon))); + std::string value = trim(line.substr(colon + 1)); + auto iter = parsed.headers.find(name); + if (iter == parsed.headers.end()){ + parsed.headers.emplace(std::move(name), std::move(value)); + }else{ + iter->second += ", " + value; + } + } + + std::string connection = to_lower(parsed.header("connection")); + if (version == "HTTP/1.0"){ + parsed.keep_alive = connection.find("keep-alive") != std::string::npos; + }else{ + parsed.keep_alive = connection.find("close") == std::string::npos; + } + + // Body + if (!parsed.header("transfer-encoding").empty()){ + return parse_error(411, "Chunked request bodies are not supported; send Content-Length."); + } + size_t body_bytes = 0; + const std::string& length_text = parsed.header("content-length"); + if (!length_text.empty()){ + if (length_text.size() > 12 || !std::all_of(length_text.begin(), length_text.end(), ::isdigit)){ + return parse_error(400, "Invalid Content-Length."); + } + body_bytes = (size_t)std::stoull(length_text); + if (body_bytes > max_body_bytes){ + return parse_error(413, "Request body is too large."); + } + } + size_t body_start = header_end + 4; + if (buffer.size() - body_start < body_bytes){ + HttpParseResult ret; + ret.expect_continue = to_lower(parsed.header("expect")) == "100-continue"; + return ret; + } + + parsed.body = buffer.substr(body_start, body_bytes); + buffer.erase(0, body_start + body_bytes); + request = std::move(parsed); + + HttpParseResult ret; + ret.status = HttpParseStatus::COMPLETE; + return ret; +} + + +std::string serialize_http_response(const HttpResponse& response, bool keep_alive){ + std::string ret = "HTTP/1.1 " + std::to_string(response.status) + " " + reason_phrase(response.status) + "\r\n"; + if (!response.content_type.empty()){ + ret += "Content-Type: " + response.content_type + "\r\n"; + } + ret += "Content-Length: " + std::to_string(response.body.size()) + "\r\n"; + ret += keep_alive ? "Connection: keep-alive\r\n" : "Connection: close\r\n"; + for (const auto& header : response.headers){ + ret += header.first + ": " + header.second + "\r\n"; + } + ret += "\r\n"; + ret += response.body; + return ret; +} + + + + +// One client connection. Only touched on the server thread. +struct HttpConnection{ + QPointer socket; + std::string buffer; + bool busy = false; // a request is being handled; later bytes wait + bool sent_continue = false; // "100 Continue" already sent for this request +}; + + +struct HttpServer::Internal{ + Internal(Logger& p_logger, HttpHandler p_handler) + : logger(p_logger) + , handler(std::move(p_handler)) + {} + + Logger& logger; + HttpHandler handler; + + QThread thread; + QObject* context = nullptr; // lives on `thread`; target for posted work + QTcpServer* server = nullptr; // lives on `thread`; child of `context` + std::map> connections; // `thread` only + + std::atomic listening{false}; + std::atomic port{0}; + + // Worker threads post their responses through `context` while holding this + // lock, and `stop()` clears `context` under it, so no worker posts to a + // deleted object. + std::mutex post_lock; + QObject* post_target = nullptr; + + // Count of handlers still running, so `stop()` can wait for them. + std::mutex inflight_lock; + std::condition_variable inflight_cv; + size_t inflight = 0; + + + // ---- Everything below runs on `thread`. ---- + + void on_new_connection(){ + while (QTcpSocket* socket = server->nextPendingConnection()){ + auto connection = std::make_shared(); + connection->socket = socket; + connections[socket] = connection; + QObject::connect(socket, &QTcpSocket::readyRead, context, [this, socket]{ + on_ready_read(socket); + }); + QObject::connect(socket, &QTcpSocket::disconnected, context, [this, socket]{ + connections.erase(socket); + socket->deleteLater(); + }); + } + } + + void on_ready_read(QTcpSocket* socket){ + auto iter = connections.find(socket); + if (iter == connections.end()){ + return; + } + QByteArray data = socket->readAll(); + iter->second->buffer.append(data.constData(), (size_t)data.size()); + process_buffer(iter->second); + } + + // Parse and dispatch the next request on `connection`, if one is complete. + void process_buffer(const std::shared_ptr& connection){ + QTcpSocket* socket = connection->socket; + if (socket == nullptr || connection->busy){ + return; + } + + HttpRequest request; + HttpParseResult result = parse_http_request(connection->buffer, request); + switch (result.status){ + case HttpParseStatus::NEED_MORE: + if (result.expect_continue && !connection->sent_continue){ + connection->sent_continue = true; + socket->write("HTTP/1.1 100 Continue\r\n\r\n"); + } + return; + case HttpParseStatus::ERROR:{ + logger.log("[AgentServer] Rejected HTTP request: " + result.error_message, COLOR_RED); + std::string bytes = serialize_http_response( + HttpResponse::text(result.error_status, result.error_message + "\n"), false + ); + socket->write(bytes.data(), (qint64)bytes.size()); + socket->disconnectFromHost(); + return; + } + case HttpParseStatus::COMPLETE: + break; + } + + connection->busy = true; + connection->sent_continue = false; + request.peer_address = socket->peerAddress().toString().toStdString(); + + { + std::lock_guard lg(inflight_lock); + inflight++; + } + std::weak_ptr weak = connection; + std::thread([this, weak, request = std::move(request)]{ + run_handler(weak, request); + }).detach(); + } + + // ---- Worker thread. ---- + + void run_handler(std::weak_ptr weak, const HttpRequest& request){ + HttpResponse response; + try{ + response = handler(request); + }catch (const std::exception& e){ + logger.log(std::string("[AgentServer] Request handler failed: ") + e.what(), COLOR_RED); + response = HttpResponse::text(500, "Internal server error.\n"); + }catch (...){ + logger.log("[AgentServer] Request handler failed.", COLOR_RED); + response = HttpResponse::text(500, "Internal server error.\n"); + } + bool keep_alive = request.keep_alive; + std::string bytes = serialize_http_response(response, keep_alive); + + { + std::lock_guard lg(post_lock); + if (post_target != nullptr){ + QMetaObject::invokeMethod(post_target, [this, weak, bytes = std::move(bytes), keep_alive]{ + send_response(weak, bytes, keep_alive); + }, Qt::QueuedConnection); + } + } + + { + std::lock_guard lg(inflight_lock); + inflight--; + } + inflight_cv.notify_all(); + } + + // ---- Back on `thread`. ---- + + void send_response(const std::weak_ptr& weak, const std::string& bytes, bool keep_alive){ + std::shared_ptr connection = weak.lock(); + if (!connection || connection->socket == nullptr){ + return; // client went away while the handler ran + } + QTcpSocket* socket = connection->socket; + socket->write(bytes.data(), (qint64)bytes.size()); + connection->busy = false; + if (!keep_alive){ + socket->disconnectFromHost(); + return; + } + // The client may have sent its next request already. + process_buffer(connection); + } +}; + + + +HttpServer::HttpServer(Logger& logger, HttpHandler handler) + : m_internal(std::make_unique(logger, std::move(handler))) +{} +HttpServer::~HttpServer(){ + stop(); +} + +std::string HttpServer::start(uint16_t port, bool bind_all_interfaces){ + Internal& data = *m_internal; + if (data.thread.isRunning()){ + return "The server is already running."; + } + + data.thread.setObjectName("AgentServer HTTP"); + data.thread.start(); + data.context = new QObject(); + data.context->moveToThread(&data.thread); + + std::string error; + bool dispatched = QMetaObject::invokeMethod(data.context, [&]{ + data.server = new QTcpServer(data.context); + QObject::connect(data.server, &QTcpServer::newConnection, data.context, [&data]{ + data.on_new_connection(); + }); + QHostAddress address = bind_all_interfaces ? QHostAddress(QHostAddress::Any) : QHostAddress(QHostAddress::LocalHost); + if (!data.server->listen(address, port)){ + error = "Unable to listen on port " + std::to_string(port) + ": " + + data.server->errorString().toStdString(); + return; + } + data.port.store(data.server->serverPort(), std::memory_order_release); + }, Qt::BlockingQueuedConnection); + if (!dispatched){ + // E.g. no QCoreApplication exists, so the server thread has no event loop. + error = "Unable to start the server thread."; + } + + if (!error.empty()){ + data.logger.log("[AgentServer] " + error, COLOR_RED); + stop(); + return error; + } + + { + std::lock_guard lg(data.post_lock); + data.post_target = data.context; + } + data.listening.store(true, std::memory_order_release); + data.logger.log( + "[AgentServer] Listening on " + std::string(bind_all_interfaces ? "all interfaces" : "127.0.0.1") + + ", port " + std::to_string(data.port.load()), + COLOR_BLUE + ); + return ""; +} + +void HttpServer::stop(){ + Internal& data = *m_internal; + if (!data.thread.isRunning()){ + return; + } + data.listening.store(false, std::memory_order_release); + + // 1. On the server thread: stop accepting, then drop and delete every + // connection and the listener. (Socket objects must be deleted on the + // thread that owns them.) + QMetaObject::invokeMethod(data.context, [&data]{ + for (auto& item : data.connections){ + QTcpSocket* socket = item.second->socket; + if (socket != nullptr){ + QObject::disconnect(socket, nullptr, data.context, nullptr); + socket->abort(); + delete socket; + } + } + data.connections.clear(); + delete data.server; // also deletes any not-yet-accepted connections + data.server = nullptr; + }, Qt::BlockingQueuedConnection); + + // 2. Stop workers from posting responses, then wait for running handlers. + { + std::lock_guard lg(data.post_lock); + data.post_target = nullptr; + } + { + std::unique_lock lg(data.inflight_lock); + while (!data.inflight_cv.wait_for(lg, std::chrono::seconds(5), [&]{ return data.inflight == 0; })){ + data.logger.log( + "[AgentServer] Waiting for " + std::to_string(data.inflight) + " request(s) to finish...", + COLOR_ORANGE + ); + } + } + + // 3. Stop the thread. `context` is a plain QObject with nothing left on it, + // so it can be deleted here once its thread has finished. + data.thread.quit(); + data.thread.wait(); + delete data.context; + data.context = nullptr; + data.port.store(0, std::memory_order_release); + data.logger.log("[AgentServer] Stopped.", COLOR_BLUE); +} + +bool HttpServer::is_listening() const{ + return m_internal->listening.load(std::memory_order_acquire); +} +uint16_t HttpServer::port() const{ + return m_internal->port.load(std::memory_order_acquire); +} + + + +} +} diff --git a/SerialPrograms/Source/Integrations/AgentServer/AgentServer_HttpServer.h b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_HttpServer.h new file mode 100644 index 0000000000..468c328cc1 --- /dev/null +++ b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_HttpServer.h @@ -0,0 +1,119 @@ +/* Agent Server: HTTP Server + * + * From: https://github.com/PokemonAutomation/ + * + * A minimal HTTP/1.1 server on Qt's QTcpServer, just enough to serve MCP's + * "Streamable HTTP" transport to AI agents (see AgentServer_McpServer.h). + * + * Qt's own QHttpServer module is not used because it is GPL-3.0-only, and this + * project does not take plain GPL dependencies. QTcpServer (Qt Network) is LGPL. + * + * Supported: one request at a time per connection, keep-alive, Content-Length + * request bodies, "Expect: 100-continue". Not supported (answered with an error): + * chunked request bodies, pipelining, upgrades. + */ + +#ifndef PokemonAutomation_AgentServer_HttpServer_H +#define PokemonAutomation_AgentServer_HttpServer_H + +#include +#include +#include +#include +#include +#include +#include + +namespace PokemonAutomation{ + class Logger; +namespace AgentServer{ + + +struct HttpRequest{ + std::string method; // e.g. "POST" + std::string path; // e.g. "/mcp" (without the query string) + std::string query; // the part after '?', if any + std::map headers; // keys are lower-case + std::string body; + std::string peer_address; // remote IP address, for logging + bool keep_alive = true; // HTTP/1.1 default, unless "Connection: close" + + // Returns the header value, or an empty string if absent. `name` must be lower-case. + const std::string& header(const std::string& name) const; +}; + +struct HttpResponse{ + int status = 200; + std::string content_type; // empty = no Content-Type header (e.g. for 202) + std::string body; + std::vector> headers; // extra headers + + static HttpResponse text(int status, std::string body); + static HttpResponse json(int status, std::string body); +}; + +// Called on a worker thread for every complete request. May block (e.g. while an +// agent's tool call runs); other connections are served meanwhile. +using HttpHandler = std::function; + + +// Result of trying to parse one request from the front of a connection's buffer. +enum class HttpParseStatus{ + NEED_MORE, // incomplete; wait for more bytes + COMPLETE, // `request` is filled and its bytes were removed from the buffer + ERROR, // malformed or unsupported; `error_status` says which HTTP error +}; +struct HttpParseResult{ + HttpParseStatus status = HttpParseStatus::NEED_MORE; + int error_status = 0; // 400, 411, 413, 431, 501 or 505 when ERROR + std::string error_message; + bool expect_continue = false; // headers ask for "100 Continue" before the body +}; + +// Parse one HTTP/1.1 request from the front of `buffer`. +// Limits: 32 KB of headers, `max_body_bytes` of body. +// This is separate from the server so it can be unit-tested. +HttpParseResult parse_http_request( + std::string& buffer, HttpRequest& request, + size_t max_body_bytes = 8 * 1024 * 1024 +); + +// Serialize a response. `keep_alive` selects the Connection header. +std::string serialize_http_response(const HttpResponse& response, bool keep_alive); + + + +// Listens on one address/port and hands each request to `handler`. +// +// Threads: the QTcpServer and all sockets live on a private thread with its own +// Qt event loop, so this works whether or not the caller runs one. (A +// QCoreApplication must exist, as it always does in the app.) Each request +// runs on its own worker thread. `stop()` (and the destructor) closes the listener +// and all connections and waits for running handlers to return. +class HttpServer{ + HttpServer(const HttpServer&) = delete; + void operator=(const HttpServer&) = delete; + +public: + HttpServer(Logger& logger, HttpHandler handler); + ~HttpServer(); + + // Start listening. `bind_all_interfaces` = false binds 127.0.0.1 only. + // Returns an empty string on success, or the error (e.g. port in use). + std::string start(uint16_t port, bool bind_all_interfaces); + + void stop(); + + bool is_listening() const; + uint16_t port() const; // the actual port (useful when started with port 0) + +private: + struct Internal; + std::unique_ptr m_internal; +}; + + + +} +} +#endif diff --git a/SerialPrograms/cmake/SourceFiles.cmake b/SerialPrograms/cmake/SourceFiles.cmake index f5e1ea9435..adb587b6d6 100644 --- a/SerialPrograms/cmake/SourceFiles.cmake +++ b/SerialPrograms/cmake/SourceFiles.cmake @@ -904,6 +904,8 @@ file(GLOB LIBRARY_SOURCES Source/Controllers/PABotBase2/SerialPABotBase_StatusThread.h Source/Controllers/SerialPort/SerialPortPollerQt.cpp Source/Controllers/SerialPort/SerialPortPollerQt.h + Source/Integrations/AgentServer/AgentServer_HttpServer.cpp + Source/Integrations/AgentServer/AgentServer_HttpServer.h Source/Integrations/DiscordIntegrationSettings.cpp Source/Integrations/DiscordIntegrationSettings.h Source/Integrations/DiscordIntegrationTable.cpp From ee66ae9561fb6a1875538b98bb5cd04d418ee450 Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:14:29 -0700 Subject: [PATCH 12/16] AgentServer: load the tool definitions from AgentTools.json - AgentTools.json is compiled into SerialProgramsLib with a new CMake helper, pa_embed_text_file() (cmake/EmbedTextFile.cmake), as a byte array, so it isn't limited by compilers' string-literal size limits. - AgentToolDefinitions parses it, keeps the tools for host "app", inlines $refs (same as the Python loader), and validates tool arguments against the schemas with a small validator for the keywords the file uses. Co-Authored-By: Claude Opus 5.5 --- SerialPrograms/CMakeLists.txt | 10 + .../AgentServer_ToolDefinitions.cpp | 290 ++++++++++++++++++ .../AgentServer/AgentServer_ToolDefinitions.h | 87 ++++++ SerialPrograms/cmake/EmbedTextFile.cmake | 73 +++++ SerialPrograms/cmake/SourceFiles.cmake | 2 + 5 files changed, 462 insertions(+) create mode 100644 SerialPrograms/Source/Integrations/AgentServer/AgentServer_ToolDefinitions.cpp create mode 100644 SerialPrograms/Source/Integrations/AgentServer/AgentServer_ToolDefinitions.h create mode 100644 SerialPrograms/cmake/EmbedTextFile.cmake diff --git a/SerialPrograms/CMakeLists.txt b/SerialPrograms/CMakeLists.txt index 8ace70c71f..d04fed9ba5 100644 --- a/SerialPrograms/CMakeLists.txt +++ b/SerialPrograms/CMakeLists.txt @@ -284,6 +284,16 @@ function(apply_common_target_properties target_name) endfunction() +# Compile the shared AI agent tool definitions into the app (see +# Source/Integrations/AgentServer/AgentTools.json). +include(cmake/EmbedTextFile.cmake) +pa_embed_text_file( + SerialProgramsLib + ${CMAKE_CURRENT_SOURCE_DIR}/Source/Integrations/AgentServer/AgentTools.json + ${CMAKE_CURRENT_BINARY_DIR}/Generated/AgentServer_AgentToolsJson.cpp + "PokemonAutomation::AgentServer::agent_tools_json_text" +) + # Apply common properties to both targets apply_common_target_properties(SerialProgramsLib) apply_common_target_properties(SerialPrograms) diff --git a/SerialPrograms/Source/Integrations/AgentServer/AgentServer_ToolDefinitions.cpp b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_ToolDefinitions.cpp new file mode 100644 index 0000000000..217d233fd8 --- /dev/null +++ b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_ToolDefinitions.cpp @@ -0,0 +1,290 @@ +/* Agent Server: Tool Definitions + * + * From: https://github.com/PokemonAutomation/ + * + */ + +#include +#include +#include +#include "AgentServer_ToolDefinitions.h" + +namespace PokemonAutomation{ +namespace AgentServer{ + +using nlohmann::json; + + + +json resolve_schema_refs(const json& schema, const json& definitions){ + if (schema.is_array()){ + json ret = json::array(); + for (const json& item : schema){ + ret.push_back(resolve_schema_refs(item, definitions)); + } + return ret; + } + if (!schema.is_object()){ + return schema; + } + auto ref = schema.find("$ref"); + if (ref != schema.end()){ + const std::string prefix = "#/definitions/"; + std::string target = ref->is_string() ? ref->get() : ""; + if (!target.starts_with(prefix)){ + throw std::runtime_error("Unsupported $ref: " + target); + } + std::string name = target.substr(prefix.size()); + auto definition = definitions.find(name); + if (definition == definitions.end()){ + throw std::runtime_error("Unknown $ref: " + target); + } + json merged = *definition; + for (auto iter = schema.begin(); iter != schema.end(); ++iter){ + if (iter.key() != "$ref"){ + merged[iter.key()] = iter.value(); + } + } + return resolve_schema_refs(merged, definitions); + } + json ret = json::object(); + for (auto iter = schema.begin(); iter != schema.end(); ++iter){ + ret[iter.key()] = resolve_schema_refs(iter.value(), definitions); + } + return ret; +} + + + +namespace{ + +std::string json_type_name(const json& value){ + if (value.is_null()) return "null"; + if (value.is_boolean()) return "boolean"; + if (value.is_number_integer()) return "integer"; + if (value.is_number()) return "number"; + if (value.is_string()) return "string"; + if (value.is_array()) return "array"; + return "object"; +} + +bool matches_type(const std::string& type, const json& value){ + if (type == "object") return value.is_object(); + if (type == "array") return value.is_array(); + if (type == "string") return value.is_string(); + if (type == "boolean") return value.is_boolean(); + if (type == "null") return value.is_null(); + if (type == "number") return value.is_number(); + if (type == "integer"){ + if (value.is_number_integer()){ + return true; + } + // JSON has one number type; 5.0 is an integer as far as schemas go. + if (value.is_number_float()){ + double x = value.get(); + return std::isfinite(x) && x == std::floor(x); + } + return false; + } + return false; +} + +std::string number_text(double x){ + std::ostringstream ss; + ss << x; + return ss.str(); +} + +// "a string", "an integer", "a string or array" +std::string with_article(const std::string& type_list){ + char first = type_list.empty() ? 'x' : type_list[0]; + bool vowel = first == 'a' || first == 'e' || first == 'i' || first == 'o' || first == 'u'; + return (vowel ? "an " : "a ") + type_list; +} + +std::string where(const std::string& path){ + return path.empty() ? "arguments" : path; +} + +} + + +std::string validate_json_schema(const json& schema, const json& value, const std::string& path){ + if (!schema.is_object()){ + return ""; + } + + auto any_of = schema.find("anyOf"); + if (any_of != schema.end() && any_of->is_array()){ + // Valid if any option accepts the value. Otherwise report the error of the + // option with the value's type (e.g. "must be <= 1" for a stick array), or + // list the accepted types if none has it. + std::string type_matched_error; + std::string expected; + bool valid = false; + for (const json& option : *any_of){ + std::string error = validate_json_schema(option, value, path); + if (error.empty()){ + valid = true; + break; + } + auto option_type = option.find("type"); + if (option_type != option.end() && option_type->is_string()){ + expected += (expected.empty() ? "" : " or ") + option_type->get(); + if (type_matched_error.empty() && matches_type(option_type->get(), value)){ + type_matched_error = error; + } + } + } + if (!valid){ + return type_matched_error.empty() + ? where(path) + " must be " + with_article(expected) + ", not " + json_type_name(value) + : type_matched_error; + } + } + + auto type = schema.find("type"); + if (type != schema.end() && type->is_string() && !matches_type(type->get(), value)){ + return where(path) + " must be " + with_article(type->get()) + ", not " + json_type_name(value); + } + + auto enum_values = schema.find("enum"); + if (enum_values != schema.end() && enum_values->is_array()){ + bool found = false; + for (const json& allowed : *enum_values){ + found |= allowed == value; + } + if (!found){ + return where(path) + " must be one of " + enum_values->dump(); + } + } + + if (value.is_number()){ + double x = value.get(); + auto minimum = schema.find("minimum"); + if (minimum != schema.end() && x < minimum->get()){ + return where(path) + " must be >= " + number_text(minimum->get()); + } + auto maximum = schema.find("maximum"); + if (maximum != schema.end() && x > maximum->get()){ + return where(path) + " must be <= " + number_text(maximum->get()); + } + } + + if (value.is_array()){ + auto min_items = schema.find("minItems"); + if (min_items != schema.end() && value.size() < min_items->get()){ + return where(path) + " must have at least " + std::to_string(min_items->get()) + " item(s)"; + } + auto max_items = schema.find("maxItems"); + if (max_items != schema.end() && value.size() > max_items->get()){ + return where(path) + " must have at most " + std::to_string(max_items->get()) + " item(s)"; + } + auto items = schema.find("items"); + if (items != schema.end()){ + for (size_t c = 0; c < value.size(); c++){ + std::string error = validate_json_schema(*items, value[c], path + "[" + std::to_string(c) + "]"); + if (!error.empty()){ + return error; + } + } + } + } + + if (value.is_object()){ + auto properties = schema.find("properties"); + auto required = schema.find("required"); + if (required != schema.end()){ + for (const json& name : *required){ + if (!value.contains(name.get())){ + return where(path) + " is missing required property \"" + name.get() + "\""; + } + } + } + auto additional = schema.find("additionalProperties"); + bool allow_additional = additional == schema.end() || !additional->is_boolean() || additional->get(); + for (auto iter = value.begin(); iter != value.end(); ++iter){ + std::string child = path.empty() ? iter.key() : path + "." + iter.key(); + if (properties != schema.end() && properties->contains(iter.key())){ + std::string error = validate_json_schema((*properties)[iter.key()], iter.value(), child); + if (!error.empty()){ + return error; + } + }else if (!allow_additional){ + return "unknown argument \"" + child + "\""; + } + } + } + return ""; +} + + + +AgentToolDefinitions::AgentToolDefinitions(const std::string& json_text, const std::string& host){ + json data = json::parse(json_text); // throws json::parse_error (a std::exception) + + m_server_name = data.at("server_name").get(); + for (const json& line : data.at("instructions")){ + if (!m_instructions.empty()){ + m_instructions += "\n"; + } + m_instructions += line.get(); + } + + json definitions = data.value("definitions", json::object()); + for (const json& tool : data.at("tools")){ + bool hosted = false; + for (const json& h : tool.at("hosts")){ + hosted |= h.get() == host; + } + if (!hosted){ + continue; + } + AgentToolDefinition definition; + definition.name = tool.at("name").get(); + definition.description = tool.at("description").get(); + definition.input_schema = resolve_schema_refs(tool.at("inputSchema"), definitions); + m_index[definition.name] = m_tools.size(); + m_tools.emplace_back(std::move(definition)); + } +} + +const AgentToolDefinition* AgentToolDefinitions::find(const std::string& name) const{ + auto iter = m_index.find(name); + return iter == m_index.end() ? nullptr : &m_tools[iter->second]; +} + +json AgentToolDefinitions::tools_list() const{ + json ret = json::array(); + for (const AgentToolDefinition& tool : m_tools){ + ret.push_back({ + {"name", tool.name}, + {"description", tool.description}, + {"inputSchema", tool.input_schema}, + }); + } + return ret; +} + +std::string AgentToolDefinitions::validate_arguments(const AgentToolDefinition& tool, const json& arguments) const{ + return validate_json_schema(tool.input_schema, arguments, ""); +} + +json AgentToolDefinitions::with_defaults(const AgentToolDefinition& tool, const json& arguments){ + json ret = arguments.is_object() ? arguments : json::object(); + auto properties = tool.input_schema.find("properties"); + if (properties == tool.input_schema.end()){ + return ret; + } + for (auto iter = properties->begin(); iter != properties->end(); ++iter){ + if (!ret.contains(iter.key()) && iter.value().contains("default")){ + ret[iter.key()] = iter.value()["default"]; + } + } + return ret; +} + + + +} +} diff --git a/SerialPrograms/Source/Integrations/AgentServer/AgentServer_ToolDefinitions.h b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_ToolDefinitions.h new file mode 100644 index 0000000000..d0141ed2ca --- /dev/null +++ b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_ToolDefinitions.h @@ -0,0 +1,87 @@ +/* Agent Server: Tool Definitions + * + * From: https://github.com/PokemonAutomation/ + * + * Loads the shared MCP interface, AgentTools.json, which both this app and the + * Python package (pokemon_automation.mcp_server) serve, so AI agents see the same + * tools whichever server they connect to. The file is compiled into the app (see + * `agent_tools_json_text()`). + * + * Also validates tool arguments against the tools' JSON Schemas. Only the schema + * keywords AgentTools.json uses are supported: type, properties, required, + * additionalProperties, items, minItems, maxItems, minimum, maximum, enum, anyOf. + */ + +#ifndef PokemonAutomation_AgentServer_ToolDefinitions_H +#define PokemonAutomation_AgentServer_ToolDefinitions_H + +#include +#include +#include +#include "3rdParty-Core/nlohmann/json.hpp" + +namespace PokemonAutomation{ +namespace AgentServer{ + + +// The contents of AgentTools.json, embedded at build time. +const std::string& agent_tools_json_text(); + + +struct AgentToolDefinition{ + std::string name; + std::string description; + nlohmann::json input_schema; // self-contained: all "$ref"s are inlined +}; + + +class AgentToolDefinitions{ +public: + // Parse the definition file and keep the tools implemented by `host` + // ("app" or "python"). Throws std::runtime_error if the JSON is malformed or a + // "$ref" names an unknown definition. + AgentToolDefinitions(const std::string& json_text, const std::string& host); + + const std::string& server_name() const{ return m_server_name; } + const std::string& instructions() const{ return m_instructions; } + const std::vector& tools() const{ return m_tools; } + + // Returns nullptr if there is no tool named `name`. + const AgentToolDefinition* find(const std::string& name) const; + + // The `tools` array for an MCP tools/list response. + nlohmann::json tools_list() const; + + // Check `arguments` against the tool's input schema. + // Returns an empty string if valid, otherwise a message for the agent such as + // "steps[0].hold_ms must be >= 1". + std::string validate_arguments(const AgentToolDefinition& tool, const nlohmann::json& arguments) const; + + // Return `arguments` with every missing top-level property that has a schema + // "default" filled in. (Defaults of nested objects, such as run_inputs steps, + // are applied by the code that parses them.) + static nlohmann::json with_defaults(const AgentToolDefinition& tool, const nlohmann::json& arguments); + +private: + std::string m_server_name; + std::string m_instructions; + std::vector m_tools; + std::map m_index; +}; + + +// Replace every {"$ref": "#/definitions/", ...siblings} in `schema` with a copy +// of that definition merged with the siblings (siblings win). Same behavior as +// `resolve_refs()` in pokemon_automation/agent_tools.py. +// Throws std::runtime_error for an unknown or unsupported reference. +nlohmann::json resolve_schema_refs(const nlohmann::json& schema, const nlohmann::json& definitions); + +// Validate `value` against `schema`. Returns an empty string if valid, otherwise a +// message that starts with `path` (the location of `value`, e.g. "steps[0]"). +std::string validate_json_schema(const nlohmann::json& schema, const nlohmann::json& value, const std::string& path); + + + +} +} +#endif diff --git a/SerialPrograms/cmake/EmbedTextFile.cmake b/SerialPrograms/cmake/EmbedTextFile.cmake new file mode 100644 index 0000000000..ecd7c1b3c1 --- /dev/null +++ b/SerialPrograms/cmake/EmbedTextFile.cmake @@ -0,0 +1,73 @@ +# pa_embed_text_file( ) +# +# Compile a text file into as a function that returns its contents: +# +# pa_embed_text_file(MyLib data/Tools.json ${CMAKE_BINARY_DIR}/Tools.cpp "My::Space::tools_json") +# +# generates Tools.cpp defining `const std::string& My::Space::tools_json()`. Declare +# that function in a header of your own. The file is embedded as a byte array rather +# than a string literal, so its size and content are not limited by compilers' +# string-literal rules (e.g. MSVC's 16 KB limit). +# +# The .cpp is regenerated when changes: the file is registered as a +# configure dependency, so the next build re-runs CMake first. + +function(pa_embed_text_file target input output function_name) + set_property(DIRECTORY APPEND PROPERTY CMAKE_CONFIGURE_DEPENDS ${input}) + + file(READ ${input} hex HEX) + string(LENGTH "${hex}" hex_length) + set(bytes "") + set(column 0) + set(position 0) + while(position LESS hex_length) + string(SUBSTRING "${hex}" ${position} 2 byte) + string(APPEND bytes "0x${byte},") + math(EXPR position "${position} + 2") + math(EXPR column "${column} + 1") + if(column EQUAL 24) + string(APPEND bytes "\n ") + set(column 0) + endif() + endwhile() + + # "A::B::name" -> namespaces A, B and function `name` + string(REPLACE "::" ";" parts "${function_name}") + list(POP_BACK parts name) + set(open_namespaces "") + set(close_namespaces "") + foreach(part ${parts}) + string(APPEND open_namespaces "namespace ${part}{\n") + string(APPEND close_namespaces "}\n") + endforeach() + + file(RELATIVE_PATH input_name ${CMAKE_CURRENT_SOURCE_DIR} ${input}) + set(content +"// Generated by cmake/EmbedTextFile.cmake from ${input_name}. Do not edit. + +#include + +${open_namespaces} + +static const unsigned char EMBEDDED_TEXT[] = { + ${bytes}0x00 +}; + +const std::string& ${name}(){ + static const std::string text((const char*)EMBEDDED_TEXT, sizeof(EMBEDDED_TEXT) - 1); + return text; +} + +${close_namespaces}") + + # Only touch the output when it changes, so unrelated reconfigures don't rebuild it. + set(existing "") + if(EXISTS ${output}) + file(READ ${output} existing) + endif() + if(NOT existing STREQUAL content) + file(WRITE ${output} "${content}") + endif() + + target_sources(${target} PRIVATE ${output}) +endfunction() diff --git a/SerialPrograms/cmake/SourceFiles.cmake b/SerialPrograms/cmake/SourceFiles.cmake index adb587b6d6..c5362c3c6b 100644 --- a/SerialPrograms/cmake/SourceFiles.cmake +++ b/SerialPrograms/cmake/SourceFiles.cmake @@ -906,6 +906,8 @@ file(GLOB LIBRARY_SOURCES Source/Controllers/SerialPort/SerialPortPollerQt.h Source/Integrations/AgentServer/AgentServer_HttpServer.cpp Source/Integrations/AgentServer/AgentServer_HttpServer.h + Source/Integrations/AgentServer/AgentServer_ToolDefinitions.cpp + Source/Integrations/AgentServer/AgentServer_ToolDefinitions.h Source/Integrations/DiscordIntegrationSettings.cpp Source/Integrations/DiscordIntegrationSettings.h Source/Integrations/DiscordIntegrationTable.cpp From 8cce6e89c9c5f7e7569669589bfa12ce5611f3ea Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:14:29 -0700 Subject: [PATCH 13/16] AgentServer: parse input steps Parse the input vocabulary of AgentTools.json (button names, d-pad, sticks, input steps) into NintendoSwitch::Button, DpadPosition and JoystickPosition. The C++ twin of pokemon_automation/buttons.py; both pass the shared AgentInputTestCases.json. Co-Authored-By: Claude Opus 5.5 --- .../AgentServer/AgentServer_InputSteps.cpp | 403 ++++++++++++++++++ .../AgentServer/AgentServer_InputSteps.h | 94 ++++ SerialPrograms/cmake/SourceFiles.cmake | 2 + 3 files changed, 499 insertions(+) create mode 100644 SerialPrograms/Source/Integrations/AgentServer/AgentServer_InputSteps.cpp create mode 100644 SerialPrograms/Source/Integrations/AgentServer/AgentServer_InputSteps.h diff --git a/SerialPrograms/Source/Integrations/AgentServer/AgentServer_InputSteps.cpp b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_InputSteps.cpp new file mode 100644 index 0000000000..086a1f865e --- /dev/null +++ b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_InputSteps.cpp @@ -0,0 +1,403 @@ +/* Agent Server: Input Steps + * + * From: https://github.com/PokemonAutomation/ + * + */ + +#include +#include +#include +#include +#include "AgentServer_InputSteps.h" + +namespace PokemonAutomation{ +namespace AgentServer{ + +using nlohmann::json; +using namespace NintendoSwitch; + + + +namespace{ + +// Same tables as BUTTON_BITS / BUTTON_ALIASES / DPAD_POSITIONS in buttons.py. +const std::vector>& button_names(){ + static const std::vector> names{ + {"Y", BUTTON_Y}, + {"B", BUTTON_B}, + {"A", BUTTON_A}, + {"X", BUTTON_X}, + {"L", BUTTON_L}, + {"R", BUTTON_R}, + {"ZL", BUTTON_ZL}, + {"ZR", BUTTON_ZR}, + {"MINUS", BUTTON_MINUS}, + {"PLUS", BUTTON_PLUS}, + {"LCLICK", BUTTON_LCLICK}, + {"RCLICK", BUTTON_RCLICK}, + {"HOME", BUTTON_HOME}, + {"CAPTURE", BUTTON_CAPTURE}, + }; + return names; +} +const std::map& button_aliases(){ + static const std::map aliases{ + {"+", "PLUS"}, + {"START", "PLUS"}, + {"-", "MINUS"}, + {"SELECT", "MINUS"}, + {"L3", "LCLICK"}, + {"LS", "LCLICK"}, + {"LSTICK", "LCLICK"}, + {"R3", "RCLICK"}, + {"RS", "RCLICK"}, + {"RSTICK", "RCLICK"}, + {"SCREENSHOT", "CAPTURE"}, + }; + return aliases; +} + +struct Direction{ + const char* name; + int x; + int y; + DpadPosition dpad; +}; +const std::vector& directions(){ + static const std::vector list{ + {"UP", 0, 1, DPAD_UP}, + {"UP_RIGHT", 1, 1, DPAD_UP_RIGHT}, + {"RIGHT", 1, 0, DPAD_RIGHT}, + {"DOWN_RIGHT", 1, -1, DPAD_DOWN_RIGHT}, + {"DOWN", 0, -1, DPAD_DOWN}, + {"DOWN_LEFT", -1, -1, DPAD_DOWN_LEFT}, + {"LEFT", -1, 0, DPAD_LEFT}, + {"UP_LEFT", -1, 1, DPAD_UP_LEFT}, + }; + return list; +} +const Direction* find_direction(const std::string& name){ + for (const Direction& direction : directions()){ + if (name == direction.name){ + return &direction; + } + } + return nullptr; +} + + +std::string trim(const std::string& text){ + size_t start = text.find_first_not_of(" \t\r\n"); + if (start == std::string::npos){ + return ""; + } + size_t end = text.find_last_not_of(" \t\r\n"); + return text.substr(start, end - start + 1); +} +std::string to_upper(std::string text){ + for (char& ch : text){ + ch = (char)std::toupper((unsigned char)ch); + } + return text; +} + +// "up-right" -> "UP_RIGHT", "UPRIGHT" -> "UP_RIGHT", "DPAD_UP" -> "UP", "a" -> "A" +std::string normalize_name(const std::string& raw){ + std::string name = to_upper(trim(raw)); + for (char& ch : name){ + if (ch == '-' || ch == ' '){ + ch = '_'; + } + } + if (name.starts_with("DPAD_")){ + name = name.substr(5); + } + for (const char* vertical : {"UP", "DOWN"}){ + for (const char* horizontal : {"LEFT", "RIGHT"}){ + std::string v = vertical, h = horizontal; + if (name == v + h || name == h + v || name == h + "_" + v){ + return v + "_" + h; + } + } + } + return name; +} + +// Split a combination into names. A lone "+" or "-" is the PLUS/MINUS button. +void split_buttons(const json& value, std::vector& out){ + if (value.is_null()){ + return; + } + if (value.is_array()){ + for (const json& item : value){ + split_buttons(item, out); + } + return; + } + if (!value.is_string()){ + throw InputError("Buttons must be a string like \"A\" or \"L+R\", or a list of names."); + } + std::string text = trim(value.get()); + if (text == "+" || text == "-"){ + out.emplace_back(text); + return; + } + for (char& ch : text){ + if (ch == ','){ + ch = '+'; + } + } + std::stringstream ss(text); + std::string part; + while (std::getline(ss, part, '+')){ + if (!trim(part).empty()){ + out.emplace_back(part); + } + } +} + +std::string valid_names_message(){ + std::string buttons, dpad; + for (const auto& item : button_names()){ + buttons += (buttons.empty() ? "" : ", ") + item.first; + } + for (const Direction& direction : directions()){ + dpad += (dpad.empty() ? "" : ", ") + std::string(direction.name); + } + return "Valid buttons: " + buttons + ", d-pad: " + dpad + "."; +} + +std::string number_text(double x){ + std::ostringstream ss; + ss << x; + return ss.str(); +} + +// Read a non-negative-or-not integer field. Accepts 5 and 5.0. +int64_t read_integer(const json& object, const char* key, int64_t default_value){ + auto iter = object.find(key); + if (iter == object.end() || iter->is_null()){ + return default_value; + } + if (iter->is_number_integer()){ + return iter->get(); + } + if (iter->is_number_float()){ + double x = iter->get(); + if (std::isfinite(x) && x == std::floor(x)){ + return (int64_t)x; + } + } + throw InputError(std::string(key) + " must be an integer."); +} + +bool is_empty_buttons(const json& value){ + return value.is_null() + || (value.is_string() && value.get().empty()) + || (value.is_array() && value.empty()); +} + +} + + + +ParsedButtons parse_buttons(const json& value){ + std::vector names; + split_buttons(value, names); + + ParsedButtons ret; + int dx = 0, dy = 0; + std::vector seen_dpad; + for (const std::string& raw : names){ + std::string name = trim(raw); + std::string key = (name == "+" || name == "-") ? name : normalize_name(name); + auto alias = button_aliases().find(key); + if (alias != button_aliases().end()){ + key = alias->second; + } + + bool found = false; + for (const auto& item : button_names()){ + if (item.first == key){ + ret.buttons |= item.second; + found = true; + break; + } + } + if (found){ + continue; + } + + const Direction* direction = find_direction(key); + if (direction != nullptr){ + if ((direction->x && dx && direction->x != dx) || (direction->y && dy && direction->y != dy)){ + std::string list; + for (const std::string& s : seen_dpad){ + list += "\"" + s + "\", "; + } + throw InputError("Contradictory d-pad directions: [" + list + "\"" + name + "\"]"); + } + dx = direction->x ? direction->x : dx; + dy = direction->y ? direction->y : dy; + seen_dpad.emplace_back(name); + continue; + } + + throw InputError("Unknown button \"" + name + "\". " + valid_names_message()); + } + + if (dx || dy){ + for (const Direction& direction : directions()){ + if (direction.x == dx && direction.y == dy){ + ret.dpad = direction.dpad; + } + } + } + return ret; +} + + +JoystickPosition parse_stick(const json& value){ + if (value.is_null()){ + return {0, 0}; + } + if (value.is_string()){ + std::string key = normalize_name(value.get()); + if (key == "NEUTRAL" || key == "CENTER" || key == "NONE"){ + return {0, 0}; + } + const Direction* direction = find_direction(key); + if (direction == nullptr){ + std::string names; + for (const Direction& d : directions()){ + std::string lower = d.name; + for (char& ch : lower){ + ch = (char)std::tolower((unsigned char)ch); + } + names += (names.empty() ? "" : ", ") + lower; + } + throw InputError( + "Unknown stick direction \"" + value.get() + "\". Use one of " + + names + ", or an [x, y] pair." + ); + } + double length = std::hypot((double)direction->x, (double)direction->y); + return {direction->x / length, direction->y / length}; + } + if (value.is_array()){ + if (value.size() != 2 || !value[0].is_number() || !value[1].is_number()){ + throw InputError("A stick position needs exactly two numbers [x, y], got " + value.dump() + "."); + } + double x = value[0].get(); + double y = value[1].get(); + if (!(x >= -1.0 && x <= 1.0 && y >= -1.0 && y <= 1.0)){ + throw InputError( + "Stick coordinates must be within [-1, 1], got (" + number_text(x) + ", " + number_text(y) + ")." + ); + } + return {x, y}; + } + throw InputError("A stick position must be a direction name or an [x, y] pair."); +} + + +InputStep parse_step(const json& object){ + if (!object.is_object()){ + throw InputError("Each step must be an object, e.g. {\"buttons\": \"A\"}."); + } + static const char* FIELDS[] = { + "buttons", "left_stick", "right_stick", "hold_ms", "release_ms", "repeat", "wait_ms" + }; + std::string unknown; + for (auto iter = object.begin(); iter != object.end(); ++iter){ + bool known = false; + for (const char* field : FIELDS){ + known |= iter.key() == field; + } + if (!known){ + unknown += (unknown.empty() ? "'" : ", '") + iter.key() + "'"; + } + } + if (!unknown.empty()){ + throw InputError("Unknown input step field(s): [" + unknown + "]"); + } + + int64_t hold = read_integer(object, "hold_ms", 80); + int64_t release = read_integer(object, "release_ms", 80); + int64_t repeat = read_integer(object, "repeat", 1); + int64_t wait = read_integer(object, "wait_ms", 0); + if (hold < 0 || release < 0 || wait < 0){ + throw InputError("Durations must not be negative."); + } + if (repeat < 1){ + throw InputError("repeat must be at least 1."); + } + + InputStep step; + step.hold_ms = (uint64_t)hold; + step.release_ms = (uint64_t)release; + step.repeat = (uint64_t)repeat; + step.wait_ms = (uint64_t)wait; + + json buttons = object.value("buttons", json()); + json left = object.value("left_stick", json()); + json right = object.value("right_stick", json()); + step.wait_only = is_empty_buttons(buttons) && left.is_null() && right.is_null(); + if (step.wait_only){ + return step; + } + if (hold == 0){ + throw InputError("hold_ms must be positive for a step that presses something."); + } + step.pressed = parse_buttons(buttons); + if (!left.is_null()){ + step.left_stick = parse_stick(left); + } + if (!right.is_null()){ + step.right_stick = parse_stick(right); + } + return step; +} + + +uint64_t InputStep::duration_ms() const{ + if (wait_only){ + return wait_ms; + } + return repeat * (hold_ms + release_ms) + wait_ms; +} + +std::string InputStep::describe() const{ + if (wait_only){ + return "wait " + std::to_string(wait_ms) + "ms"; + } + std::string ret; + if (pressed.has_buttons()){ + ret += button_to_string(pressed.buttons); + } + if (pressed.has_dpad()){ + ret += (ret.empty() ? "" : " + ") + std::string("d-pad ") + dpad_to_string(pressed.dpad); + } + auto stick_text = [](const char* side, const JoystickPosition& p){ + return std::string(side) + " stick (" + number_text(p.x) + ", " + number_text(p.y) + ")"; + }; + if (left_stick){ + ret += (ret.empty() ? "" : " + ") + stick_text("left", *left_stick); + } + if (right_stick){ + ret += (ret.empty() ? "" : " + ") + stick_text("right", *right_stick); + } + if (ret.empty()){ + ret = "neutral"; + } + ret += " " + std::to_string(hold_ms) + "ms"; + if (repeat > 1){ + ret += " x" + std::to_string(repeat); + } + return ret; +} + + + +} +} diff --git a/SerialPrograms/Source/Integrations/AgentServer/AgentServer_InputSteps.h b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_InputSteps.h new file mode 100644 index 0000000000..5d7bff160d --- /dev/null +++ b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_InputSteps.h @@ -0,0 +1,94 @@ +/* Agent Server: Input Steps + * + * From: https://github.com/PokemonAutomation/ + * + * Parses the input vocabulary of AgentTools.json (button names, d-pad directions, + * stick positions, input steps) into Nintendo Switch controller values. + * + * This is the C++ twin of pokemon_automation/buttons.py and `InputStep` in + * pokemon_automation/controller.py. Both are tested against the shared cases in + * AgentInputTestCases.json, so an agent's inputs mean the same thing whichever + * server it talks to. Keep them in sync. + * + * Vocabulary: + * - Buttons: A B X Y L R ZL ZR PLUS ("+", START) MINUS ("-", SELECT) HOME + * CAPTURE LCLICK (L3, LS) RCLICK (R3, RS). Case-insensitive. + * - D-pad: UP DOWN LEFT RIGHT and diagonals (UP_RIGHT, "up-right", UPRIGHT, + * DPAD_UP...). "UP" and "RIGHT" together also mean up-right. + * - Combinations: "L+R", "b, x", or a JSON array ["ZL", "A"]. + * - Sticks: a direction name ("up", "down_left", "neutral") or [x, y] in [-1, 1], + * +y = up. Diagonal names are normalized to length 1. + */ + +#ifndef PokemonAutomation_AgentServer_InputSteps_H +#define PokemonAutomation_AgentServer_InputSteps_H + +#include +#include +#include +#include +#include +#include "3rdParty-Core/nlohmann/json.hpp" +#include "Controllers/Joystick.h" +#include "NintendoSwitch/Controllers/NintendoSwitch_ControllerButtons.h" + +namespace PokemonAutomation{ +namespace AgentServer{ + + +// Thrown for input the agent got wrong. The message is meant for the agent. +class InputError : public std::runtime_error{ +public: + using std::runtime_error::runtime_error; +}; + + +struct ParsedButtons{ + NintendoSwitch::Button buttons = NintendoSwitch::BUTTON_NONE; + NintendoSwitch::DpadPosition dpad = NintendoSwitch::DPAD_NONE; + + bool has_buttons() const{ return buttons != NintendoSwitch::BUTTON_NONE; } + bool has_dpad() const{ return dpad != NintendoSwitch::DPAD_NONE; } +}; + +// Parse a button combination: a string ("A", "L+R", "zl, up") or a JSON array of +// strings. null or "" means nothing pressed. +// Throws InputError for unknown names or contradictory d-pad directions ("up+down"). +ParsedButtons parse_buttons(const nlohmann::json& value); + +// Parse a joystick position: a direction name or an [x, y] array. +// Throws InputError for unknown names, wrong array sizes or values outside [-1, 1]. +JoystickPosition parse_stick(const nlohmann::json& value); + + +// One step of an input sequence (a `run_inputs` step). +// +// A step either waits (nothing pressed: only `wait_ms`), or holds the given +// buttons, d-pad and sticks together for `hold_ms`, releases everything for +// `release_ms`, and repeats that `repeat` times, then waits `wait_ms`. +struct InputStep{ + ParsedButtons pressed; + std::optional left_stick; + std::optional right_stick; + uint64_t hold_ms = 80; + uint64_t release_ms = 80; + uint64_t repeat = 1; + uint64_t wait_ms = 0; + bool wait_only = false; // nothing to press: the step is just `wait_ms` + + // Total time the step occupies on the controller. + uint64_t duration_ms() const; + // Short description for logs, e.g. "A x3" or "left stick (0, 1) 2000ms". + std::string describe() const; +}; + +// Parse a step object, e.g. {"buttons": "A", "hold_ms": 100}. +// Throws InputError for unknown fields, negative durations, repeat < 1, or +// hold_ms = 0 on a step that presses something. +InputStep parse_step(const nlohmann::json& object); + + + +} +} +#endif diff --git a/SerialPrograms/cmake/SourceFiles.cmake b/SerialPrograms/cmake/SourceFiles.cmake index c5362c3c6b..75c0db60d2 100644 --- a/SerialPrograms/cmake/SourceFiles.cmake +++ b/SerialPrograms/cmake/SourceFiles.cmake @@ -906,6 +906,8 @@ file(GLOB LIBRARY_SOURCES Source/Controllers/SerialPort/SerialPortPollerQt.h Source/Integrations/AgentServer/AgentServer_HttpServer.cpp Source/Integrations/AgentServer/AgentServer_HttpServer.h + Source/Integrations/AgentServer/AgentServer_InputSteps.cpp + Source/Integrations/AgentServer/AgentServer_InputSteps.h Source/Integrations/AgentServer/AgentServer_ToolDefinitions.cpp Source/Integrations/AgentServer/AgentServer_ToolDefinitions.h Source/Integrations/DiscordIntegrationSettings.cpp From 8b95e1e66e910874148000a5d0cfe8acfabebc79 Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:14:29 -0700 Subject: [PATCH 14/16] AgentServer: add the MCP protocol server MCP over Streamable HTTP at POST /mcp: JSON-RPC 2.0 with the initialize handshake (protocol versions 2024-11-05 to 2025-11-25; newer clients' probe gets "method not found" and falls back to it), ping, tools/list and tools/call, sessions, and plain JSON responses. Security: an optional bearer token, and DNS-rebinding protection (Origin and Host checks). The tools themselves are implemented by an McpToolHandler. Co-Authored-By: Claude Opus 5.5 --- .../AgentServer/AgentServer_McpServer.cpp | 419 ++++++++++++++++++ .../AgentServer/AgentServer_McpServer.h | 128 ++++++ SerialPrograms/cmake/SourceFiles.cmake | 2 + 3 files changed, 549 insertions(+) create mode 100644 SerialPrograms/Source/Integrations/AgentServer/AgentServer_McpServer.cpp create mode 100644 SerialPrograms/Source/Integrations/AgentServer/AgentServer_McpServer.h diff --git a/SerialPrograms/Source/Integrations/AgentServer/AgentServer_McpServer.cpp b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_McpServer.cpp new file mode 100644 index 0000000000..b95b29d8e8 --- /dev/null +++ b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_McpServer.cpp @@ -0,0 +1,419 @@ +/* Agent Server: MCP Server + * + * From: https://github.com/PokemonAutomation/ + * + */ + +#include +#include +#include +#include "Common/Cpp/Color.h" +#include "Common/Cpp/Logging/AbstractLogger.h" +#include "AgentServer_McpServer.h" + +namespace PokemonAutomation{ +namespace AgentServer{ + +using nlohmann::json; + + + +McpContent McpContent::make_text(std::string text){ + McpContent ret; + ret.type = Type::TEXT; + ret.text = std::move(text); + return ret; +} +McpContent McpContent::make_image(std::string base64_data, std::string mime_type){ + McpContent ret; + ret.type = Type::IMAGE; + ret.base64_data = std::move(base64_data); + ret.mime_type = std::move(mime_type); + return ret; +} +McpToolResult McpToolResult::text(std::string message){ + McpToolResult ret; + ret.content.emplace_back(McpContent::make_text(std::move(message))); + return ret; +} +McpToolResult McpToolResult::error(std::string message){ + McpToolResult ret = text(std::move(message)); + ret.is_error = true; + return ret; +} + + + +namespace{ + +// JSON-RPC 2.0 error codes. +const int PARSE_ERROR = -32700; +const int INVALID_REQUEST = -32600; +const int METHOD_NOT_FOUND = -32601; +const int INVALID_PARAMS = -32602; + +json rpc_error(const json& id, int code, const std::string& message){ + return { + {"jsonrpc", "2.0"}, + {"id", id}, + {"error", {{"code", code}, {"message", message}}}, + }; +} +json rpc_result(const json& id, json result){ + return { + {"jsonrpc", "2.0"}, + {"id", id}, + {"result", std::move(result)}, + }; +} + +HttpResponse json_response(int status, const json& body){ + return HttpResponse::json(status, body.dump()); +} + +std::string to_lower(std::string text){ + for (char& ch : text){ + ch = (char)std::tolower((unsigned char)ch); + } + return text; +} + +// "localhost:8765" -> "localhost", "[::1]:8765" -> "[::1]", "127.0.0.1" -> "127.0.0.1" +std::string host_without_port(const std::string& host){ + if (host.starts_with("[")){ + size_t end = host.find(']'); + return end == std::string::npos ? host : host.substr(0, end + 1); + } + size_t colon = host.find(':'); + return colon == std::string::npos ? host : host.substr(0, colon); +} +bool is_local_host_name(const std::string& host){ + std::string name = to_lower(host_without_port(host)); + return name == "localhost" || name == "127.0.0.1" || name == "[::1]"; +} + +// Compare without an early exit, so response timing doesn't reveal the token. +bool constant_time_equals(const std::string& a, const std::string& b){ + if (a.size() != b.size()){ + return false; + } + unsigned char diff = 0; + for (size_t c = 0; c < a.size(); c++){ + diff |= (unsigned char)(a[c] ^ b[c]); + } + return diff == 0; +} + +std::string random_hex(size_t bytes){ + static std::mutex lock; + static std::mt19937_64 rng(std::random_device{}()); + std::lock_guard lg(lock); + const char* digits = "0123456789abcdef"; + std::string ret; + for (size_t c = 0; c < bytes; c++){ + uint8_t byte = (uint8_t)(rng() & 0xff); + ret += digits[byte >> 4]; + ret += digits[byte & 15]; + } + return ret; +} + +json to_json(const McpToolResult& result){ + json content = json::array(); + for (const McpContent& block : result.content){ + if (block.type == McpContent::Type::IMAGE){ + content.push_back({ + {"type", "image"}, + {"data", block.base64_data}, + {"mimeType", block.mime_type}, + }); + }else{ + content.push_back({ + {"type", "text"}, + {"text", block.text}, + }); + } + } + return { + {"content", std::move(content)}, + {"isError", result.is_error}, + }; +} + +const size_t MAX_SESSIONS = 64; + +} + + + +const std::vector& McpServer::supported_protocol_versions(){ + static const std::vector versions{ + "2025-11-25", + "2025-06-18", + "2025-03-26", + "2024-11-05", + }; + return versions; +} + + +McpServer::McpServer( + Logger& logger, + const AgentToolDefinitions& definitions, + McpToolHandler& handler, + McpServerConfig config +) + : m_logger(logger) + , m_definitions(definitions) + , m_handler(handler) + , m_config(std::move(config)) +{} + +size_t McpServer::active_sessions() const{ + std::lock_guard lg(m_lock); + return m_sessions.size(); +} + + +std::optional McpServer::check_access(const HttpRequest& request) const{ + // A web page can make the browser send requests here. Browsers always send + // Origin with those; MCP clients generally don't send one at all. + const std::string& origin = request.header("origin"); + if (!origin.empty()){ + size_t scheme_end = origin.find("://"); + std::string host = scheme_end == std::string::npos ? "" : origin.substr(scheme_end + 3); + if (!is_local_host_name(host)){ + m_logger.log("[AgentServer] Rejected request from web origin: " + origin, COLOR_RED); + return HttpResponse::text(403, "Requests from web pages are not allowed.\n"); + } + } + + // DNS rebinding: a remote site's domain resolving to 127.0.0.1 still carries + // that domain in the Host header. + const std::string& host = request.header("host"); + if (m_config.localhost_only && !host.empty() && !is_local_host_name(host)){ + m_logger.log("[AgentServer] Rejected request for host: " + host, COLOR_RED); + return HttpResponse::text(403, "Invalid Host header.\n"); + } + + if (!m_config.access_token.empty()){ + const std::string& authorization = request.header("authorization"); + if (!constant_time_equals(authorization, "Bearer " + m_config.access_token)){ + m_logger.log("[AgentServer] Rejected request without a valid access token from " + request.peer_address, COLOR_RED); + return json_response(401, rpc_error(nullptr, INVALID_REQUEST, + "Missing or wrong access token. Send \"Authorization: Bearer \" " + "with the token shown in SerialPrograms' AI Agent Server program." + )); + } + } + return std::nullopt; +} + + +HttpResponse McpServer::handle(const HttpRequest& request){ + if (request.path != "/mcp" && request.path != "/mcp/"){ + return HttpResponse::text(404, "Not found. The MCP endpoint is /mcp.\n"); + } + if (std::optional denied = check_access(request)){ + return *denied; + } + + const std::string& session_id = request.header("mcp-session-id"); + + if (request.method == "DELETE"){ + std::lock_guard lg(m_lock); + if (session_id.empty() || m_sessions.erase(session_id) == 0){ + return HttpResponse::text(404, "Session not found.\n"); + } + return HttpResponse::text(200, "Session ended.\n"); + } + if (request.method != "POST"){ + // GET would open a server-to-client event stream, which this server doesn't offer. + HttpResponse response = HttpResponse::text(405, "Use POST.\n"); + response.headers.emplace_back("Allow", "POST, DELETE"); + return response; + } + + std::string content_type = to_lower(request.header("content-type")); + if (!content_type.empty() && !content_type.starts_with("application/json")){ + return HttpResponse::text(415, "Content-Type must be application/json.\n"); + } + + if (!session_id.empty()){ + std::lock_guard lg(m_lock); + if (m_sessions.find(session_id) == m_sessions.end()){ + // The spec's signal for "session expired; initialize again". + return json_response(404, rpc_error(nullptr, INVALID_REQUEST, "Session not found.")); + } + } + + json body; + try{ + body = json::parse(request.body); + }catch (const json::parse_error&){ + return json_response(400, rpc_error(nullptr, PARSE_ERROR, "Parse error: the body is not valid JSON.")); + } + + std::string new_session_id; + json reply; + if (body.is_array()){ + // JSON-RPC batch (protocol 2025-03-26). + if (body.empty()){ + return json_response(400, rpc_error(nullptr, INVALID_REQUEST, "Empty batch.")); + } + reply = json::array(); + for (const json& message : body){ + json response = handle_message(message, new_session_id); + if (!response.is_null()){ + reply.push_back(std::move(response)); + } + } + if (reply.empty()){ + reply = nullptr; + } + }else{ + reply = handle_message(body, new_session_id); + } + + HttpResponse response; + if (reply.is_null()){ + response.status = 202; // only notifications/responses: nothing to return + }else{ + response = json_response(200, reply); + } + if (!new_session_id.empty()){ + response.headers.emplace_back("Mcp-Session-Id", new_session_id); + } + return response; +} + + +json McpServer::handle_message(const json& message, std::string& new_session_id){ + if (!message.is_object() || message.value("jsonrpc", "") != "2.0"){ + return rpc_error(nullptr, INVALID_REQUEST, "Not a JSON-RPC 2.0 message."); + } + auto method_iter = message.find("method"); + if (method_iter == message.end()){ + return nullptr; // a response to a server request; this server sends none + } + if (!method_iter->is_string()){ + return rpc_error(message.value("id", json()), INVALID_REQUEST, "\"method\" must be a string."); + } + const std::string method = method_iter->get(); + + if (!message.contains("id")){ + // Notification, e.g. notifications/initialized or notifications/cancelled. + return nullptr; + } + const json& id = message["id"]; + json params = message.value("params", json::object()); + + try{ + if (method == "initialize"){ + return rpc_result(id, handle_initialize(params, new_session_id)); + } + if (method == "ping"){ + return rpc_result(id, json::object()); + } + if (method == "tools/list"){ + return rpc_result(id, {{"tools", m_definitions.tools_list()}}); + } + if (method == "tools/call"){ + if (!params.is_object() || !params.contains("name") || !params["name"].is_string()){ + return rpc_error(id, INVALID_PARAMS, "tools/call needs a \"name\"."); + } + std::string name = params["name"].get(); + if (m_definitions.find(name) == nullptr){ + return rpc_error(id, INVALID_PARAMS, "Unknown tool: " + name); + } + return rpc_result(id, handle_tools_call(params)); + } + }catch (const std::exception& e){ + m_logger.log(std::string("[AgentServer] Error handling ") + method + ": " + e.what(), COLOR_RED); + return rpc_error(id, -32603, std::string("Internal error: ") + e.what()); + } + + // Includes "server/discover" (protocol 2026-07-28): "method not found" tells + // newer clients to fall back to the initialize handshake. + return rpc_error(id, METHOD_NOT_FOUND, "Method not found: " + method); +} + + +json McpServer::handle_initialize(const json& params, std::string& new_session_id){ + std::string requested = params.is_object() ? params.value("protocolVersion", "") : ""; + const std::vector& supported = supported_protocol_versions(); + std::string version = supported.front(); + for (const std::string& v : supported){ + if (v == requested){ + version = v; + } + } + + std::string client = "unknown client"; + if (params.is_object() && params.contains("clientInfo") && params["clientInfo"].is_object()){ + const json& info = params["clientInfo"]; + client = info.value("name", "unknown client") + " " + info.value("version", ""); + } + + new_session_id = random_hex(16); + { + std::lock_guard lg(m_lock); + m_sessions[new_session_id] = ++m_session_counter; + // Clients that vanish never send DELETE; drop the oldest sessions. + while (m_sessions.size() > MAX_SESSIONS){ + auto oldest = m_sessions.begin(); + for (auto iter = m_sessions.begin(); iter != m_sessions.end(); ++iter){ + if (iter->second < oldest->second){ + oldest = iter; + } + } + m_sessions.erase(oldest); + } + } + m_logger.log("[AgentServer] Agent connected: " + client + " (protocol " + version + ")", COLOR_BLUE); + + return { + {"protocolVersion", version}, + {"capabilities", {{"tools", {{"listChanged", false}}}}}, + {"serverInfo", { + {"name", m_definitions.server_name()}, + {"version", m_config.server_version}, + }}, + {"instructions", m_definitions.instructions()}, + }; +} + + +json McpServer::handle_tools_call(const json& params){ + const AgentToolDefinition& tool = *m_definitions.find(params["name"].get()); + json arguments = params.value("arguments", json::object()); + if (arguments.is_null()){ + arguments = json::object(); + } + + std::string error = m_definitions.validate_arguments(tool, arguments); + if (!error.empty()){ + // A tool error (not a protocol error), so the agent can read it and retry. + return to_json(McpToolResult::error("Invalid arguments for " + tool.name + ": " + error)); + } + arguments = AgentToolDefinitions::with_defaults(tool, arguments); + + auto start = std::chrono::steady_clock::now(); + McpToolResult result; + try{ + result = m_handler.call_tool(tool.name, arguments); + }catch (const std::exception& e){ + result = McpToolResult::error(e.what()); + } + auto millis = std::chrono::duration_cast(std::chrono::steady_clock::now() - start).count(); + m_logger.log( + "[AgentServer] " + tool.name + (result.is_error ? " failed" : " done") + " (" + std::to_string(millis) + " ms)", + result.is_error ? COLOR_RED : COLOR_DARKGREEN + ); + return to_json(result); +} + + + +} +} diff --git a/SerialPrograms/Source/Integrations/AgentServer/AgentServer_McpServer.h b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_McpServer.h new file mode 100644 index 0000000000..96deb49e91 --- /dev/null +++ b/SerialPrograms/Source/Integrations/AgentServer/AgentServer_McpServer.h @@ -0,0 +1,128 @@ +/* Agent Server: MCP Server + * + * From: https://github.com/PokemonAutomation/ + * + * The Model Context Protocol (MCP) over "Streamable HTTP", served at POST /mcp. + * Any MCP client (Claude Code, Claude Desktop, Codex, ...) can connect to it and + * call the tools in AgentTools.json. This class does the protocol; the tools + * themselves are implemented by a `McpToolHandler` (the "AI Agent Server" program). + * + * Protocol support: + * - The initialize handshake, protocol versions 2024-11-05 through 2025-11-25. + * Newer clients first probe `server/discover` (protocol 2026-07-28); that gets + * "method not found", which tells them to fall back to the handshake. + * - Requests: initialize, ping, tools/list, tools/call. Notifications are accepted + * and ignored. Responses are plain JSON (no server-sent-event streams). + * - GET /mcp (a server-to-client event stream) is not offered: 405. + * - DELETE /mcp ends the session. + * + * Security (the server can press any button on the user's console): + * - Bearer token: when `McpServerConfig::access_token` is set, every request needs + * "Authorization: Bearer ". + * - DNS-rebinding protection: requests from web pages (with an Origin header) are + * rejected unless the origin is this machine, and in localhost-only mode the Host + * header must name this machine. + */ + +#ifndef PokemonAutomation_AgentServer_McpServer_H +#define PokemonAutomation_AgentServer_McpServer_H + +#include +#include +#include +#include +#include +#include "3rdParty-Core/nlohmann/json.hpp" +#include "AgentServer_HttpServer.h" +#include "AgentServer_ToolDefinitions.h" + +namespace PokemonAutomation{ + class Logger; +namespace AgentServer{ + + +// One content block of a tool result. +struct McpContent{ + enum class Type{ TEXT, IMAGE }; + Type type = Type::TEXT; + std::string text; // TEXT + std::string base64_data; // IMAGE + std::string mime_type; // IMAGE, e.g. "image/jpeg" + + static McpContent make_text(std::string text); + static McpContent make_image(std::string base64_data, std::string mime_type); +}; + +// What a tool returns. `is_error` marks a failure the agent should see and react +// to (bad arguments, the user has taken control, no video, ...). +struct McpToolResult{ + std::vector content; + bool is_error = false; + + static McpToolResult text(std::string message); + static McpToolResult error(std::string message); +}; + + +// Implements the tools. `call_tool()` runs on an HTTP worker thread and may block +// while inputs execute. It is only called for tools in the definitions, with +// arguments that passed schema validation and have top-level defaults filled in. +// Exceptions are caught and reported to the agent as tool errors. +class McpToolHandler{ +public: + virtual ~McpToolHandler() = default; + virtual McpToolResult call_tool(const std::string& name, const nlohmann::json& arguments) = 0; +}; + + +struct McpServerConfig{ + std::string server_version; // reported to clients, e.g. the app version + std::string access_token; // empty = no authentication + bool localhost_only = true; // check the Host header (see above) +}; + + +class McpServer{ +public: + McpServer( + Logger& logger, + const AgentToolDefinitions& definitions, + McpToolHandler& handler, + McpServerConfig config + ); + + // Handle one HTTP request. Thread-safe; use as the `HttpServer` handler. + HttpResponse handle(const HttpRequest& request); + + // Protocol versions accepted in initialize, newest first. + static const std::vector& supported_protocol_versions(); + + size_t active_sessions() const; + +private: + // Returns a response if the request is not allowed, else nothing. + std::optional check_access(const HttpRequest& request) const; + + // Handle one JSON-RPC message. Returns the response, or null for a + // notification or response (which get no reply). + nlohmann::json handle_message(const nlohmann::json& message, std::string& new_session_id); + + nlohmann::json handle_initialize(const nlohmann::json& params, std::string& new_session_id); + nlohmann::json handle_tools_call(const nlohmann::json& params); + +private: + Logger& m_logger; + const AgentToolDefinitions& m_definitions; + McpToolHandler& m_handler; + const McpServerConfig m_config; + + mutable std::mutex m_lock; + std::map m_sessions; // session ID -> creation counter + uint64_t m_session_counter = 0; +}; + + + +} +} +#endif diff --git a/SerialPrograms/cmake/SourceFiles.cmake b/SerialPrograms/cmake/SourceFiles.cmake index 75c0db60d2..d283afc08c 100644 --- a/SerialPrograms/cmake/SourceFiles.cmake +++ b/SerialPrograms/cmake/SourceFiles.cmake @@ -908,6 +908,8 @@ file(GLOB LIBRARY_SOURCES Source/Integrations/AgentServer/AgentServer_HttpServer.h Source/Integrations/AgentServer/AgentServer_InputSteps.cpp Source/Integrations/AgentServer/AgentServer_InputSteps.h + Source/Integrations/AgentServer/AgentServer_McpServer.cpp + Source/Integrations/AgentServer/AgentServer_McpServer.h Source/Integrations/AgentServer/AgentServer_ToolDefinitions.cpp Source/Integrations/AgentServer/AgentServer_ToolDefinitions.h Source/Integrations/DiscordIntegrationSettings.cpp From 3b95c24a3989a634beb2469ad0acc9a5ddbf3a4b Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:14:30 -0700 Subject: [PATCH 15/16] GameConsole: tell listeners about the user's controller input - ConsoleSystemSession::Listener gets on_controller_input(), called after the user's keyboard input was sent to the console's controllers (not for input suppressed because the console isn't focused or its controllers are locked). It runs outside the session lock; the forwarding code itself is unchanged. - ConsoleHandle::system_session() gives programs access to their session, so they can register such a listener. Used by the AI Agent Server program to notice when the user takes over. Co-Authored-By: Claude Opus 5.5 --- .../Source/GameConsole/ConsoleHandle.cpp | 3 ++ .../Source/GameConsole/ConsoleHandle.h | 4 +++ .../Framework/ConsoleSystemSession.cpp | 35 ++++++++++--------- .../Framework/ConsoleSystemSession.h | 6 ++++ 4 files changed, 32 insertions(+), 16 deletions(-) diff --git a/SerialPrograms/Source/GameConsole/ConsoleHandle.cpp b/SerialPrograms/Source/GameConsole/ConsoleHandle.cpp index 54633abbfd..702b4c1237 100644 --- a/SerialPrograms/Source/GameConsole/ConsoleHandle.cpp +++ b/SerialPrograms/Source/GameConsole/ConsoleHandle.cpp @@ -83,6 +83,9 @@ ConsoleHandle::ConsoleHandle(ConsoleSystemSession& session) size_t ConsoleHandle::index() const{ return m_data->m_index; } +ConsoleSystemSession& ConsoleHandle::system_session(){ + return m_data->m_session; +} size_t ConsoleHandle::controllers() const{ return m_data->m_session.controllers(); } diff --git a/SerialPrograms/Source/GameConsole/ConsoleHandle.h b/SerialPrograms/Source/GameConsole/ConsoleHandle.h index 08d226fea3..ceed9c07cd 100644 --- a/SerialPrograms/Source/GameConsole/ConsoleHandle.h +++ b/SerialPrograms/Source/GameConsole/ConsoleHandle.h @@ -32,6 +32,10 @@ class ConsoleHandle : public VideoStream{ size_t index() const; + // The console session this handle belongs to. Programs use it to listen for + // session events, e.g. the user's keyboard input (see ConsoleSystemSession::Listener). + ConsoleSystemSession& system_session(); + operator Logger&(){ return logger(); } operator VideoFeed&(){ return video(); } operator VideoOverlay&(){ return overlay(); } diff --git a/SerialPrograms/Source/GameConsole/Framework/ConsoleSystemSession.cpp b/SerialPrograms/Source/GameConsole/Framework/ConsoleSystemSession.cpp index 9bed1fa736..02e56c6179 100644 --- a/SerialPrograms/Source/GameConsole/Framework/ConsoleSystemSession.cpp +++ b/SerialPrograms/Source/GameConsole/Framework/ConsoleSystemSession.cpp @@ -236,25 +236,28 @@ void ConsoleSystemSession::on_focus_out(){ m_listeners.run_method(&Listener::on_input_status_change, status); } void ConsoleSystemSession::run_controller_input(ControllerInputState& state){ - std::lock_guard lg(m_lock); - if (!m_focused){ - m_logger.log("Keyboard Command Suppressed: Not in focus.", COLOR_RED); - return; - } - if (!m_lock_controllers_reason.empty() && !allow_commands_while_locked()){ - m_logger.log("Keyboard Command Suppressed: " + m_lock_controllers_reason, COLOR_RED); - return; - } + { + std::lock_guard lg(m_lock); + if (!m_focused){ + m_logger.log("Keyboard Command Suppressed: Not in focus.", COLOR_RED); + return; + } + if (!m_lock_controllers_reason.empty() && !allow_commands_while_locked()){ + m_logger.log("Keyboard Command Suppressed: " + m_lock_controllers_reason, COLOR_RED); + return; + } -// cout << "ConsoleSystemSession::run_controller_input()" << endl; - for (ControllerEntry& controller : m_controllers){ - std::string error = controller.session.try_run([&](AbstractController& controller){ - controller.run_controller_input(state); - }); - if (!error.empty()){ - controller.session.logger().log("Keyboard Command Failed: " + error, COLOR_RED); +// cout << "ConsoleSystemSession::run_controller_input()" << endl; + for (ControllerEntry& controller : m_controllers){ + std::string error = controller.session.try_run([&](AbstractController& controller){ + controller.run_controller_input(state); + }); + if (!error.empty()){ + controller.session.logger().log("Keyboard Command Failed: " + error, COLOR_RED); + } } } + m_listeners.run_method(&Listener::on_controller_input, state); } diff --git a/SerialPrograms/Source/GameConsole/Framework/ConsoleSystemSession.h b/SerialPrograms/Source/GameConsole/Framework/ConsoleSystemSession.h index db00c25105..ea37d5ef09 100644 --- a/SerialPrograms/Source/GameConsole/Framework/ConsoleSystemSession.h +++ b/SerialPrograms/Source/GameConsole/Framework/ConsoleSystemSession.h @@ -52,6 +52,12 @@ class ConsoleSystemSession final virtual void on_input_status_change(const std::string& status){} virtual void on_lock_controllers(){} virtual void on_unlock_controllers(){} + + // Called after the user's keyboard (or other controller input) was sent to + // this console's controllers, e.g. so a program can notice that the user is + // steering by hand. Not called for input that was suppressed (console not in + // focus, or controllers locked). Runs on the UI thread; keep it quick. + virtual void on_controller_input(const ControllerInputState& state){} }; void add_listener(Listener& listener){ From 9e22f986f380c19eff461047ecbcdef451a16cbd Mon Sep 17 00:00:00 2001 From: Gin <> Date: Mon, 5 Oct 2026 23:14:30 -0700 Subject: [PATCH 16/16] ML: add the AI Agent Server program A program in the ML tab (developer mode only) that hosts the MCP server, so AI agents (Claude Code, Codex, ...) can control the Switch through SerialPrograms while the user watches. Agent inputs run on the program's controller context; screenshots come from its video feed; read_text uses the app's Tesseract OCR. Pressing a key hands control to the user: running agent inputs are interrupted and new ones refused until the user clicks "Return control to agent" (or after an optional idle time). Stop shuts the server down and releases the controller. Co-Authored-By: Claude Opus 5.5 --- .../Source/Integrations/AgentServer/README.md | 52 + SerialPrograms/Source/ML/ML_Panels.cpp | 2 + .../Source/ML/Programs/ML_AgentServer.cpp | 942 ++++++++++++++++++ .../Source/ML/Programs/ML_AgentServer.h | 96 ++ .../Source/PythonBindings/README.md | 6 +- SerialPrograms/cmake/SourceFiles.cmake | 2 + 6 files changed, 1099 insertions(+), 1 deletion(-) create mode 100644 SerialPrograms/Source/Integrations/AgentServer/README.md create mode 100644 SerialPrograms/Source/ML/Programs/ML_AgentServer.cpp create mode 100644 SerialPrograms/Source/ML/Programs/ML_AgentServer.h diff --git a/SerialPrograms/Source/Integrations/AgentServer/README.md b/SerialPrograms/Source/Integrations/AgentServer/README.md new file mode 100644 index 0000000000..03ac080a91 --- /dev/null +++ b/SerialPrograms/Source/Integrations/AgentServer/README.md @@ -0,0 +1,52 @@ +# AI Agent Server (MCP) + +Lets AI agents (Claude Code, Claude Desktop, Codex, and other +[MCP](https://modelcontextprotocol.io) clients) see and control the Switch through +SerialPrograms, while you watch in the app and can take over with the keyboard. + +## Using it + +1. Enable developer mode, open **ML → AI Agent Server**, pick the controller and the + capture card like any program, and press **Start**. +2. Connect the agent to the URL shown in the program (default + `http://127.0.0.1:8765/mcp`, transport "Streamable HTTP") with the header + `Authorization: Bearer `. The program shows a ready-to-paste command, e.g. + for Claude Code: + + ```bash + claude mcp add --transport http switch http://127.0.0.1:8765/mcp --header "Authorization: Bearer " + ``` +3. Watch the agent in the video panel; each action is shown on the overlay and in + the log. **Press any mapped key** to take over: the agent's inputs are refused + (and told why) until you click **Return control to agent**, or after the optional + idle time. **Stop** shuts the server down and releases the controller. + +Only this computer can connect by default. The access token keeps other local +programs and web pages from pressing buttons; see the options for turning it off or +allowing other computers (not recommended). + +## Tools + +The tools, their argument schemas and the instructions sent to agents are defined in +[`AgentTools.json`](AgentTools.json), which is shared with the Python MCP server +(`Source/PythonBindings`, `python -m pokemon_automation.mcp_server`). Each tool lists +the hosts that implement it (`app`, `python`), so both servers expose the same +interface. The input vocabulary (button names, sticks, step fields) is tested against +[`AgentInputTestCases.json`](AgentInputTestCases.json) in both languages. + +App tools: `switch_status`, `screenshot`, `wait_and_observe`, `read_text` (the app's +Tesseract OCR), `press_buttons`, `move_stick`, `run_inputs`, +`cancel_all_commands_blocking`, `get_logs`. + +## Code + +| File | Purpose | +|---|---| +| `AgentServer_HttpServer.*` | Minimal HTTP/1.1 server on Qt's `QTcpServer` (Qt's `QHttpServer` is GPL-only, so it's not used). | +| `AgentServer_McpServer.*` | MCP over Streamable HTTP: JSON-RPC, initialize handshake (protocol 2024-11-05 … 2025-11-25), sessions, token/Origin/Host checks. | +| `AgentServer_ToolDefinitions.*` | Loads `AgentTools.json` (compiled in via `cmake/EmbedTextFile.cmake`), validates tool arguments. | +| `AgentServer_InputSteps.*` | Button/stick/step parsing; the C++ twin of `pokemon_automation/buttons.py`. | +| `ML/Programs/ML_AgentServer.*` | The program: runs input tools on the program thread, screenshots, OCR, user takeover. | + +When changing a tool, edit `AgentTools.json` and both implementations; the Python +tests (`tests/test_shared_interface.py`) check that the Python server matches the file. diff --git a/SerialPrograms/Source/ML/ML_Panels.cpp b/SerialPrograms/Source/ML/ML_Panels.cpp index 41945ac10e..8f8e06d2de 100644 --- a/SerialPrograms/Source/ML/ML_Panels.cpp +++ b/SerialPrograms/Source/ML/ML_Panels.cpp @@ -7,6 +7,7 @@ #include "CommonFramework/StaticGlobals.h" #include "CommonFramework/Panels/PanelTools.h" #include "GameConsole/ConsolePanel.h" +#include "Programs/ML_AgentServer.h" #include "Programs/ML_LabelImages.h" #include "Programs/ML_RunYOLO.h" #include "NintendoSwitch/NintendoSwitch_SingleSwitchProgram.h" @@ -29,6 +30,7 @@ std::vector PanelListFactory::make_panels() const{ ret.emplace_back(GameConsole::make_ConsolePanel()); // ret.emplace_back(make_panel()); ret.emplace_back(NintendoSwitch::make_SingleSwitchProgram()); + ret.emplace_back(NintendoSwitch::make_SingleSwitchProgram()); // ret.emplace_back(make_SingleSwitchProgram()); } diff --git a/SerialPrograms/Source/ML/Programs/ML_AgentServer.cpp b/SerialPrograms/Source/ML/Programs/ML_AgentServer.cpp new file mode 100644 index 0000000000..fbc43b6a81 --- /dev/null +++ b/SerialPrograms/Source/ML/Programs/ML_AgentServer.cpp @@ -0,0 +1,942 @@ +/* ML AI Agent Server Program + * + * From: https://github.com/PokemonAutomation/ + * + */ + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include "Common/Cpp/Exceptions.h" +#include "Common/Cpp/ScopeExit.h" +#include "CommonFramework/Globals.h" +#include "CommonFramework/Language.h" +#include "CommonFramework/ImageTools/ImageBoxes.h" +#include "CommonFramework/ProgramStats/StatsTracking.h" +#include "CommonFramework/VideoPipeline/VideoFeed.h" +#include "CommonFramework/VideoPipeline/VideoOverlay.h" +#include "CommonTools/OCR/OCR_RawTesseractOCR.h" +#include "GameConsole/Framework/ConsoleSystemSession.h" +#include "Integrations/AgentServer/AgentServer_HttpServer.h" +#include "Integrations/AgentServer/AgentServer_InputSteps.h" +#include "Integrations/AgentServer/AgentServer_McpServer.h" +#include "Integrations/AgentServer/AgentServer_ToolDefinitions.h" +#include "NintendoSwitch/Commands/NintendoSwitch_Commands_PushButtons.h" +#include "NintendoSwitch/Controllers/Procon/NintendoSwitch_ProController.h" +#include "ML_AgentServer.h" + +namespace PokemonAutomation{ +namespace ML{ + +using namespace NintendoSwitch; +using AgentServer::InputStep; +using AgentServer::McpContent; +using AgentServer::McpToolResult; +using nlohmann::json; + + + +AgentServer_Descriptor::AgentServer_Descriptor() + : SingleSwitchProgramDescriptor( + "ML:AgentServer", + "ML", "AI Agent Server", + "", + "Let AI agents (Claude, Codex, ...) control the Switch over MCP while you watch " + "and take over with the keyboard at any time.", + ProgramControllerClass::StandardController_NoRestrictions, + FeedbackType::REQUIRED, + AllowCommandsWhenRunning::ENABLE_COMMANDS + ) +{} + +struct AgentServer_Descriptor::Stats : public StatsTracker{ + Stats() + : tool_calls(m_stats["Tool Calls"]) + , inputs(m_stats["Input Calls"]) + , screenshots(m_stats["Screenshots"]) + , takeovers(m_stats["User Takeovers"]) + , errors(m_stats["Errors"]) + { + m_display_order.emplace_back("Tool Calls"); + m_display_order.emplace_back("Input Calls"); + m_display_order.emplace_back("Screenshots"); + m_display_order.emplace_back("User Takeovers", HIDDEN_IF_ZERO); + m_display_order.emplace_back("Errors", HIDDEN_IF_ZERO); + } + std::atomic& tool_calls; + std::atomic& inputs; + std::atomic& screenshots; + std::atomic& takeovers; + std::atomic& errors; +}; +std::unique_ptr AgentServer_Descriptor::make_stats() const{ + return std::unique_ptr(new Stats()); +} + + + +namespace{ + +std::string random_token(){ + std::random_device rd; + std::mt19937_64 rng(((uint64_t)rd() << 32) ^ rd()); + const char* digits = "0123456789abcdef"; + std::string ret; + for (int c = 0; c < 32; c++){ + ret += digits[rng() & 15]; + } + return ret; +} + +int64_t now_ms(){ + return std::chrono::duration_cast( + std::chrono::steady_clock::now().time_since_epoch() + ).count(); +} + +// Encode an image as JPEG and return it base64-encoded, for an MCP image block. +std::string encode_jpeg_base64(const ImageViewRGB32& image, int quality){ + // ImageRGB32 pixels are 0xAARRGGBB, i.e. QImage::Format_RGB32 (alpha ignored). + QImage qimage( + (const uchar*)image.data(), + (int)image.width(), (int)image.height(), + (qsizetype)image.bytes_per_row(), + QImage::Format_RGB32 + ); + QByteArray bytes; + QBuffer buffer(&bytes); + buffer.open(QIODevice::WriteOnly); + qimage.save(&buffer, "JPG", quality); + return bytes.toBase64().toStdString(); +} + +ImageFloatBox to_box(const json& value){ + if (!value.is_array() || value.size() != 4){ + return ImageFloatBox(0, 0, 1, 1); + } + return ImageFloatBox(value[0].get(), value[1].get(), value[2].get(), value[3].get()); +} + +} + + + +// The state of one run of the program: the MCP tool implementations, the queue +// that moves input work onto the program thread, and who is in control. +// +// Threads: +// - Program thread: `run()` executes input jobs on the program's controller context, +// so agent inputs go through the same scheduler as any automation program. +// - HTTP worker threads: `call_tool()`. Observation tools (screenshot, read_text, +// status, logs) run here directly; input tools are queued for the program thread. +// `cancel_all_commands_blocking` runs here directly so it works while an input job is running. +// - UI thread: `on_controller_input()` (keyboard) and the program's buttons. +class AgentSession final : public AgentServer::McpToolHandler, + public GameConsole::ConsoleSystemSession::Listener{ +public: + enum class Control{ AGENT, USER }; + + AgentSession(AgentServerProgram& options, SingleSwitchProgramEnvironment& env, ProControllerContext& context) + : m_options(options) + , m_env(env) + , m_context(context) + , m_stats(env.current_stats()) + , m_settle_ms(options.SETTLE_MS) + , m_max_hold_ms(options.MAX_HOLD_MS) + , m_max_sequence_ms(options.MAX_SEQUENCE_MS) + , m_screenshot_width(options.SCREENSHOT_WIDTH) + , m_jpeg_quality(options.JPEG_QUALITY) + { + m_env.console.system_session().add_listener(*this); + } + ~AgentSession(){ + m_env.console.system_session().remove_listener(*this); + } + + // Program thread: execute queued input jobs until the program is stopped. + // Throws (like any program) when the user presses Stop. + void run(); + + // Stop accepting tool calls and fail the queued ones. Any thread. + void shutdown(); + + // Give control back to the agent (button, or idle timeout). Any thread. + void return_control(const std::string& reason); + + // HTTP worker threads. + virtual McpToolResult call_tool(const std::string& name, const json& arguments) override; + + // UI thread: the user's keyboard input was sent to this console's controller. + virtual void on_controller_input(const ControllerInputState& state) override; + + +private: + struct Job{ + std::function work; + std::promise result; + }; + + McpToolResult call_tool_unchecked(const std::string& name, const json& arguments); + McpToolResult run_on_program_thread(std::function work); + + McpToolResult tool_status(); + McpToolResult tool_cancel_all_commands_blocking(); + McpToolResult tool_get_logs(const json& arguments); + McpToolResult tool_screenshot(const json& box, uint32_t max_width, uint64_t wait_ms); + McpToolResult tool_read_text(const json& arguments); + McpToolResult tool_inputs(std::vector steps, const json& arguments, std::string description); + + // Program thread. + McpToolResult execute_inputs(const std::vector& steps, bool observe, uint64_t settle_ms, const std::string& description); + void issue_step(const InputStep& step); + + // A frame captured at or after `not_before`, cropped to `box`, downscaled to at + // most `max_width`, as [text, image] content. Throws InternalProgramError if the + // video is unavailable. + std::vector capture(const json& box, uint32_t max_width, WallClock not_before); + + // Sleep `ms`, waking early if the session shuts down. Returns false if it did. + bool sleep_unless_stopped(uint64_t ms); + + // Log to the program log, the video overlay and the `get_logs` history. + void add_event(const std::string& text, Color color); + + std::string no_control_message() const{ + return "The user has taken control of the Switch, so your inputs were not sent. " + "Stop sending inputs; call switch_status later to see when control returns to you."; + } + +private: + AgentServerProgram& m_options; + SingleSwitchProgramEnvironment& m_env; + ProControllerContext& m_context; + AgentServer_Descriptor::Stats& m_stats; + + // Option values, read once: options are locked while the program runs. + const uint32_t m_settle_ms; + const uint32_t m_max_hold_ms; + const uint32_t m_max_sequence_ms; + const uint32_t m_screenshot_width; + const uint32_t m_jpeg_quality; + + std::mutex m_lock; + std::condition_variable m_cv; + std::deque> m_jobs; + bool m_stopping = false; + + std::atomic m_control{Control::AGENT}; + // Incremented by cancel_all_commands_blocking and user takeovers; an input job stops issuing + // inputs as soon as it sees a change. + std::atomic m_abort_generation{0}; + std::atomic m_last_keyboard_ms{0}; + std::atomic m_keyboard_neutral{true}; + std::atomic m_stats_dirty{false}; + + std::mutex m_events_lock; + std::deque m_events; +}; + + + +void AgentSession::add_event(const std::string& text, Color color){ + m_env.log("[Agent] " + text, color); + m_env.console.overlay().add_log(text, color); + std::lock_guard lg(m_events_lock); + m_events.emplace_back(current_time_to_str() + " - " + text); + if (m_events.size() > 500){ + m_events.pop_front(); + } +} + + +void AgentSession::run(){ + while (true){ + m_context.throw_if_cancelled(); + + if (m_stats_dirty.exchange(false)){ + m_env.update_stats(); + } + + // Optional: return control after the keyboard has been idle for a while. + uint32_t auto_return = m_options.AUTO_RETURN_SECONDS; + if (auto_return > 0 && + m_control.load() == Control::USER && + m_keyboard_neutral.load() && + now_ms() - m_last_keyboard_ms.load() > (int64_t)auto_return * 1000 + ){ + return_control("no keyboard input for " + std::to_string(auto_return) + " seconds"); + } + + std::shared_ptr job; + { + std::unique_lock lg(m_lock); + m_cv.wait_for(lg, std::chrono::milliseconds(100), [&]{ return !m_jobs.empty() || m_stopping; }); + if (m_stopping){ + return; + } + if (m_jobs.empty()){ + continue; + } + job = std::move(m_jobs.front()); + m_jobs.pop_front(); + } + + try{ + job->result.set_value(job->work()); + }catch (ProgramCancelledException&){ + job->result.set_value(McpToolResult::error("The AI Agent Server was stopped.")); + throw; + }catch (OperationCancelledException&){ + job->result.set_value(McpToolResult::error("The AI Agent Server was stopped.")); + throw; + }catch (Exception& e){ + m_stats.errors++; + job->result.set_value(McpToolResult::error(e.message())); + }catch (std::exception& e){ + m_stats.errors++; + job->result.set_value(McpToolResult::error(e.what())); + } + m_env.update_stats(); + } +} + +void AgentSession::shutdown(){ + std::deque> jobs; + { + std::lock_guard lg(m_lock); + m_stopping = true; + jobs.swap(m_jobs); + } + m_cv.notify_all(); + for (auto& job : jobs){ + job->result.set_value(McpToolResult::error("The AI Agent Server was stopped.")); + } +} + +void AgentSession::return_control(const std::string& reason){ + if (m_control.exchange(Control::AGENT) == Control::USER){ + add_event("Control returned to the agent (" + reason + ").", COLOR_BLUE); + } +} + +void AgentSession::on_controller_input(const ControllerInputState& state){ + m_last_keyboard_ms.store(now_ms()); + bool neutral = state.is_neutral(); + m_keyboard_neutral.store(neutral); + if (neutral){ + return; + } + if (m_control.exchange(Control::USER) == Control::AGENT){ + // The keyboard already overrides the controller's queued commands; this + // stops a running input job from issuing more and refuses new ones. + m_abort_generation++; + m_stats.takeovers++; + m_stats_dirty.store(true); + add_event("You took control. Agent inputs are paused until you click \"Return control to agent\".", COLOR_ORANGE); + } +} + + +McpToolResult AgentSession::run_on_program_thread(std::function work){ + auto job = std::make_shared(); + job->work = std::move(work); + std::future future = job->result.get_future(); + { + std::lock_guard lg(m_lock); + if (m_stopping){ + return McpToolResult::error("The AI Agent Server is stopping."); + } + m_jobs.emplace_back(job); + } + m_cv.notify_all(); + return future.get(); +} + + +McpToolResult AgentSession::call_tool(const std::string& name, const json& arguments){ + m_stats.tool_calls++; + m_stats_dirty.store(true); + try{ + McpToolResult result = call_tool_unchecked(name, arguments); + if (result.is_error){ + m_stats.errors++; + } + return result; + }catch (AgentServer::InputError& e){ + m_stats.errors++; + return McpToolResult::error(e.what()); + }catch (Exception& e){ + m_stats.errors++; + return McpToolResult::error(e.message()); + } +} + +McpToolResult AgentSession::call_tool_unchecked(const std::string& name, const json& arguments){ + { + std::lock_guard lg(m_lock); + if (m_stopping){ + return McpToolResult::error("The AI Agent Server is stopping."); + } + } + + if (name == "switch_status"){ + return tool_status(); + } + if (name == "cancel_all_commands_blocking"){ + return tool_cancel_all_commands_blocking(); + } + if (name == "get_logs"){ + return tool_get_logs(arguments); + } + if (name == "screenshot"){ + return tool_screenshot(arguments.value("box", json()), arguments.value("max_width", m_screenshot_width), 0); + } + if (name == "wait_and_observe"){ + return tool_screenshot(arguments.value("box", json()), m_screenshot_width, arguments["duration_ms"].get()); + } + if (name == "read_text"){ + return tool_read_text(arguments); + } + + // Input tools + if (name == "press_buttons"){ + json step = { + {"buttons", arguments["buttons"]}, + {"hold_ms", arguments["hold_ms"]}, + {"release_ms", arguments["release_ms"]}, + {"repeat", arguments["repeat"]}, + }; + uint64_t repeat = arguments["repeat"].get(); + std::string description = "press " + arguments["buttons"].get() + + (repeat > 1 ? " x" + std::to_string(repeat) : ""); + return tool_inputs({AgentServer::parse_step(step)}, arguments, description); + } + if (name == "move_stick"){ + std::string stick = arguments["stick"].get(); + json step = { + {stick + "_stick", arguments["direction"]}, + {"hold_ms", arguments["duration_ms"]}, + {"release_ms", 0}, + }; + if (arguments.contains("buttons")){ + step["buttons"] = arguments["buttons"]; + } + AgentServer::parse_stick(arguments["direction"]); // clear error for a bad direction + std::string description = stick + " stick " + arguments["direction"].dump() + " for " + + std::to_string(arguments["duration_ms"].get()) + " ms" + + (arguments.contains("buttons") ? " holding " + arguments["buttons"].get() : ""); + return tool_inputs({AgentServer::parse_step(step)}, arguments, description); + } + if (name == "run_inputs"){ + std::vector steps; + for (const json& step : arguments["steps"]){ + steps.emplace_back(AgentServer::parse_step(step)); + } + std::string description = std::to_string(steps.size()) + " step(s)"; + return tool_inputs(std::move(steps), arguments, description); + } + + return McpToolResult::error("Tool " + name + " is not available in SerialPrograms."); +} + + +McpToolResult AgentSession::tool_status(){ + AbstractController& controller = m_env.console.controller(); + VideoSnapshot snapshot = m_env.console.video().snapshot(); + + json status; + status["control"] = m_control.load() == Control::AGENT ? "agent" : "user"; + status["controller"] = { + {"ready", controller.is_ready()}, + {"name", controller.name()}, + }; + if (snapshot){ + status["video"] = { + {"available", true}, + {"resolution", {snapshot.frame->width(), snapshot.frame->height()}}, + }; + }else{ + status["video"] = {{"available", false}}; + } + status["limits"] = { + {"max_hold_ms", m_max_hold_ms}, + {"max_sequence_ms", m_max_sequence_ms}, + }; + status["default_settle_ms"] = m_settle_ms; + status["host"] = "SerialPrograms " + PROGRAM_VERSION; + return McpToolResult::text(status.dump(2)); +} + + +McpToolResult AgentSession::tool_cancel_all_commands_blocking(){ + m_abort_generation++; + bool confirmed = m_env.console.controller().cancel_all_commands_blocking(Milliseconds(1000)); + add_event(std::string("cancel_all_commands_blocking (confirmed: ") + (confirmed ? "yes" : "no") + ")", COLOR_ORANGE); + if (confirmed){ + return McpToolResult::text("All inputs released; the device confirmed the neutral state."); + } + return McpToolResult::text( + "Release requested, but the device did not confirm within 1000 ms. " + "Check switch_status; the controller may be disconnected or the Switch asleep." + ); +} + + +McpToolResult AgentSession::tool_get_logs(const json& arguments){ + size_t count = arguments["count"].get(); + std::string text; + { + std::lock_guard lg(m_events_lock); + size_t start = m_events.size() > count ? m_events.size() - count : 0; + for (size_t c = start; c < m_events.size(); c++){ + text += m_events[c] + "\n"; + } + } + return McpToolResult::text(text.empty() ? "(no events yet)" : text); +} + + +bool AgentSession::sleep_unless_stopped(uint64_t ms){ + std::unique_lock lg(m_lock); + return !m_cv.wait_for(lg, std::chrono::milliseconds(ms), [&]{ return m_stopping; }); +} + + +std::vector AgentSession::capture(const json& box, uint32_t max_width, WallClock not_before){ + VideoSnapshot snapshot = m_env.console.video().snapshot(); + WallClock deadline = current_time() + std::chrono::seconds(1); + while (snapshot && snapshot.timestamp < not_before && current_time() < deadline){ + std::this_thread::sleep_for(std::chrono::milliseconds(15)); + snapshot = m_env.console.video().snapshot(); + } + if (!snapshot){ + throw InternalProgramError( + nullptr, PA_CURRENT_FUNCTION, + "No video. Select the capture card in the AI Agent Server program's video panel." + ); + } + + ImageViewRGB32 view = *snapshot.frame; + if (!box.is_null()){ + view = extract_box_reference(view, to_box(box)); + } + std::string base64; + size_t width = view.width(); + size_t height = view.height(); + if (width > max_width && max_width > 0){ + height = std::max(1, height * max_width / width); + width = max_width; + ImageRGB32 scaled = view.scale_to(width, height); + base64 = encode_jpeg_base64(scaled, (int)m_jpeg_quality); + }else{ + base64 = encode_jpeg_base64(view, (int)m_jpeg_quality); + } + m_stats.screenshots++; + m_stats_dirty.store(true); + + std::string text = "Screenshot " + std::to_string(width) + "x" + std::to_string(height) + + " (full frame " + std::to_string(snapshot.frame->width()) + "x" + std::to_string(snapshot.frame->height()) + ")"; + return { + McpContent::make_text(std::move(text)), + McpContent::make_image(std::move(base64), "image/jpeg"), + }; +} + + +McpToolResult AgentSession::tool_screenshot(const json& box, uint32_t max_width, uint64_t wait_ms){ + if (wait_ms > 0 && !sleep_unless_stopped(wait_ms)){ + return McpToolResult::error("The AI Agent Server was stopped."); + } + McpToolResult result; + result.content = capture(box, max_width, current_time()); + return result; +} + + +McpToolResult AgentSession::tool_read_text(const json& arguments){ + std::string code = arguments["language"].get(); + Language language; + try{ + language = language_code_to_enum(code); + }catch (Exception&){ + return McpToolResult::error("Unknown OCR language \"" + code + "\". Use a Tesseract code such as \"eng\" or \"jpn\"."); + } + if (!OCR::tesseract_language_available(language)){ + return McpToolResult::error( + "OCR data for language \"" + code + "\" is not installed. Ask the user to download the " + "\"Tesseract\" resource in SerialPrograms' Settings (resource downloads)." + ); + } + + std::string mode = arguments["mode"].get(); + OCR::PageSegMode psm = OCR::PageSegMode::SINGLE_BLOCK; + if (mode == "line"){ + psm = OCR::PageSegMode::SINGLE_LINE; + }else if (mode == "word"){ + psm = OCR::PageSegMode::SINGLE_WORD; + }else if (mode == "sparse"){ + psm = (OCR::PageSegMode)11; // Tesseract PSM_SPARSE_TEXT + } + + VideoSnapshot snapshot = m_env.console.video().snapshot(); + if (!snapshot){ + return McpToolResult::error("No video. Select the capture card in the AI Agent Server program's video panel."); + } + ImageViewRGB32 view = *snapshot.frame; + if (arguments.contains("box")){ + view = extract_box_reference(view, to_box(arguments["box"])); + } + std::string text = OCR::tesseract_ocr_read(language, view, psm); + size_t start = text.find_first_not_of(" \t\r\n"); + size_t end = text.find_last_not_of(" \t\r\n"); + text = start == std::string::npos ? "" : text.substr(start, end - start + 1); + return McpToolResult::text(text.empty() ? "(no text found)" : text); +} + + +McpToolResult AgentSession::tool_inputs(std::vector steps, const json& arguments, std::string description){ + if (m_control.load() == Control::USER){ + return McpToolResult::error(no_control_message()); + } + + uint64_t total = 0; + for (const InputStep& step : steps){ + if (!step.wait_only && step.hold_ms > m_max_hold_ms){ + return McpToolResult::error( + "hold_ms " + std::to_string(step.hold_ms) + " exceeds the limit of " + + std::to_string(m_max_hold_ms) + " ms." + ); + } + total += step.duration_ms(); + } + if (total > m_max_sequence_ms){ + return McpToolResult::error( + "The sequence lasts " + std::to_string(total) + " ms, over the limit of " + + std::to_string(m_max_sequence_ms) + " ms. Split it into several calls." + ); + } + + bool observe = arguments["observe"].get(); + uint64_t settle_ms = arguments.contains("settle_ms") ? arguments["settle_ms"].get() : m_settle_ms; + return run_on_program_thread([this, steps = std::move(steps), observe, settle_ms, description, total]{ + McpToolResult result = execute_inputs(steps, observe, settle_ms, description); + if (!result.is_error && !result.content.empty()){ + result.content[0].text = "Done: " + description + " (" + std::to_string(total) + " ms of input)." + + (result.content.size() > 1 ? " " + result.content[0].text : ""); + } + return result; + }); +} + + +void AgentSession::issue_step(const InputStep& step){ + if (step.wait_only){ + if (step.wait_ms > 0){ + pbf_wait(m_context, Milliseconds(step.wait_ms)); + } + return; + } + const bool buttons = step.pressed.has_buttons(); + const bool dpad = step.pressed.has_dpad(); + const bool left = step.left_stick.has_value(); + const bool right = step.right_stick.has_value(); + Milliseconds hold(step.hold_ms); + Milliseconds release(step.release_ms); + Milliseconds cycle(step.hold_ms + step.release_ms); + + for (uint64_t c = 0; c < step.repeat; c++){ + // Single-kind inputs use the dedicated commands, which let the scheduler + // apply per-button cooldowns exactly like `pbf_*()` automation. + if (buttons && !dpad && !left && !right){ + m_context->issue_buttons(&m_context, cycle, hold, release, step.pressed.buttons); + }else if (dpad && !buttons && !left && !right){ + m_context->issue_dpad(&m_context, cycle, hold, release, step.pressed.dpad); + }else if (left && !buttons && !dpad && !right){ + m_context->issue_left_joystick(&m_context, cycle, hold, release, *step.left_stick); + }else if (right && !buttons && !dpad && !left){ + m_context->issue_right_joystick(&m_context, cycle, hold, release, *step.right_stick); + }else{ + m_context->issue_full_controller_state( + &m_context, true, hold, + step.pressed.buttons, step.pressed.dpad, + step.left_stick.value_or(JoystickPosition{}), + step.right_stick.value_or(JoystickPosition{}) + ); + if (step.release_ms > 0){ + pbf_wait(m_context, release); + } + } + } + if (step.wait_ms > 0){ + pbf_wait(m_context, Milliseconds(step.wait_ms)); + } +} + + +McpToolResult AgentSession::execute_inputs( + const std::vector& steps, bool observe, uint64_t settle_ms, const std::string& description +){ + if (m_control.load() == Control::USER){ + return McpToolResult::error(no_control_message()); + } + const uint64_t generation = m_abort_generation.load(); + add_event("Agent: " + description, COLOR_DARKGREEN); + m_stats.inputs++; + m_stats_dirty.store(true); + + WallClock issue_start = current_time(); + uint64_t total_ms = 0; + for (const InputStep& step : steps){ + if (m_abort_generation.load() != generation){ + break; + } + issue_step(step); + total_ms += step.duration_ms(); + } + + // Wait for the inputs to play out. Don't just call wait_for_all_requests(): + // it holds the controller's issue lock for the whole wait, and the user's + // keyboard input needs that lock, so the user couldn't take over (and the UI + // would freeze) until a long agent input finished. Sleep in short slices + // instead, stopping early on a takeover or cancel_all_commands_blocking, and sync with the + // device only at the end, when little or nothing is left to wait for. + WallClock expected_end = issue_start + Milliseconds(total_ms); + while (m_abort_generation.load() == generation && current_time() < expected_end){ + m_context.wait_for(std::min( + Milliseconds(20), + std::chrono::duration_cast(expected_end - current_time()) + Milliseconds(1) + )); + } + if (m_abort_generation.load() == generation){ + m_context.wait_for_all_requests(); + } + if (m_abort_generation.load() != generation){ + return McpToolResult::error( + m_control.load() == Control::USER + ? "Interrupted: " + no_control_message() + : std::string("Interrupted: cancel_all_commands_blocking was called while the inputs were running.") + ); + } + + McpToolResult result = McpToolResult::text(""); + if (observe){ + WallClock after = current_time() + Milliseconds(settle_ms); + m_context.wait_for(Milliseconds(settle_ms)); + std::vector screenshot = capture(json(), m_screenshot_width, after); + result.content[0].text = screenshot[0].text + ", taken " + std::to_string(settle_ms) + " ms after the inputs finished."; + result.content.emplace_back(std::move(screenshot[1])); + } + return result; +} + + + + +AgentServerProgram::~AgentServerProgram(){ + PORT.remove_listener(*this); + REQUIRE_TOKEN.remove_listener(*this); + ACCESS_TOKEN.remove_listener(*this); + ALLOW_NETWORK.remove_listener(*this); + NEW_TOKEN.remove_listener(static_cast(*this)); + RETURN_CONTROL.remove_listener(static_cast(*this)); +} +AgentServerProgram::AgentServerProgram() + : PORT( + "Port:
The local port the MCP server listens on.", + LockMode::LOCK_WHILE_RUNNING, + 8765, 1024, 65535 + ) + , ALLOW_NETWORK( + "Allow connections from other computers:
" + "Off (recommended): only programs on this computer can connect.
" + "On: anything on your network that has the access token can control your Switch.", + LockMode::LOCK_WHILE_RUNNING, + false + ) + , REQUIRE_TOKEN( + "Require access token:
" + "Agents must send the token below. Keeps other programs and web pages from " + "controlling your Switch.", + LockMode::LOCK_WHILE_RUNNING, + true + ) + , ACCESS_TOKEN( + false, + "Access token:
Generated automatically if empty.", + LockMode::LOCK_WHILE_RUNNING, + "", + "Generated when the server starts" + ) + , NEW_TOKEN("", "Generate new access token") + , CONNECTION_INFO("") + , SETTLE_MS( + "Default settle time (ms):
" + "After an agent's inputs finish, wait this long before taking the screenshot " + "that shows their effect. Agents can override it per call.", + LockMode::LOCK_WHILE_RUNNING, + 500, 0, 10000 + ) + , MAX_HOLD_MS( + "Max hold per input (ms):
Longest time an agent may hold an input in one step.", + LockMode::LOCK_WHILE_RUNNING, + 10000, 1, 600000 + ) + , MAX_SEQUENCE_MS( + "Max input per call (ms):
Longest total input an agent may send in one call.", + LockMode::LOCK_WHILE_RUNNING, + 60000, 1, 3600000 + ) + , SCREENSHOT_WIDTH( + "Screenshot width:
Screenshots sent to agents are downscaled to this width.", + LockMode::LOCK_WHILE_RUNNING, + 1280, 64, 3840 + ) + , JPEG_QUALITY( + "Screenshot JPEG quality:", + LockMode::LOCK_WHILE_RUNNING, + 75, 10, 100 + ) + , AUTO_RETURN_SECONDS( + "Auto-return control (seconds):
" + "When you take over with the keyboard, the agent's inputs are paused. Give control " + "back automatically after this many seconds without keyboard input. 0 = only with " + "the button below.", + LockMode::UNLOCK_WHILE_RUNNING, + 0, 0, 3600 + ) + , RETURN_CONTROL( + "Agent control:
After you take over with the keyboard, click to let the agent continue.", + "Return control to agent" + ) +{ + PA_ADD_OPTION(PORT); + PA_ADD_OPTION(ALLOW_NETWORK); + PA_ADD_OPTION(REQUIRE_TOKEN); + PA_ADD_OPTION(ACCESS_TOKEN); + PA_ADD_OPTION(NEW_TOKEN); + PA_ADD_OPTION(CONNECTION_INFO); + PA_ADD_OPTION(SETTLE_MS); + PA_ADD_OPTION(MAX_HOLD_MS); + PA_ADD_OPTION(MAX_SEQUENCE_MS); + PA_ADD_OPTION(SCREENSHOT_WIDTH); + PA_ADD_OPTION(JPEG_QUALITY); + PA_ADD_OPTION(AUTO_RETURN_SECONDS); + PA_ADD_OPTION(RETURN_CONTROL); + + update_connection_info(); + PORT.add_listener(*this); + REQUIRE_TOKEN.add_listener(*this); + ACCESS_TOKEN.add_listener(*this); + ALLOW_NETWORK.add_listener(*this); + NEW_TOKEN.add_listener(static_cast(*this)); + RETURN_CONTROL.add_listener(static_cast(*this)); +} + + +void AgentServerProgram::update_connection_info(){ + std::string url = "http://127.0.0.1:" + std::to_string((uint16_t)PORT) + "/mcp"; + std::string token = ACCESS_TOKEN; + bool require_token = REQUIRE_TOKEN; + + std::string text = "Connect an agent (while this program is running):
"; + text += "MCP server URL: " + url + " (Streamable HTTP)
"; + if (require_token){ + text += "Header: Authorization: Bearer " + (token.empty() ? std::string("<token>") : token) + "
"; + } + text += "Claude Code: claude mcp add --transport http switch " + url; + if (require_token){ + text += " --header \"Authorization: Bearer " + (token.empty() ? std::string("<token>") : token) + "\""; + } + text += ""; + if (require_token && token.empty()){ + text += "
(The token is generated when you press Start.)"; + } + CONNECTION_INFO.set_text(std::move(text)); +} + +void AgentServerProgram::on_config_value_changed(void* object){ + update_connection_info(); +} + +void AgentServerProgram::on_press(ButtonCell& button){ + if (&button == &RETURN_CONTROL){ + std::lock_guard lg(m_session_lock); + if (m_session != nullptr){ + m_session->return_control("you clicked the button"); + } + return; + } + if (&button == &NEW_TOKEN){ + std::lock_guard lg(m_session_lock); + if (m_session == nullptr){ // can't change it while agents are connected + ACCESS_TOKEN.set(random_token()); + } + return; + } +} + + +void AgentServerProgram::program(SingleSwitchProgramEnvironment& env, ProControllerContext& context){ + if (REQUIRE_TOKEN && ((std::string)ACCESS_TOKEN).empty()){ + ACCESS_TOKEN.set(random_token()); + } + + static const AgentServer::AgentToolDefinitions definitions( + AgentServer::agent_tools_json_text(), "app" + ); + + AgentServer::McpServerConfig config; + config.server_version = PROGRAM_VERSION; + config.access_token = REQUIRE_TOKEN ? (std::string)ACCESS_TOKEN : ""; + config.localhost_only = !ALLOW_NETWORK; + + AgentSession session(*this, env, context); + AgentServer::McpServer mcp(env.logger(), definitions, session, config); + AgentServer::HttpServer http(env.logger(), [&mcp](const AgentServer::HttpRequest& request){ + return mcp.handle(request); + }); + + { + std::lock_guard lg(m_session_lock); + m_session = &session; + } + NEW_TOKEN.set_enabled(false); // agents are using the current token + // On any exit (Stop, error): refuse new tool calls, fail queued ones, then stop + // the server, which waits for in-flight requests to return. + // This runs during stack unwinding, so it must not throw. + ScopeExit cleanup([&]{ + session.shutdown(); + http.stop(); + { + std::lock_guard lg(m_session_lock); + m_session = nullptr; + } + NEW_TOKEN.set_enabled(true); + try{ + env.console.controller().cancel_all_commands_blocking(Milliseconds(500)); + }catch (...){} + }); + + std::string error = http.start(PORT, ALLOW_NETWORK); + if (!error.empty()){ + throw UserSetupError(env.logger(), "Unable to start the AI agent server: " + error); + } + env.console.overlay().add_log( + "AI agent server on port " + std::to_string(http.port()) + ". Waiting for an agent...", + COLOR_WHITE + ); + + session.run(); +} + + + + +} +} diff --git a/SerialPrograms/Source/ML/Programs/ML_AgentServer.h b/SerialPrograms/Source/ML/Programs/ML_AgentServer.h new file mode 100644 index 0000000000..9e9d1d4e92 --- /dev/null +++ b/SerialPrograms/Source/ML/Programs/ML_AgentServer.h @@ -0,0 +1,96 @@ +/* ML AI Agent Server Program + * + * From: https://github.com/PokemonAutomation/ + * + * Lets AI agents control the Switch through SerialPrograms. While running, this + * program hosts an MCP (Model Context Protocol) server on a local port. Any MCP + * client (Claude Code, Claude Desktop, Codex, ...) can connect to it and press + * buttons, move sticks, take screenshots and read on-screen text, using this + * console's controller and video. The tools are defined in the shared + * Integrations/AgentServer/AgentTools.json, which the Python MCP server + * (pokemon_automation.mcp_server) serves too. + * + * The user watches the agent in this program's video panel and can take over at any + * time with keyboard control: the agent's inputs are refused until the user clicks + * "Return control to agent" (or, optionally, after some idle seconds). + * Stopping the program stops the server and releases the controller. + */ + +#ifndef PokemonAutomation_ML_AgentServer_H +#define PokemonAutomation_ML_AgentServer_H + +#include +#include +#include "Common/Cpp/Options/BooleanCheckBoxOption.h" +#include "Common/Cpp/Options/ButtonOption.h" +#include "Common/Cpp/Options/SimpleIntegerOption.h" +#include "Common/Cpp/Options/StaticTextOption.h" +#include "Common/Cpp/Options/StringOption.h" +#include "NintendoSwitch/NintendoSwitch_SingleSwitchProgram.h" + +namespace PokemonAutomation{ +namespace ML{ + + +class AgentServer_Descriptor : public NintendoSwitch::SingleSwitchProgramDescriptor{ +public: + AgentServer_Descriptor(); + + struct Stats; + virtual std::unique_ptr make_stats() const override; +}; + + +class AgentSession; + +class AgentServerProgram : public NintendoSwitch::SingleSwitchProgramInstance, + public ButtonListener, + public ConfigOption::Listener{ +public: + using Descriptor = AgentServer_Descriptor; + + ~AgentServerProgram(); + AgentServerProgram(); + + virtual void program( + NintendoSwitch::SingleSwitchProgramEnvironment& env, + NintendoSwitch::ProControllerContext& context + ) override; + + // "Return control to agent" and "New access token" buttons. + virtual void on_press(ButtonCell& button) override; + // Keeps the connection instructions in sync with the port and token options. + virtual void on_config_value_changed(void* object) override; + +private: + // Refresh CONNECTION_INFO from the current options. + void update_connection_info(); + +public: + SimpleIntegerOption PORT; + BooleanCheckBoxOption ALLOW_NETWORK; + BooleanCheckBoxOption REQUIRE_TOKEN; + StringOption ACCESS_TOKEN; + ButtonOption NEW_TOKEN; + StaticTextOption CONNECTION_INFO; + + SimpleIntegerOption SETTLE_MS; + SimpleIntegerOption MAX_HOLD_MS; + SimpleIntegerOption MAX_SEQUENCE_MS; + SimpleIntegerOption SCREENSHOT_WIDTH; + SimpleIntegerOption JPEG_QUALITY; + SimpleIntegerOption AUTO_RETURN_SECONDS; + ButtonOption RETURN_CONTROL; + +private: + // The running session, if any. Buttons are pressed on the UI thread. + std::mutex m_session_lock; + AgentSession* m_session = nullptr; +}; + + + + +} +} +#endif diff --git a/SerialPrograms/Source/PythonBindings/README.md b/SerialPrograms/Source/PythonBindings/README.md index 1f926fc954..02b77fdc52 100644 --- a/SerialPrograms/Source/PythonBindings/README.md +++ b/SerialPrograms/Source/PythonBindings/README.md @@ -111,7 +111,11 @@ Conventions: ## MCP server -The tools are defined in the shared `Source/Integrations/AgentServer/AgentTools.json`. +The same MCP interface is also served by the SerialPrograms app itself (**ML → AI +Agent Server**, see `Source/Integrations/AgentServer/README.md`), so an agent can +drive the Switch through the app while you watch and take over with the keyboard. +Both servers load their tools from the shared +`Source/Integrations/AgentServer/AgentTools.json`. ```bash python -m pokemon_automation.mcp_server --serial /dev/cu.usbserial-0001 --video MiraBox diff --git a/SerialPrograms/cmake/SourceFiles.cmake b/SerialPrograms/cmake/SourceFiles.cmake index d283afc08c..ef97ad6a03 100644 --- a/SerialPrograms/cmake/SourceFiles.cmake +++ b/SerialPrograms/cmake/SourceFiles.cmake @@ -1193,6 +1193,8 @@ file(GLOB LIBRARY_SOURCES Source/ML/Models/ML_OrtEnv.h Source/ML/Models/ML_YOLOv5Model.cpp Source/ML/Models/ML_YOLOv5Model.h + Source/ML/Programs/ML_AgentServer.cpp + Source/ML/Programs/ML_AgentServer.h Source/ML/Programs/ML_LabelImages.cpp Source/ML/Programs/ML_LabelImages.h Source/ML/Programs/ML_LabelImagesOverlayManager.cpp