diff --git a/SerialPrograms/CMakeLists.txt b/SerialPrograms/CMakeLists.txt index d456953b4b..d04fed9ba5 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() @@ -108,6 +115,7 @@ if(WIN32 AND QT_DEPLOY_FILES) Qml Quick QuickWidgets + Network ) else() # Find all subdirectories in the Qt base directory @@ -120,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}*") @@ -157,6 +166,7 @@ else() Qml Quick QuickWidgets + Network ) endif() @@ -265,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 @@ -273,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) @@ -846,3 +867,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/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/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){ 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/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/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/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/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/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/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/Integrations/PybindSwitchController.cpp b/SerialPrograms/Source/Integrations/PybindSwitchController.cpp index 77aae9d4a1..e0b034b36b 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(){ @@ -140,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 decbe6a252..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 @@ -32,6 +38,10 @@ 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. In particular +// `cancel_all_commands_blocking()` may be called while another thread is blocked in +// `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; @@ -58,10 +68,29 @@ 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(); + // 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. 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/.gitignore b/SerialPrograms/Source/PythonBindings/.gitignore new file mode 100644 index 0000000000..a6ed7df9c4 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/.gitignore @@ -0,0 +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 new file mode 100644 index 0000000000..c3f4657167 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/PythonBindings.cmake @@ -0,0 +1,67 @@ +# 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) +# 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}/ + 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/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/README.md b/SerialPrograms/Source/PythonBindings/README.md new file mode 100644 index 0000000000..02b77fdc52 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/README.md @@ -0,0 +1,176 @@ +# 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 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 +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/pokemon_automation/__init__.py b/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py new file mode 100644 index 0000000000..9e9712cd28 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/__init__.py @@ -0,0 +1,42 @@ +"""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`. +- 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 + + 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", + "VideoSource", + "encode_image", + "list_serial_ports", + "list_video_devices", + "ocr_image", + "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/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/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/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/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/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 new file mode 100644 index 0000000000..c811d2951b --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pokemon_automation/fake.py @@ -0,0 +1,135 @@ +"""Fake controller and video backends for testing without hardware. + +`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: + """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"] + + +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/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/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/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 new file mode 100644 index 0000000000..8f3b8c38f1 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/pyproject.toml @@ -0,0 +1,38 @@ +[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 .`. 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." +readme = "README.md" +requires-python = ">=3.10" +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] +mcp = ["mcp>=1.2"] +ocr = ["pytesseract>=0.3.10"] # also needs the tesseract program installed +serial = ["pyserial>=3.5"] +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"] + +[tool.setuptools.package-data] +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_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)) diff --git a/SerialPrograms/Source/PythonBindings/tests/test_controller.py b/SerialPrograms/Source/PythonBindings/tests/test_controller.py new file mode 100644 index 0000000000..fb76f0a32f --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/tests/test_controller.py @@ -0,0 +1,92 @@ +import pytest + +from pokemon_automation import Console, InputStep, SwitchController, VideoSource +from pokemon_automation import buttons as btn +from pokemon_automation.fake import FakeController, FakeVideoCapture + + +@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_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() + 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 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/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 new file mode 100644 index 0000000000..3d03b991b2 --- /dev/null +++ b/SerialPrograms/Source/PythonBindings/tests/test_shared_interface.py @@ -0,0 +1,94 @@ +"""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 +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"] + + +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() + + +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"] 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 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. # 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 f5e1ea9435..ef97ad6a03 100644 --- a/SerialPrograms/cmake/SourceFiles.cmake +++ b/SerialPrograms/cmake/SourceFiles.cmake @@ -904,6 +904,14 @@ 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/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 Source/Integrations/DiscordIntegrationSettings.h Source/Integrations/DiscordIntegrationTable.cpp @@ -1185,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