diff --git a/Cargo.lock b/Cargo.lock index 09447d2979..d80a26cd60 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -204,6 +204,7 @@ dependencies = [ "insta", "rayon", "rstest 0.24.0", + "schemars", "serde", "serde_derive", "serde_json", @@ -250,7 +251,7 @@ dependencies = [ "regex", "rustc-hash", "shlex", - "syn", + "syn 2.0.104", ] [[package]] @@ -339,7 +340,7 @@ dependencies = [ "proc-macro2", "quote", "rustversion", - "syn", + "syn 2.0.104", ] [[package]] @@ -487,7 +488,7 @@ dependencies = [ "heck", "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -693,7 +694,7 @@ dependencies = [ "proc-macro2", "quote", "strsim", - "syn", + "syn 2.0.104", ] [[package]] @@ -704,7 +705,7 @@ checksum = "fc34b93ccb385b40dc71c6fceac4b2ad23662c7eeb248cf10d529b7e055b6ead" dependencies = [ "darling_core", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -757,7 +758,7 @@ checksum = "6edb4b64a43d977b8e99788fe3a04d483834fba1215a7e02caa415b626497f7f" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -817,7 +818,7 @@ checksum = "97369cbbc041bc366949bc74d34658d6cda5621039731c6310521892a3a20ae0" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -869,6 +870,12 @@ dependencies = [ "zstd", ] +[[package]] +name = "dyn-clone" +version = "1.0.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d0881ea181b1df73ff77ffaaf9c7544ecc11e82fba9b5f27b262a3c73a332555" + [[package]] name = "either" version = "1.15.0" @@ -1068,7 +1075,7 @@ checksum = "162ee34ebcb7c64a8abebc059ce0fee27c2262618d7b60ed8faf72fef13c3650" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -1703,7 +1710,7 @@ checksum = "ed3955f1a9c7c0c15e092f9c887db08b1fc683305fdf6eb6684f22555355e202" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -1734,7 +1741,7 @@ dependencies = [ "proc-macro-crate", "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -2027,7 +2034,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "061c1221631e079b26479d25bbf2275bfe5917ae8419cd7e34f13bfc2aa7539a" dependencies = [ "proc-macro2", - "syn", + "syn 2.0.104", ] [[package]] @@ -2078,7 +2085,7 @@ dependencies = [ "itertools 0.14.0", "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -2182,6 +2189,26 @@ dependencies = [ "thiserror 2.0.18", ] +[[package]] +name = "ref-cast" +version = "1.0.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7e440fb4e4b4147295338efb76001ab9e4efc0e5839df2c47fc5ac2381d365c3" +dependencies = [ + "ref-cast-impl", +] + +[[package]] +name = "ref-cast-impl" +version = "1.0.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92ecd8964f8453721699a1ed72037b0db49ce2f5a5138486ee89bed6f67cdf3a" +dependencies = [ + "proc-macro2", + "quote", + "syn 3.0.6", +] + [[package]] name = "regex" version = "1.11.1" @@ -2290,7 +2317,7 @@ dependencies = [ "regex", "relative-path", "rustc_version", - "syn", + "syn 2.0.104", "unicode-ident", ] @@ -2308,7 +2335,7 @@ dependencies = [ "regex", "relative-path", "rustc_version", - "syn", + "syn 2.0.104", "unicode-ident", ] @@ -2439,6 +2466,31 @@ dependencies = [ "sdd", ] +[[package]] +name = "schemars" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "687274d293b6cdc6e73e0fee520bf2049650090d7164f87672d212a3c530cf4a" +dependencies = [ + "dyn-clone", + "ref-cast", + "schemars_derive", + "serde", + "serde_json", +] + +[[package]] +name = "schemars_derive" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d98c67716b46af2f0b8cf752abc930f6f9aecfbf671ecfb531db8a31dbe4e2ba" +dependencies = [ + "proc-macro2", + "quote", + "serde_derive_internals", + "syn 3.0.6", +] + [[package]] name = "scopeguard" version = "1.2.0" @@ -2468,7 +2520,7 @@ checksum = "1783eabc414609e28a5ba76aee5ddd52199f7107a0b24c2e9746a1ecc34a683d" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -2531,7 +2583,18 @@ checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", +] + +[[package]] +name = "serde_derive_internals" +version = "0.30.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f852137cce035d6a4df67ccce505ff6b3e9fd3a10e3e52b24dc71e650bb1a9bd" +dependencies = [ + "proc-macro2", + "quote", + "syn 3.0.6", ] [[package]] @@ -2580,7 +2643,7 @@ checksum = "5d69265a08751de7844521fd15003ae0a888e035773ba05695c5c759a6f89eef" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -2641,7 +2704,7 @@ checksum = "0eb01866308440fc64d6c44d9e86c5cc17adfe33c4d6eed55da9145044d0ffc1" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -2716,6 +2779,17 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "syn" +version = "3.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8593e8e72159ed2257d083c7a454a85cbf854f37a0966d8d483aff8c8a3ebcee" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + [[package]] name = "synstructure" version = "0.13.2" @@ -2724,7 +2798,7 @@ checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -2776,7 +2850,7 @@ checksum = "4fee6c4efc90059e10f81e6d42c60a18f76588c3d74cb83a0b242a2b6c7504c1" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -2787,7 +2861,7 @@ checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -2887,7 +2961,7 @@ checksum = "7490cfa5ec963746568740651ac6781f701c9c5ea257c58e057f3ba8cf69e8da" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -3151,7 +3225,7 @@ dependencies = [ "log", "proc-macro2", "quote", - "syn", + "syn 2.0.104", "wasm-bindgen-shared", ] @@ -3173,7 +3247,7 @@ checksum = "8ae87ea40c9f689fc23f209965b6fb8a99ad69aeeb0231408be24920604395de" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", "wasm-bindgen-backend", "wasm-bindgen-shared", ] @@ -3551,7 +3625,7 @@ checksum = "b659052874eb698efe5b9e8cf382204678a0086ebf46982b79d6ca3182927e5d" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", "synstructure", ] @@ -3582,7 +3656,7 @@ checksum = "125139de3f6b9d625c39e2efdd73d41bdac468ccd556556440e322be0e1bbd91" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -3593,7 +3667,7 @@ checksum = "9ecf5b4cc5364572d7f4c329661bcc82724222973f2cab6f050a4e5c22f75181" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] @@ -3613,7 +3687,7 @@ checksum = "d71e5d6e06ab090c67b5e44993ec16b72dcbaabc526db883a360057678b48502" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", "synstructure", ] @@ -3647,7 +3721,7 @@ checksum = "eadce39539ca5cb3985590102671f2567e659fca9666581ad3411d59207951f3" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.104", ] [[package]] diff --git a/api-docs/cppdocs/Doxyfile-Docset b/api-docs/cppdocs/Doxyfile-Docset index eb3774bd62..f93deb4603 100644 --- a/api-docs/cppdocs/Doxyfile-Docset +++ b/api-docs/cppdocs/Doxyfile-Docset @@ -913,7 +913,10 @@ EXCLUDE_SYMBOLS = BinaryNinja::LowLevelILInstructionAccessor \ BinaryNinja::MediumLevelILInstructionAccessor* \ BinaryNinja::HighLevelILInstructionAccessor \ BinaryNinja::HighLevelILInstructionAccessor* \ - + BinaryNinja::detail \ + BinaryNinja::detail::* \ + BinaryNinja::MCP::detail \ + BinaryNinja::MCP::detail::* #EXCLUDE_SYMBOLS = QProgressIndicator #Broke Exhale diff --git a/api-docs/cppdocs/Doxyfile-HTML b/api-docs/cppdocs/Doxyfile-HTML index 22250a5a5a..386ab29ead 100644 --- a/api-docs/cppdocs/Doxyfile-HTML +++ b/api-docs/cppdocs/Doxyfile-HTML @@ -912,7 +912,11 @@ EXCLUDE_SYMBOLS = BinaryNinja::LowLevelILInstructionAccessor \ BinaryNinja::MediumLevelILInstructionAccessor \ BinaryNinja::MediumLevelILInstructionAccessor* \ BinaryNinja::HighLevelILInstructionAccessor \ - BinaryNinja::HighLevelILInstructionAccessor* + BinaryNinja::HighLevelILInstructionAccessor* \ + BinaryNinja::detail \ + BinaryNinja::detail::* \ + BinaryNinja::MCP::detail \ + BinaryNinja::MCP::detail::* #EXCLUDE_SYMBOLS = QProgressIndicator #Broke Exhale diff --git a/api-docs/source/conf.py b/api-docs/source/conf.py index c7a4d866b1..192d61b93a 100644 --- a/api-docs/source/conf.py +++ b/api-docs/source/conf.py @@ -172,6 +172,11 @@ def get_docstring_summary(name, ref): directive_start = summary.rfind(':py:', 0, truncate_at) if directive_start != -1: truncate_at = directive_start + # Check for an unclosed ``literal``. Its backticks come in pairs, so the single-backtick check misses it. + if summary[:truncate_at].count('``') % 2 != 0: + last_literal = summary.rfind('``', 0, truncate_at) + if last_literal > 0: + truncate_at = summary.rfind(' ', 0, last_literal) # Check for unclosed backticks if summary[:truncate_at].count('`') % 2 != 0: last_backtick = summary.rfind('`', 0, truncate_at) @@ -344,21 +349,26 @@ def generaterst(): module_contents.write('\n') # Generate individual sections with proper headers - for (classname, classref) in members: - # Only include classes that actually belong to this module - if inspect.getmodule(classref).__name__ == module.__name__: - _, directive, needs_members = get_autodoc_info(classname, classref) - module_contents.write(f'''{classname} -{"-" * len(classname)} + own_members = [(name, ref) for name, ref in members if inspect.getmodule(ref).__name__ == module.__name__] + # autosectionlabel lowercases section titles, so title a function whose name differs from another member's only + # by case with parentheses, such as tool() beside Tool. + lowercase_names = [name.lower() for name, _ in own_members] + for (classname, classref) in own_members: + _, directive, needs_members = get_autodoc_info(classname, classref) + title = classname + if inspect.isfunction(classref) and lowercase_names.count(classname.lower()) > 1: + title = f"{classname}()" + module_contents.write(f'''{title} +{"-" * len(title)} .. {directive}:: {module.__name__}.{classname} ''') - if needs_members: - module_contents.write(''' :members: + if needs_members: + module_contents.write(''' :members: :undoc-members: :show-inheritance: ''') - module_contents.write('\n') + module_contents.write('\n') module_contents.write(stats) new_module_contents = module_contents.getvalue() diff --git a/binaryninjaapi.h b/binaryninjaapi.h index eae5334dae..961818c713 100644 --- a/binaryninjaapi.h +++ b/binaryninjaapi.h @@ -2122,7 +2122,7 @@ namespace BinaryNinja { size_t result; if (!BNCoreEnumFromString(name.c_str(), value.c_str(), &result)) return std::nullopt; - return result; + return static_cast(result); } diff --git a/binaryninjacore.h b/binaryninjacore.h index c53f23db44..18641d8b49 100644 --- a/binaryninjacore.h +++ b/binaryninjacore.h @@ -37,7 +37,7 @@ // Current ABI version for linking to the core. This is incremented any time // there are changes to the API that affect linking, including new functions, // new types, or modifications to existing functions or types. -#define BN_CURRENT_CORE_ABI_VERSION 188 +#define BN_CURRENT_CORE_ABI_VERSION 189 // Minimum ABI version that is supported for loading of plugins. Plugins that // are linked to an ABI version less than this will not be able to load and @@ -347,6 +347,9 @@ extern "C" typedef struct BNUndoAction BNUndoAction; typedef struct BNUndoEntry BNUndoEntry; typedef struct BNDemangler BNDemangler; + typedef struct BNMcpTool BNMcpTool; + typedef struct BNMcpToolCall BNMcpToolCall; + typedef struct BNMcpToolResult BNMcpToolResult; typedef struct BNFirmwareNinja BNFirmwareNinja; typedef struct BNFirmwareNinjaReferenceNode BNFirmwareNinjaReferenceNode; typedef struct BNFirmwareNinjaRelationship BNFirmwareNinjaRelationship; @@ -4063,6 +4066,46 @@ extern "C" void (*freeResult)(void* ctxt, BNDemanglerResult* result); } BNDemanglerCallbacks; + BN_ENUM(uint8_t, BNMcpToolScope) + { + GlobalScope, + BinaryViewScope + }; + + BN_OPTIONS(uint32_t, BNMcpToolAnnotation) + { + ReadOnlyHint = 1, + DestructiveHint = 2, + IdempotentHint = 4, + OpenWorldHint = 8 + }; + + typedef struct BNMcpToolDefinition + { + const char* name; + const char* title; + const char* description; + const char* inputSchema; + const char* outputSchema; + BNMcpToolScope scope; + uint32_t annotations; + } BNMcpToolDefinition; + + typedef struct BNMcpToolCallbacks + { + void* context; + void (*invoke)(void* ctxt, BNMcpToolCall* call, const char* arguments, BNMcpToolResult* result); + void (*freeObject)(void* ctxt); + } BNMcpToolCallbacks; + + typedef struct BNMcpToolCallCallbacks + { + void* context; + BNBinaryView* (*getBinaryView)(void* ctxt); + bool (*isCancelled)(void* ctxt); + void (*reportProgress)(void* ctxt, double progress, double total, const char* message); + } BNMcpToolCallCallbacks; + BN_ENUM(uint8_t, BNScopeType) { OneLineScopeType, @@ -8866,6 +8909,50 @@ extern "C" BINARYNINJACOREAPI bool BNPromoteDemangler(const BNDemangler* demangler); BINARYNINJACOREAPI bool BNIsDemanglerMangledName(const BNDemangler* demangler, const char* name); + // MCP tools + BINARYNINJACOREAPI BNMcpTool* BNRegisterMcpTool( + const BNMcpToolDefinition* definition, const BNMcpToolCallbacks* callbacks); + BINARYNINJACOREAPI BNMcpTool* BNNewMcpToolReference(BNMcpTool* tool); + BINARYNINJACOREAPI void BNFreeMcpTool(BNMcpTool* tool); + BINARYNINJACOREAPI BNMcpTool** BNGetMcpToolList(size_t* count); + BINARYNINJACOREAPI void BNFreeMcpToolList(BNMcpTool** tools, size_t count); + BINARYNINJACOREAPI BNMcpTool* BNGetMcpToolByName(const char* name); + BINARYNINJACOREAPI char* BNGetMcpToolName(BNMcpTool* tool); + BINARYNINJACOREAPI char* BNGetMcpToolTitle(BNMcpTool* tool); + BINARYNINJACOREAPI char* BNGetMcpToolDescription(BNMcpTool* tool); + BINARYNINJACOREAPI char* BNGetMcpToolInputSchema(BNMcpTool* tool); + BINARYNINJACOREAPI char* BNGetMcpToolOutputSchema(BNMcpTool* tool); + BINARYNINJACOREAPI BNMcpToolScope BNGetMcpToolScope(BNMcpTool* tool); + BINARYNINJACOREAPI uint32_t BNGetMcpToolAnnotations(BNMcpTool* tool); + + BINARYNINJACOREAPI BNMcpToolCall* BNNewMcpToolCallReference(BNMcpToolCall* call); + BINARYNINJACOREAPI void BNFreeMcpToolCall(BNMcpToolCall* call); + BINARYNINJACOREAPI BNBinaryView* BNGetMcpToolCallBinaryView(BNMcpToolCall* call); + BINARYNINJACOREAPI bool BNIsMcpToolCallCancelled(BNMcpToolCall* call); + BINARYNINJACOREAPI void BNReportMcpToolCallProgress( + BNMcpToolCall* call, double progress, double total, const char* message); + BINARYNINJACOREAPI bool BNParseMcpToolCallAddress( + BNMcpToolCall* call, const char* value, uint64_t* result, uint64_t here, char** errorMessage); + BINARYNINJACOREAPI bool BNParseMcpToolCallInteger( + BNMcpToolCall* call, const char* value, uint64_t* result, uint64_t here, char** errorMessage); + + BINARYNINJACOREAPI BNMcpToolResult* BNNewMcpToolResultReference(BNMcpToolResult* result); + BINARYNINJACOREAPI void BNFreeMcpToolResult(BNMcpToolResult* result); + BINARYNINJACOREAPI void BNAddMcpToolResultText(BNMcpToolResult* result, const char* text); + BINARYNINJACOREAPI bool BNSetMcpToolResultStructuredContent(BNMcpToolResult* result, const char* json); + BINARYNINJACOREAPI void BNSetMcpToolResultError( + BNMcpToolResult* result, const char* code, const char* message, const char* detailsJson); + BINARYNINJACOREAPI void BNAddMcpToolResultWarning(BNMcpToolResult* result, const char* code, const char* message); + + // MCP servers + BINARYNINJACOREAPI BNMcpTool* BNCreateMcpTool( + const BNMcpToolDefinition* definition, const BNMcpToolCallbacks* callbacks); + BINARYNINJACOREAPI BNMcpToolCall* BNCreateMcpToolCall(const BNMcpToolCallCallbacks* callbacks); + BINARYNINJACOREAPI BNMcpToolResult* BNCreateMcpToolResult(void); + BINARYNINJACOREAPI void BNInvokeMcpTool( + BNMcpTool* tool, BNMcpToolCall* call, const char* arguments, BNMcpToolResult* result); + BINARYNINJACOREAPI char* BNGetMcpToolResultJson(BNMcpToolResult* result); + // Plugin repository APIs BINARYNINJACOREAPI char** BNPluginGetApis(BNPlugin* p, size_t* count); BINARYNINJACOREAPI const char* BNPluginGetAuthor(BNPlugin* p); diff --git a/docs/dev/index.md b/docs/dev/index.md index 94772b2532..c1becc0493 100644 --- a/docs/dev/index.md +++ b/docs/dev/index.md @@ -12,6 +12,7 @@ The Python API is the most common third-party API and is used in many [public pl - [Writing Python Plugins](plugins.md) - [Container Transforms](containertransforms.md) - Creating custom container/archive decoders + - [Writing MCP Tools](mcp-tools.md) - Adding tools to Binary Ninja's MCP server, in Python, C++ or Rust - [Applying Annotations](annotation.md) - [Script Cookbook](cookbook.md) with common examples and concepts explained - [Python API Reference](https://api.binary.ninja/) (available offline via the Help menu) diff --git a/docs/dev/mcp-tools.md b/docs/dev/mcp-tools.md new file mode 100644 index 0000000000..ae75f5201b --- /dev/null +++ b/docs/dev/mcp-tools.md @@ -0,0 +1,191 @@ +# Writing MCP Tools + +Plugins can add tools to Binary Ninja's [MCP server](../guide/mcp.md). Both the GUI's MCP server and the headless `binaryninja_mcp` server offer the tools you register alongside the built-in `bn_*` tools. Tools can be written in Python, C++ or Rust. + +## Concepts + +A tool has a name, a description, a JSON Schema describing its input, and a handler that receives the arguments and returns a result. + +- **Names** must be 1 to 64 letters, digits, `_` or `-`, and unique in the process. Prefix yours with your plugin's name, such as `myplugin_find_crypto`. The `bn_` prefix is used by Binary Ninja's built-in tools. +- **Scope.** A BinaryView-scoped tool, the default, runs against the MCP session's active BinaryView. The server resolves that view before the handler runs and reports `no_active_binary_view` itself when there is none, so the handler can rely on having one. A global tool runs without a view, and can still ask for one if the session has it. +- **Annotations** mark a tool as read-only, destructive, idempotent or open-world. Clients use them to decide how to present a tool and whether to ask the user before running it. Every hint is sent, so a tool without a flag is advertised as one that may modify its environment, is not destructive, is not idempotent, and does not reach outside Binary Ninja. +- **Results** are text, structured content (a JSON object), or both. When a tool returns structured content and no text, clients that ignore structured content see the JSON as text. An error result carries a machine-readable `errorCode` and an `errorMessage` as structured content. A tool with an output schema returns them as text instead, since the error does not match the schema. +- **Return one representation.** Some MCP clients show the model only the structured content, and others only the text. A tool that returns both must put everything the model needs in each, so prefer returning one or the other. +- **Warnings** are advisory notes about a result, such as analysis that has not finished, each with a `code` and a `message`. Binary Ninja adds them to structured content as a `warnings` array and lists them as plain text after any text. The `warnings` property is reserved, so an output schema must not declare it. Registration adds it to every output schema. Facts about the result, such as how many bytes were read, belong in the result itself. +- **Argument checking.** The typed interfaces in each language build the input schema from your declarations and check every argument before the handler runs, except one whose schema you supply yourself. A missing, mistyped or undeclared argument produces an `invalid_params` error without calling the handler. A null argument for an optional parameter is treated as absent. +- **Address expressions.** Parameters declared as addresses accept Binary Ninja expression strings such as `0x401000`, `main` or `.text + 0x10`, and your handler receives the evaluated address. Parameters declared as integer expressions accept either a JSON integer or an expression string. See [Address Expressions](../guide/mcp.md#address-expressions). +- **Threading.** Handlers run on the MCP server's thread, not the UI's main thread. Use `execute_on_main_thread_and_wait` (Python), `ExecuteOnMainThreadAndWait` (C++) or `main_thread::execute_on_main_thread_and_wait` (Rust) for anything that touches the UI. +- **Cancellation and progress.** Binary Ninja's MCP servers don't yet pass cancellation requests or progress between clients and tools. Until they do, a call never reports that it has been cancelled while the handler runs, and progress a tool reports is discarded. A long-running tool should still check for cancellation and report its progress, so that both work once the servers support them. Once the handler returns, the call is detached. It then has no binary view and reports itself cancelled, so a background thread that keeps the call stops with it. + +Clients cache the tool list, so register tools when your plugin loads. A client that connected before your plugin registered its tools sees them after it reconnects. + +## Python + +The `binaryninja.mcp` module's `tool` decorator builds the tool from the function. The first parameter receives the `ToolCall`. Every other parameter becomes a tool parameter, and needs a type annotation and a `:param name:` line in the docstring. The docstring's leading paragraph becomes the tool's description. + +```python +from typing import List, Literal, Optional +from binaryninja import mcp + +@mcp.tool(read_only=True) +def myplugin_find_strings( + call: mcp.ToolCall, + min_length: int = 4, + section: Optional[str] = None, +) -> dict: + """Find strings in the active binary view. + + :param min_length: Minimum string length in bytes. + :param section: Only search this section. + """ + bv = call.binary_view + strings = [s for s in bv.strings if s.length >= min_length] + if section is not None: + target = bv.get_section_by_name(section) + if target is None: + raise mcp.ToolError("section_not_found", f"No section named {section}") + strings = [s for s in strings if target.start <= s.start < target.end] + return {"strings": [{"address": hex(s.start), "value": s.value} for s in strings]} +``` + +The decorator supports these annotations: + +| Annotation | Schema | The function receives | +| --- | --- | --- | +| `str`, `int`, `float`, `bool` | `string`, `integer`, `number`, `boolean` | the value | +| `Annotated[int, mcp.Minimum(0), mcp.Maximum(100)]` | `integer` with those bounds | the value, after rejecting values outside the bounds | +| `Annotated[int, mcp.ClampTo(0, 1000)]` | `integer` with those bounds | the value, moved to the nearest bound when outside them | +| `Annotated[str, mcp.NonEmpty()]` | `string` with `minLength` 1 | the value, after rejecting an empty string | +| `mcp.Address` | `string` | the evaluated address, as an `int` | +| `mcp.IntegerExpression` | `integer` or `string` | the integer or evaluated expression, as an `int` | +| `Annotated[mcp.IntegerExpression, mcp.RelativeTo("address")]` | `integer` or `string` | the same, with `$here` set to an earlier integer parameter's value, or its default when it is absent, or 0 when it has none | +| a Binary Ninja enum, such as `SectionSemantics` | `string` with the member names | the enum member | +| `Literal["a", "b"]` | `string` with those choices | the string | +| `List[T]` | `array` of `T` | a list | +| `Annotated[T, mcp.Schema({...})]` | the schema you supply | the decoded JSON, unchecked | + +`Optional[T]`, `T | None` or a default value makes a parameter optional. A default value also appears in the schema. + +Return a `str` for text, a `dict` for structured content, `None` for an empty result, or an `mcp.ToolResult` to combine text and structured content, and call its `add_warning(code, message)` for a warning. Raise `mcp.ToolError(code, message, details)` to return an error result. Any other exception, or a result that cannot be serialized as JSON, produces an `internal_error` result, and its traceback goes to the log. + +The decorator also accepts `name` (defaults to the function's name), `title`, `scope` (`McpToolScope.GlobalScope` for a global tool), `read_only`, `destructive`, `idempotent`, `open_world` and `output_schema`. + +To supply a schema yourself, use `mcp.register_tool(name, description, input_schema, handler)`, where `handler(call, arguments)` receives the decoded arguments object. + +`mcp.Tool` lists every registered tool, and `mcp.Tool.by_name(name).invoke(arguments, view)` runs one as an MCP server would, which is useful in tests. + +## C++ + +Include `mcp.h` and describe each parameter with the builder. Each parameter's declaration produces both its schema and its conversion, and the handler's signature is checked against the parameters at compile time. + +```cpp +#include "mcp.h" + +using namespace BinaryNinja; +using namespace BinaryNinja::MCP; + +extern "C" BINARYNINJAPLUGIN bool CorePluginInit() +{ + MakeTool({"myplugin_set_comment", "Set Comment", "Set the comment at an address in the active binary view.", + Scope::BinaryView, IdempotentHint}) + .Param(Address("address", "Address expression of the comment.")) + .Param(String("text", "Comment text.")) + .Param(Bool("replace", "Replace an existing comment.").Default(true)) + .Register([](ToolCall& call, uint64_t address, const std::string& text, bool replace) { + Ref view = call.GetBinaryView(); + if (!replace && !view->GetCommentForAddress(address).empty()) + return ToolResult::Error("comment_exists", "The address already has a comment"); + view->SetCommentForAddress(address, text); + return ToolResult::Text("Comment set"); + }); + return true; +} +``` + +| Parameter | Schema | The handler receives | +| --- | --- | --- | +| `String` | `string` | `std::string`. `.NonEmpty()` rejects an empty string. | +| `Bool` | `boolean` | `bool` | +| `UInt` | `integer` | `uint64_t`. `.Maximum(n)` rejects larger values and `.ClampTo(n)` caps them. | +| `Address` | `string` | `uint64_t`, the evaluated address | +| `IntegerExpression` | `integer` or `string` | `uint64_t`, the integer or evaluated expression. `.RelativeTo("address")` sets `$here` to an earlier integer parameter's value, or its default when it is absent, such as for a length measured from an address. | +| `Int` | `integer` | `int64_t`. `.Minimum(n)` and `.Maximum(n)` reject values outside them, and `.ClampTo(min, max)` moves them to the nearest bound. | +| `Number` | `number` | `double` | +| `List` | `array` of `Item` | `std::vector`, such as `std::vector` for `List`. Pass an item, such as `Choice("levels", "", {...})`, for a kind that needs settings. | +| `Choice` | `string` with fixed choices | `std::string` | +| `Enum` | `string` with the listed core enum values | `T` | +| `JsonValue` | the schema fragment you supply | `const rapidjson::Value*` | + +`.Optional()` makes the handler receive `std::optional`, empty when the argument is absent or null. `.Default(value)` makes it receive `value` when the argument is absent or null, and advertises the default in the schema. `JsonValue` has no default, and its `.Optional()` keeps the handler's `const rapidjson::Value*`, which is null when the argument is absent or null. + +Return `ToolResult::Text`, `ToolResult::Structured(rapidjson::Value)` or `ToolResult::Error(code, message, details)`, and use `AddText` to combine text with structured content. A handler that throws produces an `internal_error` result. + +`Register` returns null when the tool is rejected, such as for two parameters with the same name or a `RelativeTo` that does not name an earlier integer parameter. + +To supply a schema yourself, call `RegisterTool(definition, inputSchema, handler)` with a handler taking `(ToolCall&, const rapidjson::Value& arguments)`. + +## Rust + +Implement `McpTool` and pass it to `register_mcp_tool`: + +```rust +use binaryninja::mcp::tool::*; +use serde_json::{json, Value}; + +struct FunctionCount; + +impl McpTool for FunctionCount { + fn definition(&self) -> McpToolDefinition { + McpToolInfo::new("myplugin_function_count", "Count the functions in the active binary view.") + .with_annotations(McpToolAnnotations::READ_ONLY) + .with_input_schema(empty_input_schema()) + } + + fn invoke(&self, call: &McpToolCall, _arguments: Value) -> Result { + let view = call.require_binary_view()?; + Ok(McpToolResult::structured(json!({ "count": view.functions().len() }))) + } +} + +#[no_mangle] +#[allow(non_snake_case)] +pub extern "C" fn CorePluginInit() -> bool { + register_mcp_tool(FunctionCount).is_some() +} +``` + +With the crate's `schemars` feature enabled, implement `TypedMcpTool` instead and pass it to `register_typed_mcp_tool`. The input schema and argument parsing then come from the arguments type, which derives `serde::Deserialize` and `schemars::JsonSchema`, so your plugin depends on `serde` with its `derive` feature and on `schemars` 1. That type is a struct with named fields, or with none for a tool that takes no arguments, since the input schema must be an object. Doc comments become parameter descriptions. `Option` fields and `#[serde(default)]` make parameters optional, and undeclared arguments are rejected. `McpAddress` and `McpIntegerExpression` hold address and integer expressions until you `resolve` them against the call. `McpNonEmptyString` rejects an empty string. + +```rust +#[derive(serde::Deserialize, schemars::JsonSchema)] +struct CommentArgs { + /// Address expression of the comment. + address: McpAddress, + /// Comment text. + text: String, +} + +struct SetComment; + +impl TypedMcpTool for SetComment { + type Args = CommentArgs; + + fn info(&self) -> McpToolInfo { + McpToolInfo::new("myplugin_set_comment", "Set the comment at an address in the active binary view.") + } + + fn invoke(&self, call: &McpToolCall, args: CommentArgs) -> Result { + let address = args.address.resolve(call)?; + call.require_binary_view()?.set_comment_at(address, &args.text); + Ok(McpToolResult::text("Comment set")) + } +} + +#[no_mangle] +#[allow(non_snake_case)] +pub extern "C" fn CorePluginInit() -> bool { + register_typed_mcp_tool(SetComment).is_some() +} +``` + +A tool that panics produces an `internal_error` result when the plugin is built with `panic = "unwind"`. Under `panic = "abort"`, which the Rust plugins bundled with Binary Ninja use, a panic ends the process. diff --git a/docs/guide/mcp.md b/docs/guide/mcp.md index 4ae8a03a1d..3c680419db 100644 --- a/docs/guide/mcp.md +++ b/docs/guide/mcp.md @@ -22,6 +22,8 @@ The exact tool list may change as the MCP server develops, but both server varia Use your MCP client's tool listing UI or command to see the complete set of tools available in your installed Binary Ninja version. +Plugins can add their own tools, which both server variants offer alongside the built-in ones. See [Writing MCP Tools](../dev/mcp-tools.md). + ## Tool Calling Conventions The MCP server exposes Binary Ninja state through a small set of identifiers and conventions. These are worth understanding because they differ from many REST APIs. diff --git a/mcp.cpp b/mcp.cpp new file mode 100644 index 0000000000..85eb2de4d9 --- /dev/null +++ b/mcp.cpp @@ -0,0 +1,858 @@ +// Copyright (c) 2026 Vector 35 Inc +// +// Permission is hereby granted, free of charge, to any person obtaining a copy +// of this software and associated documentation files (the "Software"), to +// deal in the Software without restriction, including without limitation the +// rights to use, copy, modify, merge, publish, distribute, sublicense, and/or +// sell copies of the Software, and to permit persons to whom the Software is +// furnished to do so, subject to the following conditions: +// +// The above copyright notice and this permission notice shall be included in +// all copies or substantial portions of the Software. +// +// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING +// FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS +// IN THE SOFTWARE. + +#include "mcp.h" + +#include "mcpserver.h" +#include "rapidjsonwrapper.h" + +#include +#include +#include +#include +#include +#include + +using namespace BinaryNinja; +using namespace BinaryNinja::MCP; + +namespace { + +Ref GetMcpLogger() +{ + static Ref logger = LogRegistry::CreateLogger("MCP"); + return logger; +} + +std::string TakeString(char* text) +{ + std::string result = text ? text : ""; + BNFreeString(text); + return result; +} + +ArgumentResult ParseWith(bool (*parse)(BNMcpToolCall*, const char*, uint64_t*, uint64_t, char**), + BNMcpToolCall* call, const rapidjson::Value& value, uint64_t here) +{ + std::string json = detail::SerializeJson(value); + uint64_t result = 0; + char* error = nullptr; + if (!parse(call, json.c_str(), &result, here, &error)) + return bn::base::unexpected(TakeString(error)); + + return result; +} + +void InvokeHandler(void* ctxt, BNMcpToolCall* callHandle, const char* arguments, BNMcpToolResult* result) +{ + auto& handler = *static_cast(ctxt); + Ref call = new ToolCall(BNNewMcpToolCallReference(callHandle)); + rapidjson::Document document; + try + { + document.Parse(arguments); + } + catch (const ParseException& e) + { + ToolResult::Error("invalid_params", fmt::format("Arguments are not valid JSON: {}", e.what())).ApplyTo(result); + return; + } + + std::optional toolResult; + try + { + toolResult = handler(*call, document); + } + catch (const std::exception& e) + { + GetMcpLogger()->LogErrorForExceptionF(e, "MCP tool failed: {}", e.what()); + toolResult = ToolResult::Error("internal_error", e.what()); + } + catch (...) + { + GetMcpLogger()->LogError("MCP tool failed with an unknown exception"); + toolResult = ToolResult::Error("internal_error", "The tool failed with an unknown exception"); + } + toolResult->ApplyTo(result); +} + +Ref CreateToolUsing(const ToolDefinition& definition, const std::string& inputSchema, ToolHandler handler, + BNMcpTool* (*create)(const BNMcpToolDefinition*, const BNMcpToolCallbacks*)) +{ + BNMcpToolDefinition apiDefinition = {}; + apiDefinition.name = definition.name.c_str(); + apiDefinition.title = definition.title.c_str(); + apiDefinition.description = definition.description.c_str(); + apiDefinition.inputSchema = inputSchema.c_str(); + apiDefinition.outputSchema = definition.outputSchema.empty() ? nullptr : definition.outputSchema.c_str(); + apiDefinition.scope = static_cast(definition.scope); + apiDefinition.annotations = definition.annotations; + + auto* context = new ToolHandler(std::move(handler)); + BNMcpToolCallbacks callbacks = {}; + callbacks.context = context; + callbacks.invoke = InvokeHandler; + callbacks.freeObject = [](void* ctxt) { delete static_cast(ctxt); }; + BNMcpTool* tool = create(&apiDefinition, &callbacks); + if (!tool) + { + delete context; + return nullptr; + } + + return new Tool(tool); +} +} // namespace + + +ToolCall::ToolCall(BNMcpToolCall* call) +{ + m_object = call; +} + + +Ref ToolCall::GetBinaryView() const +{ + BNBinaryView* view = BNGetMcpToolCallBinaryView(m_object); + return view ? new BinaryView(view) : nullptr; +} + + +bool ToolCall::IsCancelled() const +{ + return BNIsMcpToolCallCancelled(m_object); +} + + +void ToolCall::ReportProgress(double progress, double total, const std::string& message) const +{ + BNReportMcpToolCallProgress(m_object, progress, total, message.c_str()); +} + + +ArgumentResult ToolCall::ParseAddress(const rapidjson::Value& value, uint64_t here) const +{ + return ParseWith(BNParseMcpToolCallAddress, m_object, value, here); +} + + +ArgumentResult ToolCall::ParseInteger(const rapidjson::Value& value, uint64_t here) const +{ + return ParseWith(BNParseMcpToolCallInteger, m_object, value, here); +} + + +ToolResult ToolResult::Text(std::string text) +{ + ToolResult result; + result.m_text.push_back(std::move(text)); + return result; +} + + +ToolResult ToolResult::Structured(const rapidjson::Value& content) +{ + ToolResult result; + result.m_structuredContent = detail::SerializeJson(content); + return result; +} + + +ToolResult ToolResult::Error(std::string code, std::string message, const rapidjson::Value* details) +{ + ToolResult result; + result.m_error = ErrorInfo {std::move(code), std::move(message), + details ? std::optional(detail::SerializeJson(*details)) : std::nullopt}; + return result; +} + + +ToolResult& ToolResult::AddText(std::string text) +{ + m_text.push_back(std::move(text)); + return *this; +} + + +ToolResult& ToolResult::AddWarning(std::string code, std::string message) +{ + m_warnings.emplace_back(std::move(code), std::move(message)); + return *this; +} + + +void ToolResult::ApplyTo(BNMcpToolResult* result) const +{ + if (m_error) + { + BNSetMcpToolResultError(result, m_error->code.c_str(), m_error->message.c_str(), + m_error->details ? m_error->details->c_str() : nullptr); + } + else if (m_structuredContent && !BNSetMcpToolResultStructuredContent(result, m_structuredContent->c_str())) + { + BNSetMcpToolResultError( + result, "internal_error", "The tool produced structured content that is not a JSON object", nullptr); + return; + } + + for (const auto& text : m_text) + BNAddMcpToolResultText(result, text.c_str()); + for (const auto& [code, message] : m_warnings) + BNAddMcpToolResultWarning(result, code.c_str(), message.c_str()); +} + + +std::string ToolResult::ToJson() const +{ + BNMcpToolResult* result = BNCreateMcpToolResult(); + ApplyTo(result); + std::string json = TakeString(BNGetMcpToolResultJson(result)); + BNFreeMcpToolResult(result); + return json; +} + + +Tool::Tool(BNMcpTool* tool) +{ + m_object = tool; +} + + +std::vector> Tool::GetList() +{ + size_t count = 0; + BNMcpTool** tools = BNGetMcpToolList(&count); + + std::vector> result; + result.reserve(count); + for (size_t i = 0; i < count; i++) + result.push_back(new Tool(BNNewMcpToolReference(tools[i]))); + + BNFreeMcpToolList(tools, count); + return result; +} + + +Ref Tool::GetByName(const std::string& name) +{ + BNMcpTool* tool = BNGetMcpToolByName(name.c_str()); + return tool ? new Tool(tool) : nullptr; +} + + +std::string Tool::GetName() const +{ + return TakeString(BNGetMcpToolName(m_object)); +} + + +std::string Tool::GetTitle() const +{ + return TakeString(BNGetMcpToolTitle(m_object)); +} + + +std::string Tool::GetDescription() const +{ + return TakeString(BNGetMcpToolDescription(m_object)); +} + + +std::string Tool::GetInputSchema() const +{ + return TakeString(BNGetMcpToolInputSchema(m_object)); +} + + +std::string Tool::GetOutputSchema() const +{ + return TakeString(BNGetMcpToolOutputSchema(m_object)); +} + + +Scope Tool::GetScope() const +{ + return static_cast(BNGetMcpToolScope(m_object)); +} + + +uint32_t Tool::GetAnnotations() const +{ + return BNGetMcpToolAnnotations(m_object); +} + + +std::string BinaryNinja::MCP::InvokeTool( + Tool& tool, ToolCallHost& host, const std::string& arguments, const ToolResult* error) +{ + BNMcpToolCallCallbacks callbacks = {}; + callbacks.context = &host; + callbacks.getBinaryView = [](void* ctxt) -> BNBinaryView* { + try + { + Ref view = static_cast(ctxt)->GetBinaryView(); + return view ? BNNewViewReference(view->GetObject()) : nullptr; + } + catch (const std::exception& e) + { + GetMcpLogger()->LogErrorForExceptionF(e, "MCP host failed to supply a binary view: {}", e.what()); + return nullptr; + } + }; + callbacks.isCancelled = [](void* ctxt) { + try + { + return static_cast(ctxt)->IsCancelled(); + } + catch (const std::exception& e) + { + GetMcpLogger()->LogErrorForExceptionF(e, "MCP host failed to report cancellation: {}", e.what()); + return false; + } + }; + callbacks.reportProgress = [](void* ctxt, double progress, double total, const char* message) { + try + { + static_cast(ctxt)->ReportProgress(progress, total, message); + } + catch (const std::exception& e) + { + GetMcpLogger()->LogErrorForExceptionF(e, "MCP host failed to receive progress: {}", e.what()); + } + }; + + BNMcpToolCall* call = BNCreateMcpToolCall(&callbacks); + BNMcpToolResult* result = BNCreateMcpToolResult(); + if (error) + error->ApplyTo(result); + BNInvokeMcpTool(tool.GetObject(), call, arguments.c_str(), result); + std::string json = TakeString(BNGetMcpToolResultJson(result)); + BNFreeMcpToolResult(result); + BNFreeMcpToolCall(call); + return json; +} + + +Ref BinaryNinja::MCP::RegisterTool( + const ToolDefinition& definition, const std::string& inputSchema, ToolHandler handler) +{ + return CreateToolUsing(definition, inputSchema, std::move(handler), BNRegisterMcpTool); +} + + +Ref BinaryNinja::MCP::CreateTool(ToolSpec spec) +{ + return CreateToolUsing(spec.definition, spec.inputSchema, std::move(spec.handler), BNCreateMcpTool); +} + + +void detail::AddProperty( + rapidjson::Value& properties, std::string_view name, rapidjson::Value& schema, Allocator& allocator) +{ + properties.AddMember(rapidjson::Value(name.data(), name.size(), allocator), schema, allocator); +} + + +rapidjson::Value detail::TypeSchema(const char* type, Allocator& allocator) +{ + rapidjson::Value schema(rapidjson::kObjectType); + schema.AddMember("type", rapidjson::StringRef(type), allocator); + return schema; +} + + +void detail::AddDescription(rapidjson::Value& schema, std::string_view description, Allocator& allocator) +{ + if (description.empty()) + return; + + schema.AddMember("description", rapidjson::Value(description.data(), description.size(), allocator), allocator); +} + + +std::string detail::MissingMessage(std::string_view kind, std::string_view name) +{ + return fmt::format("Expected {} parameter '{}'", kind, name); +} + + +std::string detail::SerializeJson(const rapidjson::Value& value) +{ + rapidjson::StringBuffer buffer; + rapidjson::Writer writer(buffer); + value.Accept(writer); + return std::string(buffer.GetString(), buffer.GetSize()); +} + + +std::optional detail::FindUnexpectedArgument( + const rapidjson::Value& arguments, const std::vector& names) +{ + for (const auto& member : arguments.GetObj()) + { + std::string_view name(member.name.GetString(), member.name.GetStringLength()); + if (std::ranges::find(names, name) == names.end()) + return fmt::format("Unexpected parameter '{}'", name); + } + return std::nullopt; +} + + +void detail::LogRejectedTool(std::string_view name, std::string_view reason) +{ + GetMcpLogger()->LogErrorF("Rejected MCP tool '{}': {}", name, reason); +} + + +detail::InputSchemaBuilder::InputSchemaBuilder(std::string tool): + m_tool(std::move(tool)), m_schema(rapidjson::kObjectType), m_properties(rapidjson::kObjectType), + m_required(rapidjson::kArrayType) +{ +} + + +void detail::InputSchemaBuilder::Add( + const std::string& name, bool required, const std::string* relativeTo, bool integer) +{ + if (std::ranges::find(m_names, name) != m_names.end()) + throw std::invalid_argument(fmt::format("MCP tool '{}' declares parameter '{}' more than once", m_tool, name)); + + if (relativeTo && std::ranges::find(m_integerNames, *relativeTo) == m_integerNames.end()) + { + throw std::invalid_argument(fmt::format( + "MCP tool '{}' parameter '{}' is relative to '{}', which must be an integer parameter declared before it", + m_tool, name, *relativeTo)); + } + + m_names.push_back(name); + if (integer) + m_integerNames.push_back(name); + + if (required) + m_required.PushBack(rapidjson::Value(name.c_str(), name.size(), GetAllocator()), GetAllocator()); +} + + +std::string detail::InputSchemaBuilder::Build() +{ + auto& allocator = GetAllocator(); + m_schema.AddMember("type", "object", allocator); + m_schema.AddMember("properties", m_properties, allocator); + if (!m_required.Empty()) + m_schema.AddMember("required", m_required, allocator); + m_schema.AddMember("additionalProperties", false, allocator); + return SerializeJson(m_schema); +} + + +String& String::NonEmpty() +{ + m_nonEmpty = true; + return *this; +} + + +rapidjson::Value String::Schema(detail::Allocator& allocator) const +{ + rapidjson::Value schema = detail::TypeSchema("string", allocator); + if (m_nonEmpty) + schema.AddMember("minLength", 1, allocator); + + detail::AddDescription(schema, m_description, allocator); + return schema; +} + + +rapidjson::Value String::DefaultJson(std::string_view value, detail::Allocator& allocator) const +{ + return rapidjson::Value(value.data(), value.size(), allocator); +} + + +std::string String::MissingMessage() const +{ + return detail::MissingMessage(m_nonEmpty ? "non-empty string" : "string", m_name); +} + + +ArgumentResult String::Convert(ToolCall&, const rapidjson::Value& value, const ConvertedArguments&) const +{ + if (!value.IsString() || (m_nonEmpty && value.GetStringLength() == 0)) + return bn::base::unexpected(MissingMessage()); + + return std::string(value.GetString(), value.GetStringLength()); +} + + +Choice::Choice(std::string name, std::string description, std::vector choices): + ValueParam(std::move(name), std::move(description)), m_choices(std::move(choices)) +{ +} + + +rapidjson::Value Choice::Schema(detail::Allocator& allocator) const +{ + rapidjson::Value schema = detail::TypeSchema("string", allocator); + rapidjson::Value choices(rapidjson::kArrayType); + for (const auto& choice : m_choices) + choices.PushBack(rapidjson::Value(choice.c_str(), choice.size(), allocator), allocator); + schema.AddMember("enum", choices, allocator); + detail::AddDescription(schema, m_description, allocator); + return schema; +} + + +rapidjson::Value Choice::DefaultJson(std::string_view value, detail::Allocator& allocator) const +{ + return rapidjson::Value(value.data(), value.size(), allocator); +} + + +std::string Choice::MissingMessage() const +{ + return detail::MissingMessage("string", m_name); +} + + +ArgumentResult Choice::Convert(ToolCall&, const rapidjson::Value& value, const ConvertedArguments&) const +{ + if (!value.IsString()) + return bn::base::unexpected(MissingMessage()); + + std::string text(value.GetString(), value.GetStringLength()); + if (std::find(m_choices.begin(), m_choices.end(), text) == m_choices.end()) + return bn::base::unexpected(fmt::format("Invalid enum value for parameter '{}'", m_name)); + return text; +} + + +rapidjson::Value Bool::Schema(detail::Allocator& allocator) const +{ + rapidjson::Value schema = detail::TypeSchema("boolean", allocator); + detail::AddDescription(schema, m_description, allocator); + return schema; +} + + +rapidjson::Value Bool::DefaultJson(bool value, detail::Allocator&) const +{ + return rapidjson::Value(value); +} + + +std::string Bool::MissingMessage() const +{ + return detail::MissingMessage("boolean", m_name); +} + + +ArgumentResult Bool::Convert(ToolCall&, const rapidjson::Value& value, const ConvertedArguments&) const +{ + if (!value.IsBool()) + return bn::base::unexpected(MissingMessage()); + + return value.GetBool(); +} + + +UInt& UInt::Maximum(uint64_t maximum) +{ + m_maximum = maximum; + m_clamp = false; + return *this; +} + + +UInt& UInt::ClampTo(uint64_t maximum) +{ + m_maximum = maximum; + m_clamp = true; + return *this; +} + + +rapidjson::Value UInt::Schema(detail::Allocator& allocator) const +{ + rapidjson::Value schema = detail::TypeSchema("integer", allocator); + schema.AddMember("minimum", 0, allocator); + if (m_maximum) + schema.AddMember("maximum", rapidjson::Value().SetUint64(*m_maximum), allocator); + + detail::AddDescription(schema, m_description, allocator); + return schema; +} + + +rapidjson::Value UInt::DefaultJson(uint64_t value, detail::Allocator&) const +{ + return rapidjson::Value(value); +} + + +std::string UInt::MissingMessage() const +{ + return detail::MissingMessage("unsigned integer", m_name); +} + + +ArgumentResult UInt::Convert(ToolCall&, const rapidjson::Value& value, const ConvertedArguments&) const +{ + if (!value.IsUint64()) + return bn::base::unexpected(MissingMessage()); + + uint64_t result = value.GetUint64(); + if (m_maximum && result > *m_maximum) + { + if (m_clamp) + return *m_maximum; + + return bn::base::unexpected( + fmt::format("Invalid unsigned integer parameter '{}': Must be at most {}", m_name, *m_maximum)); + } + return result; +} + + +Int& Int::Minimum(int64_t minimum) +{ + m_minimum = minimum; + m_clamp = false; + return *this; +} + + +Int& Int::Maximum(int64_t maximum) +{ + m_maximum = maximum; + m_clamp = false; + return *this; +} + + +Int& Int::ClampTo(int64_t minimum, int64_t maximum) +{ + m_minimum = minimum; + m_maximum = maximum; + m_clamp = true; + return *this; +} + + +rapidjson::Value Int::Schema(detail::Allocator& allocator) const +{ + rapidjson::Value schema = detail::TypeSchema("integer", allocator); + if (m_minimum) + schema.AddMember("minimum", rapidjson::Value().SetInt64(*m_minimum), allocator); + if (m_maximum) + schema.AddMember("maximum", rapidjson::Value().SetInt64(*m_maximum), allocator); + detail::AddDescription(schema, m_description, allocator); + return schema; +} + + +rapidjson::Value Int::DefaultJson(int64_t value, detail::Allocator&) const +{ + return rapidjson::Value(value); +} + + +std::string Int::MissingMessage() const +{ + return detail::MissingMessage("integer", m_name); +} + + +ArgumentResult Int::Convert(ToolCall&, const rapidjson::Value& value, const ConvertedArguments&) const +{ + if (!value.IsInt64()) + return bn::base::unexpected(MissingMessage()); + + int64_t result = value.GetInt64(); + if (m_minimum && result < *m_minimum) + { + if (m_clamp) + return *m_minimum; + return bn::base::unexpected( + fmt::format("Invalid integer parameter '{}': Must be at least {}", m_name, *m_minimum)); + } + if (m_maximum && result > *m_maximum) + { + if (m_clamp) + return *m_maximum; + return bn::base::unexpected( + fmt::format("Invalid integer parameter '{}': Must be at most {}", m_name, *m_maximum)); + } + return result; +} + + +rapidjson::Value Number::Schema(detail::Allocator& allocator) const +{ + rapidjson::Value schema = detail::TypeSchema("number", allocator); + detail::AddDescription(schema, m_description, allocator); + return schema; +} + + +rapidjson::Value Number::DefaultJson(double value, detail::Allocator&) const +{ + return rapidjson::Value(value); +} + + +std::string Number::MissingMessage() const +{ + return detail::MissingMessage("number", m_name); +} + + +ArgumentResult Number::Convert(ToolCall&, const rapidjson::Value& value, const ConvertedArguments&) const +{ + if (!value.IsNumber()) + return bn::base::unexpected(MissingMessage()); + + return value.GetDouble(); +} + + +rapidjson::Value Address::Schema(detail::Allocator& allocator) const +{ + rapidjson::Value schema = detail::TypeSchema("string", allocator); + detail::AddDescription(schema, m_description, allocator); + return schema; +} + + +rapidjson::Value Address::DefaultJson(uint64_t value, detail::Allocator& allocator) const +{ + std::string text = fmt::format("{:#x}", value); + return rapidjson::Value(text.c_str(), text.size(), allocator); +} + + +std::string Address::MissingMessage() const +{ + return detail::MissingMessage("address expression", m_name); +} + + +ArgumentResult Address::Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments&) const +{ + auto address = call.ParseAddress(value); + if (!address) + return bn::base::unexpected( + fmt::format("Invalid address expression parameter '{}': {}", m_name, address.error())); + + return address; +} + + +rapidjson::Value IntegerExpression::Schema(detail::Allocator& allocator) const +{ + rapidjson::Value schema(rapidjson::kObjectType); + rapidjson::Value types(rapidjson::kArrayType); + types.PushBack(rapidjson::StringRef("integer"), allocator); + types.PushBack(rapidjson::StringRef("string"), allocator); + schema.AddMember("type", types, allocator); + schema.AddMember("minimum", 0, allocator); + detail::AddDescription(schema, m_description, allocator); + return schema; +} + + +rapidjson::Value IntegerExpression::DefaultJson(uint64_t value, detail::Allocator&) const +{ + return rapidjson::Value(value); +} + + +std::string IntegerExpression::MissingMessage() const +{ + return detail::MissingMessage("unsigned integer", m_name); +} + + +IntegerExpression& IntegerExpression::RelativeTo(std::string name) +{ + m_relativeTo = std::move(name); + return *this; +} + + +ArgumentResult IntegerExpression::Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments& converted) const +{ + uint64_t here = m_relativeTo ? converted.Get(*m_relativeTo).value_or(0) : 0; + auto integer = call.ParseInteger(value, here); + if (!integer) + return bn::base::unexpected( + fmt::format("Invalid unsigned integer parameter '{}': {}", m_name, integer.error())); + return integer; +} + + +JsonValue::JsonValue(std::string name, std::string schema): Param(std::move(name), ""), m_schema(std::move(schema)) {} + + +JsonValue JsonValue::Optional() const +{ + JsonValue result = *this; + result.m_required = false; + return result; +} + + +void JsonValue::AddSchema(rapidjson::Value& properties, detail::Allocator& allocator) const +{ + rapidjson::Document fragment; + try + { + fragment.Parse(m_schema.c_str()); + } + catch (const ParseException&) + { + throw std::invalid_argument(fmt::format("The schema for MCP tool parameter '{}' is not valid JSON", m_name)); + } + rapidjson::Value schema(fragment, allocator); + detail::AddProperty(properties, m_name, schema, allocator); +} + + +ArgumentResult JsonValue::Absent() const +{ + if (m_required) + return bn::base::unexpected(MissingMessage()); + + return nullptr; +} + + +std::string JsonValue::MissingMessage() const +{ + return fmt::format("Expected parameter '{}'", m_name); +} + + +ArgumentResult JsonValue::Convert( + ToolCall&, const rapidjson::Value& value, const ConvertedArguments&) const +{ + return &value; +} diff --git a/mcp.h b/mcp.h new file mode 100644 index 0000000000..f76d4d26ee --- /dev/null +++ b/mcp.h @@ -0,0 +1,1018 @@ +// Copyright (c) 2026 Vector 35 Inc +// +// Permission is hereby granted, free of charge, to any person obtaining a copy +// of this software and associated documentation files (the "Software"), to +// deal in the Software without restriction, including without limitation the +// rights to use, copy, modify, merge, publish, distribute, sublicense, and/or +// sell copies of the Software, and to permit persons to whom the Software is +// furnished to do so, subject to the following conditions: +// +// The above copyright notice and this permission notice shall be included in +// all copies or substantial portions of the Software. +// +// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING +// FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS +// IN THE SOFTWARE. + +#pragma once + +#include "base/expected.h" +#include "binaryninjaapi.h" +#include "rapidjsonwrapper.h" + +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace BinaryNinja::MCP { + +enum class Scope : uint8_t +{ + Global = GlobalScope, + BinaryView = BinaryViewScope, +}; + +struct ToolDefinition +{ + std::string name; + std::string title; + std::string description; + Scope scope = Scope::BinaryView; + /*! BNMcpToolAnnotation flags. */ + uint32_t annotations = 0; + /*! Empty for no output schema. */ + std::string outputSchema; +}; + +/*! A converted argument, or the message explaining why the argument is invalid. */ +template +using ArgumentResult = bn::base::expected; + +/*! One invocation of a tool, as seen by the tool. The call's binary view and cancellation state + are only available while the tool is running. +*/ +class ToolCall: public CoreRefCountObject +{ +public: + explicit ToolCall(BNMcpToolCall* call); + + /*! Never null while a BinaryView-scoped tool is running. */ + Ref GetBinaryView() const; + bool IsCancelled() const; + void ReportProgress(double progress, double total, const std::string& message = "") const; + + /*! An address is an expression string. An integer is an unsigned integer or an expression + string. Expressions are evaluated against the call's binary view, with \c here as the value + of the current address. + */ + ArgumentResult ParseAddress(const rapidjson::Value& value, uint64_t here = 0) const; + ArgumentResult ParseInteger(const rapidjson::Value& value, uint64_t here = 0) const; +}; + +class ToolResult +{ + struct ErrorInfo + { + std::string code; + std::string message; + std::optional details; + }; + + std::vector m_text; + std::optional m_structuredContent; + std::optional m_error; + std::vector> m_warnings; + +public: + /*! An empty, successful result. */ + ToolResult() = default; + + static ToolResult Text(std::string text); + /*! When no text is added, clients that ignore structured content see it serialized as text. */ + static ToolResult Structured(const rapidjson::Value& content); + static ToolResult Error(std::string code, std::string message, const rapidjson::Value* details = nullptr); + + ToolResult& AddText(std::string text); + /*! An advisory warning about the result, such as analysis that has not finished. Clients see it + in the structured content's reserved \c warnings member, and after any text. + */ + ToolResult& AddWarning(std::string code, std::string message); + void ApplyTo(BNMcpToolResult* result) const; + /*! Serialized as an MCP CallToolResult. */ + std::string ToJson() const; +}; + +class Tool: public CoreRefCountObject +{ +public: + explicit Tool(BNMcpTool* tool); + + /*! Sorted by name. */ + static std::vector> GetList(); + static Ref GetByName(const std::string& name); + + std::string GetName() const; + std::string GetTitle() const; + std::string GetDescription() const; + std::string GetInputSchema() const; + /*! Empty when the tool has no output schema. */ + std::string GetOutputSchema() const; + Scope GetScope() const; + uint32_t GetAnnotations() const; +}; + +using ToolHandler = std::function; + +/*! A complete tool that has not been registered or created, as ToolBuilder::Build returns it. */ +struct ToolSpec +{ + ToolDefinition definition; + /*! A JSON Schema object with \c "type": \c "object". */ + std::string inputSchema; + ToolHandler handler; +}; + +/*! Registers a tool for the life of the process. Returns null when the definition is invalid or + its name is already registered. \c inputSchema is a JSON Schema object with \c "type": \c "object". + A handler that throws produces an \c internal_error result. +*/ +Ref RegisterTool(const ToolDefinition& definition, const std::string& inputSchema, ToolHandler handler); + + +/*! The unsigned integer arguments converted so far in a call, in declaration order, for parameters + evaluated relative to an earlier one. +*/ +class ConvertedArguments +{ + std::vector> m_values; + +public: + void Set(std::string_view name, uint64_t value) { m_values.emplace_back(name, value); } + + std::optional Get(std::string_view name) const + { + auto found = std::ranges::find(m_values, name, &std::pair::first); + if (found == m_values.end()) + return std::nullopt; + return found->second; + } +}; + +namespace detail { + +using Allocator = rapidjson::Document::AllocatorType; + +void AddProperty(rapidjson::Value& properties, std::string_view name, rapidjson::Value& schema, Allocator& allocator); +rapidjson::Value TypeSchema(const char* type, Allocator& allocator); +// Call after the type and constraints so that a schema reads as its type, its constraints, then its description. +void AddDescription(rapidjson::Value& schema, std::string_view description, Allocator& allocator); +std::string MissingMessage(std::string_view kind, std::string_view name); +std::string SerializeJson(const rapidjson::Value& value); + +// Whether an argument can be the current address of a later IntegerExpression. +template +constexpr bool IsIntegerValued = std::is_same_v || std::is_same_v + || std::is_same_v> || std::is_same_v>; + +// A negative value wraps, as it does in an expression. +template +std::optional CurrentAddressValue(const T& value) +{ + if constexpr (std::is_same_v || std::is_same_v) + return static_cast(value); + else if constexpr (IsIntegerValued) + return value ? CurrentAddressValue(*value) : std::nullopt; + else + return std::nullopt; +} +// The message for the first argument that is not one of names. +std::optional FindUnexpectedArgument( + const rapidjson::Value& arguments, const std::vector& names); +void LogRejectedTool(std::string_view name, std::string_view reason); + +class InputSchemaBuilder +{ + std::string m_tool; + rapidjson::Document m_schema; + rapidjson::Value m_properties; + rapidjson::Value m_required; + std::vector m_names; + std::vector m_integerNames; + +public: + explicit InputSchemaBuilder(std::string tool); + + // Throws std::invalid_argument for a duplicate name, or for a relativeTo that does not name an earlier integer + // parameter. + void Add(const std::string& name, bool required, const std::string* relativeTo, bool integer); + rapidjson::Value& Properties() { return m_properties; } + Allocator& GetAllocator() { return m_schema.GetAllocator(); } + + const std::vector& GetNames() const { return m_names; } + + // Call once, after every parameter is declared. + std::string Build(); +}; + +// Reads a call's arguments. Read them in declaration order, since a RelativeTo parameter uses an earlier value. +class ArgumentReader +{ + ToolCall& m_call; + const rapidjson::Value& m_arguments; + ConvertedArguments m_converted; + +public: + ArgumentReader(ToolCall& call, const rapidjson::Value& arguments): m_call(call), m_arguments(arguments) {} + + ToolCall& GetCall() const { return m_call; } + + template + ArgumentResult Read(const P& param) + { + // Clients may send null for an optional argument they leave unset. + auto member = m_arguments.FindMember(param.GetName().c_str()); + bool absent = member == m_arguments.MemberEnd() || (!param.IsRequired() && member->value.IsNull()); + ArgumentResult result = + absent ? param.Absent() : param.Convert(m_call, member->value, m_converted); + if (!result) + return result; + + if (std::optional here = CurrentAddressValue(*result)) + m_converted.Set(param.GetName(), *here); + return result; + } +}; + +// Implements Declare, ParseArguments and Finish for a parameter kind that reads a single argument. +template +class SingleProperty +{ + const Derived& AsDerived() const { return static_cast(*this); } + +public: + void Declare(InputSchemaBuilder& schema) const + { + const Derived& param = AsDerived(); + schema.Add( + param.GetName(), param.IsRequired(), param.GetRelativeTo(), IsIntegerValued); + param.AddSchema(schema.Properties(), schema.GetAllocator()); + } + + auto ParseArguments(ArgumentReader& reader) const { return reader.Read(AsDerived()); } + + template + bn::base::expected Finish(ToolCall&, Parsed& parsed) const + { + return std::move(parsed); + } +}; + +// Parses each parameter's arguments in declaration order. Stops at the first invalid argument. +template +ArgumentResult> ParseArguments( + const std::tuple& params, ArgumentReader& reader) +{ + using Parsed = std::tuple; + + std::tuple...> parsed; + std::optional error; + auto parse = [&](auto& slot, const auto& param) { + if (error) + return; + + auto result = param.ParseArguments(reader); + if (result) + slot.emplace(std::move(*result)); + else + error = std::move(result.error()); + }; + return [&](std::index_sequence) -> ArgumentResult { + (parse(std::get(parsed), std::get(params)), ...); + if (error) + return bn::base::unexpected(std::move(*error)); + + return Parsed {std::move(*std::get(parsed))...}; + }(std::index_sequence_for {}); +} + +template +concept ExpectedWithError = std::same_as>; + +template +concept ParseFunction = requires(const Fn& parse, ToolCall& call, Args&... arguments) { + { parse(call, arguments...) } -> ExpectedWithError; +}; + +template +concept LookupFunction = requires(const Fn& lookup, ToolCall& call, Args&... arguments) { + { lookup(call, arguments...) } -> ExpectedWithError; +}; + +// Used as a group's parse function when it has none. Returns the arguments unchanged, as a tuple. +struct KeepArguments +{ + template + ArgumentResult> operator()(ToolCall&, Values&... values) const + { + return std::tuple {std::move(values)...}; + } +}; + +// Used as a group's lookup function when it has none. Returns the parsed value unchanged. +struct KeepParsed +{ + template + bn::base::expected operator()(ToolCall&, Parsed& parsed) const + { + return std::move(parsed); + } +}; + +} // namespace detail + +/*! Parameters for MakeTool. Each parameter kind declares the JSON Schema it advertises and how its + argument is checked and converted before the handler runs. \c Value is the type the handler + receives. A parameter is required unless it is made optional or given a default. + + A kind provides \c Schema, \c Convert and \c MissingMessage, and a \c DefaultJson to support + \c Default. \c Convert turns one argument into a value, and its error is reported as + \c invalid_params. \c converted holds the values of earlier integer arguments, which + IntegerExpression::RelativeTo uses. +*/ +template +class Param: public detail::SingleProperty +{ +protected: + std::string m_name; + std::string m_description; + + const Derived& Self() const { return static_cast(*this); } + +public: + using Value = ValueType; + using Parsed = ValueType; + + Param(std::string name, std::string description): m_name(std::move(name)), m_description(std::move(description)) {} + + const std::string& GetName() const { return m_name; } + bool IsRequired() const { return true; } + /*! The earlier parameter this one is evaluated relative to, if any. */ + const std::string* GetRelativeTo() const { return nullptr; } + + void AddSchema(rapidjson::Value& properties, detail::Allocator& allocator) const + { + rapidjson::Value schema = Self().Schema(allocator); + detail::AddProperty(properties, m_name, schema, allocator); + } + + ArgumentResult Absent() const { return bn::base::unexpected(Self().MissingMessage()); } +}; + +/*! The handler receives std::optional, empty when the argument is absent or null. */ +template +class OptionalParam: public detail::SingleProperty> +{ + Kind m_param; + +public: + using Value = std::optional; + using Parsed = Value; + + explicit OptionalParam(Kind param): m_param(std::move(param)) {} + + const std::string& GetName() const { return m_param.GetName(); } + bool IsRequired() const { return false; } + const std::string* GetRelativeTo() const { return m_param.GetRelativeTo(); } + + void AddSchema(rapidjson::Value& properties, detail::Allocator& allocator) const + { + m_param.AddSchema(properties, allocator); + } + + ArgumentResult Absent() const { return Value {}; } + + ArgumentResult Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments& converted) const + { + auto result = m_param.Convert(call, value, converted); + if (!result) + return bn::base::unexpected(std::move(result.error())); + + return Value {std::move(*result)}; + } +}; + +/*! The handler receives the default when the argument is absent or null. The schema advertises the + default. +*/ +template +class DefaultParam: public detail::SingleProperty> +{ + Kind m_param; + typename Kind::Value m_default; + +public: + using Value = typename Kind::Value; + using Parsed = Value; + + DefaultParam(Kind param, Value defaultValue): m_param(std::move(param)), m_default(std::move(defaultValue)) {} + + const std::string& GetName() const { return m_param.GetName(); } + bool IsRequired() const { return false; } + const std::string* GetRelativeTo() const { return m_param.GetRelativeTo(); } + + void AddSchema(rapidjson::Value& properties, detail::Allocator& allocator) const + { + rapidjson::Value schema = m_param.Schema(allocator); + rapidjson::Value defaultValue = m_param.DefaultJson(m_default, allocator); + schema.AddMember("default", defaultValue, allocator); + detail::AddProperty(properties, m_param.GetName(), schema, allocator); + } + + ArgumentResult Absent() const { return m_default; } + ArgumentResult Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments& converted) const + { + return m_param.Convert(call, value, converted); + } +}; + +/*! A parameter kind that can be made optional or given a default. */ +template +class ValueParam: public Param +{ +public: + using Param::Param; + + OptionalParam Optional() const { return OptionalParam(this->Self()); } + DefaultParam Default(ValueType value) const { return DefaultParam(this->Self(), std::move(value)); } +}; + +/*! A parameter kind that can describe and convert any one JSON value, such as String or UInt, so it + can also describe and convert each item of a List. +*/ +template +concept ParamKind = requires(const Kind& kind, ToolCall& call, const rapidjson::Value& value, + const ConvertedArguments& converted, detail::Allocator& allocator) { + typename Kind::Value; + { kind.Schema(allocator) } -> std::same_as; + { kind.MissingMessage() } -> std::same_as; + { kind.GetRelativeTo() } -> std::same_as; + { kind.Convert(call, value, converted) } -> std::same_as>; +}; + +/*! A parameter for ToolBuilder::Param, or an item kind for List, that accepts a JSON string. */ +class String: public ValueParam +{ + bool m_nonEmpty = false; + +public: + using ValueParam::ValueParam; + + /*! Rejects an empty string. */ + String& NonEmpty(); + + rapidjson::Value Schema(detail::Allocator& allocator) const; + rapidjson::Value DefaultJson(std::string_view value, detail::Allocator& allocator) const; + std::string MissingMessage() const; + ArgumentResult Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments& converted) const; +}; + +/*! A parameter for ToolBuilder::Param, or an item kind for List, that accepts one of a fixed set + of strings. +*/ +class Choice: public ValueParam +{ + std::vector m_choices; + +public: + Choice(std::string name, std::string description, std::vector choices); + + rapidjson::Value Schema(detail::Allocator& allocator) const; + rapidjson::Value DefaultJson(std::string_view value, detail::Allocator& allocator) const; + std::string MissingMessage() const; + ArgumentResult Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments& converted) const; +}; + +/*! A parameter for ToolBuilder::Param, or an item kind for List, that accepts a JSON boolean. */ +class Bool: public ValueParam +{ +public: + using ValueParam::ValueParam; + + rapidjson::Value Schema(detail::Allocator& allocator) const; + rapidjson::Value DefaultJson(bool value, detail::Allocator& allocator) const; + std::string MissingMessage() const; + ArgumentResult Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments& converted) const; +}; + +/*! A parameter for ToolBuilder::Param, or an item kind for List, that accepts a JSON unsigned + integer. +*/ +class UInt: public ValueParam +{ + std::optional m_maximum; + bool m_clamp = false; + +public: + using ValueParam::ValueParam; + + /*! Rejects larger values. */ + UInt& Maximum(uint64_t maximum); + /*! Silently reduces larger values to the maximum. */ + UInt& ClampTo(uint64_t maximum); + + rapidjson::Value Schema(detail::Allocator& allocator) const; + rapidjson::Value DefaultJson(uint64_t value, detail::Allocator& allocator) const; + std::string MissingMessage() const; + ArgumentResult Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments& converted) const; +}; + +/*! A parameter for ToolBuilder::Param, or an item kind for List, that accepts a JSON integer, + which may be negative. +*/ +class Int: public ValueParam +{ + std::optional m_minimum; + std::optional m_maximum; + bool m_clamp = false; + +public: + using ValueParam::ValueParam; + + /*! Rejects smaller values. */ + Int& Minimum(int64_t minimum); + /*! Rejects larger values. */ + Int& Maximum(int64_t maximum); + /*! Silently moves values outside the range to the nearest bound. */ + Int& ClampTo(int64_t minimum, int64_t maximum); + + rapidjson::Value Schema(detail::Allocator& allocator) const; + rapidjson::Value DefaultJson(int64_t value, detail::Allocator& allocator) const; + std::string MissingMessage() const; + ArgumentResult Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments& converted) const; +}; + +/*! A parameter for ToolBuilder::Param, or an item kind for List, that accepts a JSON number. */ +class Number: public ValueParam +{ +public: + using ValueParam::ValueParam; + + rapidjson::Value Schema(detail::Allocator& allocator) const; + rapidjson::Value DefaultJson(double value, detail::Allocator& allocator) const; + std::string MissingMessage() const; + ArgumentResult Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments& converted) const; +}; + +/*! A parameter for ToolBuilder::Param, or an item kind for List, that accepts an address expression + string, evaluated against the call's binary view. +*/ +class Address: public ValueParam +{ +public: + using ValueParam::ValueParam; + + rapidjson::Value Schema(detail::Allocator& allocator) const; + rapidjson::Value DefaultJson(uint64_t value, detail::Allocator& allocator) const; + std::string MissingMessage() const; + ArgumentResult Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments& converted) const; +}; + +/*! A parameter for ToolBuilder::Param, or an item kind for List, that accepts an unsigned integer, + or an expression string evaluated against the call's binary view. +*/ +class IntegerExpression: public ValueParam +{ + std::optional m_relativeTo; + +public: + using ValueParam::ValueParam; + + /*! Evaluates expressions with the value of the earlier \c Address, \c IntegerExpression, \c UInt + or \c Int parameter \c name as the current address. When that argument is absent, its default + is used, or 0 when it has none. Registration fails if \c name is not such a parameter declared + before this one. + */ + IntegerExpression& RelativeTo(std::string name); + const std::string* GetRelativeTo() const { return m_relativeTo ? &*m_relativeTo : nullptr; } + + rapidjson::Value Schema(detail::Allocator& allocator) const; + rapidjson::Value DefaultJson(uint64_t value, detail::Allocator& allocator) const; + std::string MissingMessage() const; + ArgumentResult Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments& converted) const; +}; + +/*! A parameter for ToolBuilder::Param, or an item kind for another List, that accepts a JSON array. + Each item is converted as \c Item would convert a lone argument, such as for + List. Pass the item explicitly for a kind that needs settings, such as a + \c Choice, and name it after the list so that its errors name the parameter. +*/ +template +class List: public ValueParam, std::vector> +{ + using Base = ValueParam, std::vector>; + + Item m_item; + +public: + List(std::string name, std::string description) + requires std::constructible_from + : Base(name, std::move(description)), m_item(std::move(name), "") + { + } + + List(std::string name, std::string description, Item item): + Base(std::move(name), std::move(description)), m_item(std::move(item)) + { + } + + const std::string* GetRelativeTo() const { return m_item.GetRelativeTo(); } + + rapidjson::Value Schema(detail::Allocator& allocator) const + { + rapidjson::Value schema = detail::TypeSchema("array", allocator); + rapidjson::Value items = m_item.Schema(allocator); + schema.AddMember("items", items, allocator); + detail::AddDescription(schema, this->m_description, allocator); + return schema; + } + + rapidjson::Value DefaultJson(const typename Base::Value& value, detail::Allocator& allocator) const + { + rapidjson::Value array(rapidjson::kArrayType); + for (const auto& item : value) + { + rapidjson::Value json = m_item.DefaultJson(item, allocator); + array.PushBack(json, allocator); + } + return array; + } + + std::string MissingMessage() const { return detail::MissingMessage("array", this->m_name); } + + ArgumentResult Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments& converted) const + { + if (!value.IsArray()) + return bn::base::unexpected(MissingMessage()); + + typename Base::Value result; + result.reserve(value.Size()); + for (const auto& json : value.GetArray()) + { + ArgumentResult item = m_item.Convert(call, json, converted); + if (!item) + return bn::base::unexpected(std::move(item.error())); + + result.push_back(std::move(*item)); + } + return result; + } +}; + +/*! A parameter for ToolBuilder::Param, or an item kind for List, that accepts a core enum value by + its enumerator name. Core enums cannot be enumerated, so the values the tool accepts are listed + explicitly. +*/ +template +class Enum: public ValueParam, T> +{ + std::vector m_values; + + static std::string ValueName(T value) + { + return CoreEnumToString(value).value_or(std::to_string(static_cast>(value))); + } + +public: + Enum(std::string name, std::string description, std::initializer_list values): + ValueParam, T>(std::move(name), std::move(description)), m_values(values) + { + } + + rapidjson::Value Schema(detail::Allocator& allocator) const + { + rapidjson::Value schema = detail::TypeSchema("string", allocator); + rapidjson::Value names(rapidjson::kArrayType); + for (T value : m_values) + { + std::string text = ValueName(value); + names.PushBack(rapidjson::Value(text.c_str(), text.size(), allocator), allocator); + } + schema.AddMember("enum", names, allocator); + detail::AddDescription(schema, this->m_description, allocator); + return schema; + } + + rapidjson::Value DefaultJson(T value, detail::Allocator& allocator) const + { + std::string text = ValueName(value); + return rapidjson::Value(text.c_str(), text.size(), allocator); + } + + std::string MissingMessage() const { return detail::MissingMessage("string", this->m_name); } + + ArgumentResult Convert(ToolCall&, const rapidjson::Value& value, const ConvertedArguments&) const + { + if (!value.IsString()) + return bn::base::unexpected(MissingMessage()); + + std::optional parsed = CoreEnumFromString(std::string(value.GetString(), value.GetStringLength())); + if (!parsed || std::find(m_values.begin(), m_values.end(), *parsed) == m_values.end()) + return bn::base::unexpected(fmt::format("Invalid enum value for parameter '{}'", this->m_name)); + + return *parsed; + } +}; + +/*! A parameter for ToolBuilder::Param that accepts any JSON value matching a caller-supplied schema + fragment, which carries its own description. It cannot be an item kind for List. The handler + receives the raw value, or null when an optional argument is absent or null. +*/ +class JsonValue: public Param +{ + std::string m_schema; + bool m_required = true; + +public: + JsonValue(std::string name, std::string schema); + + JsonValue Optional() const; + bool IsRequired() const { return m_required; } + + void AddSchema(rapidjson::Value& properties, detail::Allocator& allocator) const; + ArgumentResult Absent() const; + std::string MissingMessage() const; + ArgumentResult Convert( + ToolCall& call, const rapidjson::Value& value, const ConvertedArguments& converted) const; +}; + +template +class GroupParam; + +/*! Parameters that are declared separately but passed to the handler as one value, such as a + function named by an address and an architecture. Call \c Parse, \c Lookup or both before passing + the group to ToolBuilder::Param. +*/ +template + requires(std::derived_from> && ...) +class Group +{ + std::tuple m_params; + +public: + explicit Group(Params... params): m_params(std::move(params)...) {} + + /*! Sets the function that combines the arguments into one value. It is called as + parse(ToolCall&, Params::Value&...) and returns an ArgumentResult. Its + error is reported as \c invalid_params. + */ + template ParseFn> + GroupParam Parse(ParseFn parse) && + { + return {std::move(m_params), std::move(parse), {}}; + } + + /*! Sets the function that finds what the arguments refer to. It is called as + lookup(ToolCall&, Params::Value&...) and returns a + bn::base::expected. Its error becomes the tool's result. It runs only + after every argument has been parsed. + */ + template LookupFn> + auto Lookup(LookupFn lookup) && + { + using Arguments = std::tuple; + auto lookupArguments = [lookup = std::move(lookup)](ToolCall& call, Arguments& arguments) { + return std::apply([&](auto&... value) { return lookup(call, value...); }, arguments); + }; + return GroupParam { + std::move(m_params), {}, std::move(lookupArguments)}; + } +}; + +/*! Creates a Group with no parameters and the given lookup function. */ +template +auto Lookup(LookupFn lookup) +{ + return Group<>().Lookup(std::move(lookup)); +} + +/*! A Group that has a parse function, a lookup function or both, created by Group::Parse or + Group::Lookup. +*/ +template +class GroupParam +{ + std::tuple m_params; + ParseFn m_parse; + LookupFn m_lookup; + +public: + using Parsed = typename std::invoke_result_t::value_type; + using Value = typename std::invoke_result_t::value_type; + + GroupParam(std::tuple params, ParseFn parse, LookupFn lookup): + m_params(std::move(params)), m_parse(std::move(parse)), m_lookup(std::move(lookup)) + { + } + + /*! Sets the lookup function, which receives the parsed value as lookup(ToolCall&, Parsed&). + Otherwise it is the same as Group::Lookup. + */ + template NextLookupFn> + GroupParam Lookup(NextLookupFn lookup) && + requires std::same_as + { + return {std::move(m_params), std::move(m_parse), std::move(lookup)}; + } + + void Declare(detail::InputSchemaBuilder& schema) const + { + std::apply([&](const auto&... param) { (param.Declare(schema), ...); }, m_params); + } + + ArgumentResult ParseArguments(detail::ArgumentReader& reader) const + { + auto arguments = detail::ParseArguments(m_params, reader); + if (!arguments) + return bn::base::unexpected(std::move(arguments.error())); + + return std::apply([&](auto&... value) { return m_parse(reader.GetCall(), value...); }, *arguments); + } + + bn::base::expected Finish(ToolCall& call, Parsed& parsed) const + { + return m_lookup(call, parsed); + } +}; + +/*! The types that ToolBuilder::Param accepts. These are the parameter kinds such as String, their + Optional and Default forms, and a Group after Group::Parse or Group::Lookup. +*/ +template +concept ToolParam = requires(const P& param, detail::InputSchemaBuilder& schema, detail::ArgumentReader& reader, + ToolCall& call, typename P::Parsed& parsed) { + typename P::Value; + param.Declare(schema); + { param.ParseArguments(reader) } -> std::same_as>; + { param.Finish(call, parsed) } -> std::same_as>; +}; + +template +class ToolBuilder +{ + using Values = std::tuple; + + ToolDefinition m_definition; + std::tuple m_params; + + template + friend class ToolBuilder; + + ToolBuilder(ToolDefinition definition, std::tuple params): + m_definition(std::move(definition)), m_params(std::move(params)) + { + } + + // Parse every argument before running any lookup, so that an invalid argument is reported as invalid_params even + // when a lookup would also fail. + template + static bn::base::expected ExtractAll(const std::tuple& params, ToolCall& call, + const rapidjson::Value& arguments, std::index_sequence) + { + detail::ArgumentReader reader(call, arguments); + auto parsed = detail::ParseArguments(params, reader); + if (!parsed) + return bn::base::unexpected(ToolResult::Error("invalid_params", std::move(parsed.error()))); + + std::tuple...> values; + std::optional error; + auto finish = [&](auto& slot, const auto& param, auto& value) { + if (error) + return; + auto result = param.Finish(call, value); + if (result) + slot.emplace(std::move(*result)); + else + error = std::move(result.error()); + }; + (finish(std::get(values), std::get(params), std::get(*parsed)), ...); + if (error) + return bn::base::unexpected(std::move(*error)); + return Values {std::move(*std::get(values))...}; + } + +public: + explicit ToolBuilder(ToolDefinition definition) + requires(sizeof...(Params) == 0) + : m_definition(std::move(definition)) + { + } + + template + ToolBuilder Param(P param) && + { + return ToolBuilder( + std::move(m_definition), std::tuple_cat(std::move(m_params), std::make_tuple(std::move(param)))); + } + + /*! Passes the builder to \c transform and continues with what it returns, so that a group of + parameters shared by several tools can be added in the middle of a chain, as in + MakeTool(...).With(ListPaginationParams).Param(...). + */ + template + auto With(Transform&& transform) && + { + return std::forward(transform)(std::move(*this)); + } + + ToolBuilder OutputSchema(std::string schema) && + { + m_definition.outputSchema = std::move(schema); + return std::move(*this); + } + + /*! Registers the tool. The handler is called as handler(ToolCall&, Params::Value&...) + only once every argument has been checked and converted. Missing, invalid and undeclared + arguments produce an \c invalid_params result instead. Returns null when the parameters or + the definition are invalid. + */ + template + Ref Register(Handler handler) && + { + std::string name = m_definition.name; + try + { + ToolSpec spec = std::move(*this).Build(std::move(handler)); + return RegisterTool(spec.definition, spec.inputSchema, std::move(spec.handler)); + } + catch (const std::invalid_argument& e) + { + detail::LogRejectedTool(name, e.what()); + return nullptr; + } + } + + /*! Completes the tool without registering it, such as for an MCP server to create with + CreateTool. The handler is called as it is for Register. Throws \c std::invalid_argument when + the parameters are invalid, such as when two share a name or a \c RelativeTo does not name an + earlier integer parameter. + */ + template + ToolSpec Build(Handler handler) && + { + static_assert(std::is_invocable_r_v, + "The handler must accept (ToolCall&, Params::Value&...) and return a ToolResult"); + + detail::InputSchemaBuilder schema(m_definition.name); + std::apply([&](const auto&... param) { (param.Declare(schema), ...); }, m_params); + + std::string inputSchema = schema.Build(); + ToolHandler toolHandler = [params = std::move(m_params), names = schema.GetNames(), + handler = std::move(handler)]( + ToolCall& call, const rapidjson::Value& arguments) mutable -> ToolResult { + if (!arguments.IsObject()) + return ToolResult::Error("invalid_params", "Expected object arguments"); + if (auto unexpected = detail::FindUnexpectedArgument(arguments, names)) + return ToolResult::Error("invalid_params", *unexpected); + + auto values = ExtractAll(params, call, arguments, std::index_sequence_for {}); + if (!values) + return std::move(values.error()); + + return std::apply([&](auto&... value) -> ToolResult { return handler(call, value...); }, *values); + }; + return {std::move(m_definition), std::move(inputSchema), std::move(toolHandler)}; + } +}; + +inline ToolBuilder<> MakeTool(ToolDefinition definition) +{ + return ToolBuilder<>(std::move(definition)); +} +} // namespace BinaryNinja::MCP diff --git a/mcpserver.h b/mcpserver.h new file mode 100644 index 0000000000..b12e297afb --- /dev/null +++ b/mcpserver.h @@ -0,0 +1,54 @@ +// Copyright (c) 2026 Vector 35 Inc +// +// Permission is hereby granted, free of charge, to any person obtaining a copy +// of this software and associated documentation files (the "Software"), to +// deal in the Software without restriction, including without limitation the +// rights to use, copy, modify, merge, publish, distribute, sublicense, and/or +// sell copies of the Software, and to permit persons to whom the Software is +// furnished to do so, subject to the following conditions: +// +// The above copyright notice and this permission notice shall be included in +// all copies or substantial portions of the Software. +// +// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING +// FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS +// IN THE SOFTWARE. + +#pragma once + +#include "binaryninjaapi.h" +#include "mcp.h" + +#include +#include + +namespace BinaryNinja::MCP { + +/*! The request state an MCP server supplies for one invocation of a tool. */ +class ToolCallHost +{ +public: + virtual ~ToolCallHost() = default; + + virtual Ref GetBinaryView() = 0; + virtual bool IsCancelled() { return false; } + virtual void ReportProgress(double progress, double total, std::string_view message) {} +}; + +/*! Creates a tool without registering it. This can be used to expose a tool that is specific to + one MCP server, such as one that manages that server's active binary view. Returns null when the + definition is invalid. +*/ +Ref CreateTool(ToolSpec spec); + +/*! Runs the tool for an MCP server and returns an MCP CallToolResult as JSON. An \c error result, when + given, is returned in place of running the tool, finished as the tool's own error would be. This + lets a server give its own reason for refusing a call, such as why it has no binary view. +*/ +std::string InvokeTool( + Tool& tool, ToolCallHost& host, const std::string& arguments, const ToolResult* error = nullptr); +} // namespace BinaryNinja::MCP diff --git a/python/__init__.py b/python/__init__.py index 9c86099b83..becfbea890 100644 --- a/python/__init__.py +++ b/python/__init__.py @@ -88,6 +88,8 @@ from .stringrecognizer import * from .unicode import * from .similarity import * +# mcp is imported only as a module because names such as Tool and Address are too generic for the binaryninja namespace. +from . import mcp # We import each of these by name to prevent conflicts between # log.py and the function 'log' which we don't import below from .log import ( diff --git a/python/mcp.py b/python/mcp.py new file mode 100644 index 0000000000..8a0021acc7 --- /dev/null +++ b/python/mcp.py @@ -0,0 +1,798 @@ +# Copyright (c) 2026 Vector 35 Inc +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to +# deal in the Software without restriction, including without limitation the +# rights to use, copy, modify, merge, publish, distribute, sublicense, and/or +# sell copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING +# FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS +# IN THE SOFTWARE. +""" +Tools for Binary Ninja's Model Context Protocol (MCP) server. + +A tool registered here is offered by every MCP server in the process, alongside the built-in ``bn_*`` +tools. The simplest way to write one is the :py:func:`tool` decorator, which builds the tool's JSON +Schema from the function's signature and docstring:: + + from binaryninja import mcp + + @mcp.tool(read_only=True) + def myplugin_function_count(call: mcp.ToolCall, min_size: int = 0) -> dict: + \"\"\"Count the functions in the active binary view. + + :param min_size: Only count functions with at least this many bytes. + \"\"\" + functions = [f for f in call.binary_view.functions if f.total_bytes >= min_size] + return {"count": len(functions)} +""" + +import ctypes +import enum +import inspect +import json +import re +import types +import typing +from typing import Any, Callable, Dict, List, Literal, Optional, Tuple, Union + +import binaryninja +from . import _binaryninjacore as core +from . import binaryview +from .enums import McpToolAnnotation, McpToolScope +from .log import log_error_for_exception + +__all__ = [ + "Address", + "ClampTo", + "IntegerExpression", + "Maximum", + "Minimum", + "NonEmpty", + "RelativeTo", + "Schema", + "Tool", + "ToolCall", + "ToolError", + "ToolResult", + "register_tool", + "tool", +] + +class Address(int): + """ + A parameter annotated ``Address`` accepts an address expression string, evaluated against the call's + binary view. The function receives an ``int``. + """ + + __slots__ = () + + +class IntegerExpression(int): + """ + A parameter annotated ``IntegerExpression`` accepts an unsigned integer or an expression string, + evaluated against the call's binary view. The function receives an ``int``. + """ + + __slots__ = () + +_TOOL_NAME = re.compile(r"^[A-Za-z0-9_-]{1,64}$") + +# The origin of a T | None annotation, which is not typing.Union on Python 3.10 and later. +_UnionType = getattr(types, "UnionType", Union) + +# Keeps each registered tool's ctypes callbacks alive until the tool is freed. +_registered_tools: List["_RegisteredTool"] = [] + + +class Schema: + """ + Supplies the JSON Schema for a parameter the other annotations cannot describe. Use it as + ``Annotated[dict, mcp.Schema({...})]``. The function receives the argument as decoded JSON. + """ + + def __init__(self, schema: dict): + self.schema = schema + + +class RelativeTo: + """ + Evaluates an :py:class:`IntegerExpression` parameter's expression with an earlier :py:class:`Address`, + :py:class:`IntegerExpression` or ``int`` parameter's value as ``$here``. When that argument is absent, + its default is used, or 0 when it has none. Use it as + ``Annotated[mcp.IntegerExpression, mcp.RelativeTo("address")]``, for example for a length measured + from an address, or on the items of a ``List`` of them. + """ + + def __init__(self, parameter: str): + self.parameter = parameter + + +class Minimum: + """Rejects an ``int`` parameter's values below ``value``. Use it as ``Annotated[int, mcp.Minimum(0)]``.""" + + def __init__(self, value: int): + self.value = value + + +class Maximum: + """Rejects an ``int`` parameter's values above ``value``. Use it as ``Annotated[int, mcp.Maximum(100)]``.""" + + def __init__(self, value: int): + self.value = value + + +class ClampTo: + """ + Clamps an ``int`` parameter's value to the range ``minimum`` to ``maximum`` instead of rejecting values + outside it. Use it as ``Annotated[int, mcp.ClampTo(0, 1000)]``. + """ + + def __init__(self, minimum: int, maximum: int): + self.minimum = minimum + self.maximum = maximum + + +class NonEmpty: + """Rejects an empty ``str`` parameter. Use it as ``Annotated[str, mcp.NonEmpty()]``.""" + + +class ToolError(Exception): + """Raised by a tool to return an error result with a machine-readable code.""" + + def __init__(self, code: str, message: str, details: Any = None): + super().__init__(message) + self.code = code + self.message = message + self.details = details + + +class ToolResult: + """ + The result of a tool. A tool may instead return a ``str`` (text), a ``dict`` (structured content) or + ``None`` (an empty result). + """ + + def __init__(self, text: Optional[Union[str, List[str]]] = None, structured: Optional[dict] = None): + if text is None: + self.text: List[str] = [] + elif isinstance(text, str): + self.text = [text] + else: + self.text = list(text) + self.structured = structured + self.error: Optional[ToolError] = None + self.warnings: List[Tuple[str, str]] = [] + + def add_warning(self, code: str, message: str) -> "ToolResult": + """ + Adds an advisory warning about the result, such as analysis that has not finished. Clients see it + in the structured content's reserved ``warnings`` member, and after any text. + """ + self.warnings.append((code, message)) + return self + + @staticmethod + def from_error(error: ToolError) -> "ToolResult": + result = ToolResult() + result.error = error + return result + + def _apply(self, handle) -> None: + # Anything that can fail does so before the result is written, so the caller can report an error in + # its place. + if not all(isinstance(text, str) for text in self.text): + raise TypeError("A ToolResult's text must be a str or a list of str") + if not all(isinstance(code, str) and isinstance(message, str) for code, message in self.warnings): + raise TypeError("A ToolResult's warning code and message must be str") + if self.error is not None: + if not isinstance(self.error.code, str) or not isinstance(self.error.message, str): + raise TypeError("A ToolError's code and message must be str") + details = None if self.error.details is None else json.dumps(self.error.details, allow_nan=False) + core.BNSetMcpToolResultError(handle, self.error.code, self.error.message, details) + elif self.structured is not None: + structured = json.dumps(self.structured, allow_nan=False) + if not core.BNSetMcpToolResultStructuredContent(handle, structured): + core.BNSetMcpToolResultError( + handle, "internal_error", "The tool produced structured content that is not a JSON object", None + ) + return + for text in self.text: + core.BNAddMcpToolResultText(handle, text) + for code, message in self.warnings: + core.BNAddMcpToolResultWarning(handle, code, message) + + @staticmethod + def _from_value(value: Any) -> "ToolResult": + if isinstance(value, ToolResult): + return value + if value is None: + return ToolResult() + if isinstance(value, str): + return ToolResult(text=value) + if isinstance(value, dict): + return ToolResult(structured=value) + raise TypeError(f"A tool must return a str, dict, ToolResult or None, not {type(value).__name__}") + + +class ToolCall: + """ + One invocation of a tool. Once the tool returns, ``binary_view`` is ``None`` and ``is_cancelled`` is + ``True``. + """ + + def __init__(self, handle): + self.handle = core.BNNewMcpToolCallReference(handle) + + def __del__(self): + if core is not None: + core.BNFreeMcpToolCall(self.handle) + + @property + def binary_view(self) -> Optional["binaryview.BinaryView"]: + """The binary view the MCP session targets. Never ``None`` while a BinaryView-scoped tool runs.""" + handle = core.BNGetMcpToolCallBinaryView(self.handle) + if not handle: + return None + return binaryview.BinaryView(handle=handle) + + @property + def is_cancelled(self) -> bool: + return core.BNIsMcpToolCallCancelled(self.handle) + + def report_progress(self, progress: float, total: float, message: str = "") -> None: + core.BNReportMcpToolCallProgress(self.handle, progress, total, message) + + def parse_address(self, value: Any, here: int = 0) -> int: + """ + Evaluates an address expression string, with ``here`` as the value of ``$here``. Raises + ``ValueError`` when it is invalid. + """ + return self._parse(core.BNParseMcpToolCallAddress, value, here) + + def parse_integer(self, value: Any, here: int = 0) -> int: + """ + Evaluates an unsigned integer or an expression string, with ``here`` as the value of ``$here``. + Raises ``ValueError`` when it is invalid. + """ + return self._parse(core.BNParseMcpToolCallInteger, value, here) + + def _parse(self, parse, value: Any, here: int) -> int: + result = ctypes.c_uint64() + error = ctypes.c_char_p() + if not parse(self.handle, json.dumps(value), result, here, error): + message = core.pyNativeStr(error.value) if error.value is not None else "" + core.free_string(error) + raise ValueError(message) + return result.value + + +class _ToolMetaclass(type): + def __iter__(cls): + binaryninja._init_plugins() + count = ctypes.c_ulonglong() + tools = core.BNGetMcpToolList(count) + try: + for i in range(count.value): + yield Tool(core.BNNewMcpToolReference(tools[i])) + finally: + core.BNFreeMcpToolList(tools, count.value) + + +class Tool(metaclass=_ToolMetaclass): + """ + A tool in the tool registry. Iterating over ``Tool`` lists every registered tool, sorted by name. + """ + + def __init__(self, handle): + self.handle = handle + + def __del__(self): + if core is not None: + core.BNFreeMcpTool(self.handle) + + @staticmethod + def list() -> List["Tool"]: + return list(Tool) + + @staticmethod + def by_name(name: str) -> Optional["Tool"]: + binaryninja._init_plugins() + handle = core.BNGetMcpToolByName(name) + return Tool(handle) if handle else None + + @property + def name(self) -> str: + return core.BNGetMcpToolName(self.handle) + + @property + def title(self) -> str: + return core.BNGetMcpToolTitle(self.handle) + + @property + def description(self) -> str: + return core.BNGetMcpToolDescription(self.handle) + + @property + def input_schema(self) -> dict: + return json.loads(core.BNGetMcpToolInputSchema(self.handle)) + + @property + def output_schema(self) -> Optional[dict]: + schema = core.BNGetMcpToolOutputSchema(self.handle) + return json.loads(schema) if schema else None + + @property + def scope(self) -> McpToolScope: + return McpToolScope(core.BNGetMcpToolScope(self.handle)) + + @property + def annotations(self) -> McpToolAnnotation: + """A combination of :py:class:`McpToolAnnotation` flags.""" + return McpToolAnnotation(core.BNGetMcpToolAnnotations(self.handle)) + + def invoke(self, arguments: Optional[dict] = None, view: Optional["binaryview.BinaryView"] = None) -> dict: + """ + Runs the tool as an MCP server would, against ``view``, and returns the MCP ``CallToolResult``. + """ + callbacks = core.BNMcpToolCallCallbacks() + callbacks.context = 0 + + def get_binary_view(ctxt): + try: + if view is None: + return None + return ctypes.cast(core.BNNewViewReference(view.handle), ctypes.c_void_p).value + except Exception: + log_error_for_exception("Unhandled Python exception in Tool.invoke") + return None + + callbacks.getBinaryView = callbacks.getBinaryView.__class__(get_binary_view) + callbacks.isCancelled = callbacks.isCancelled.__class__(lambda ctxt: False) + callbacks.reportProgress = callbacks.reportProgress.__class__(lambda ctxt, progress, total, message: None) + + call = core.BNCreateMcpToolCall(callbacks) + result = core.BNCreateMcpToolResult() + try: + core.BNInvokeMcpTool(self.handle, call, json.dumps(arguments or {}), result) + return json.loads(core.BNGetMcpToolResultJson(result)) + finally: + core.BNFreeMcpToolResult(result) + core.BNFreeMcpToolCall(call) + + def __repr__(self): + return f"" + + def __eq__(self, other): + return isinstance(other, Tool) and self.name == other.name + + def __hash__(self): + return hash(self.name) + + +class _RegisteredTool: + def __init__(self, handler: Callable[[ToolCall, dict], Any]): + self.handler = handler + self.callbacks = core.BNMcpToolCallbacks() + self.callbacks.context = 0 + self.callbacks.invoke = self.callbacks.invoke.__class__(self._invoke) + self.callbacks.freeObject = self.callbacks.freeObject.__class__(self._free_object) + + def _free_object(self, ctxt): + try: + _registered_tools.remove(self) + except Exception: + log_error_for_exception("Unhandled Python exception freeing an MCP tool") + + def _invoke(self, ctxt, call_handle, arguments, result_handle): + try: + try: + decoded = json.loads(core.pyNativeStr(arguments)) + result = ToolResult._from_value(self.handler(ToolCall(call_handle), decoded)) + except ToolError as error: + result = ToolResult.from_error(error) + result._apply(result_handle) + except Exception as error: + log_error_for_exception("Unhandled Python exception in MCP tool") + try: + ToolResult.from_error(ToolError("internal_error", str(error)))._apply(result_handle) + except Exception: + log_error_for_exception("Unhandled Python exception applying an MCP tool result") + + +def register_tool( + name: str, + description: str, + input_schema: dict, + handler: Callable[[ToolCall, dict], Any], + *, + title: str = "", + scope: McpToolScope = McpToolScope.BinaryViewScope, + annotations: McpToolAnnotation = McpToolAnnotation(0), + output_schema: Optional[dict] = None, +) -> Tool: + """ + Registers a tool for the life of the process. + + :param handler: Called as ``handler(call, arguments)`` with the decoded arguments object. Returns a + ``str``, ``dict``, :py:class:`ToolResult` or ``None``, or raises :py:class:`ToolError`. + :raises ValueError: The definition is invalid or its name is already registered. + """ + registered = _RegisteredTool(handler) + definition = core.BNMcpToolDefinition() + definition.name = name + definition.title = title + definition.description = description + definition.inputSchema = json.dumps(input_schema) + definition.outputSchema = json.dumps(output_schema) if output_schema is not None else None + definition.scope = scope + definition.annotations = annotations + handle = core.BNRegisterMcpTool(definition, registered.callbacks) + if not handle: + raise ValueError(f"Binary Ninja rejected the MCP tool '{name}'; see the log for why") + _registered_tools.append(registered) + return Tool(handle) + + +# Converts one argument, given the arguments converted before it. +_Converter = Callable[[ToolCall, Any, Dict[str, Any]], Any] + + +class _Parameter: + def __init__( + self, name: str, schema: dict, kind: str, convert: _Converter, required: bool, default: Any, + relative_to: Optional[str] + ): + self.name = name + self.schema = schema + self.kind = kind + self.convert = convert + self.required = required + self.default = default + self.relative_to = relative_to + + +def _is_integer(value: Any) -> bool: + return isinstance(value, int) and not isinstance(value, bool) + + +def _relative_to(annotation) -> Optional[str]: + """ + Returns the parameter named by an ``Annotated[IntegerExpression, RelativeTo(...)]`` annotation, or by + the item annotation of a ``List``. + """ + origin = typing.get_origin(annotation) + if origin in (list, List): + arguments = typing.get_args(annotation) + return _relative_to(arguments[0]) if arguments else None + if origin is not typing.Annotated: + return None + base, *metadata = typing.get_args(annotation) + names = [item.parameter for item in metadata if isinstance(item, RelativeTo)] + if not names: + return None + if base is not IntegerExpression: + raise TypeError("RelativeTo only applies to parameters annotated mcp.IntegerExpression") + return names[0] + + +def _describe(annotation, name: str, relative_to: Optional[str] = None) -> Tuple[dict, str, _Converter]: + """Returns the schema, the kind named in error messages, and the conversion for one annotation.""" + + def expect(kind: str, check: Callable[[Any], bool], convert: Callable[[Any], Any] = lambda value: value): + def run(call: ToolCall, value: Any, converted: Dict[str, Any]) -> Any: + if not check(value): + raise ToolError("invalid_params", f"Expected {kind} parameter '{name}'") + return convert(value) + + return run + + if annotation is Address: + + def address(call: ToolCall, value: Any, converted: Dict[str, Any]) -> int: + try: + return call.parse_address(value) + except ValueError as error: + raise ToolError("invalid_params", f"Invalid address expression parameter '{name}': {error}") + + return {"type": "string"}, "address expression", address + + if annotation is IntegerExpression: + + def integer(call: ToolCall, value: Any, converted: Dict[str, Any]) -> int: + here = converted.get(relative_to) if relative_to is not None else None + try: + # A negative int parameter used as $here wraps to 64 bits. + return call.parse_integer(value, (here or 0) & 0xFFFFFFFFFFFFFFFF) + except ValueError as error: + raise ToolError("invalid_params", f"Invalid unsigned integer parameter '{name}': {error}") + + return {"type": ["integer", "string"], "minimum": 0}, "unsigned integer", integer + + if annotation is str: + return {"type": "string"}, "string", expect("string", lambda value: isinstance(value, str)) + if annotation is bool: + return {"type": "boolean"}, "boolean", expect("boolean", lambda value: isinstance(value, bool)) + if annotation is int: + return {"type": "integer"}, "integer", expect("integer", _is_integer) + if annotation is float: + return ( + {"type": "number"}, + "number", + expect("number", lambda value: _is_integer(value) or isinstance(value, float), float), + ) + if inspect.isclass(annotation) and issubclass(annotation, enum.Enum): + names = [member.name for member in annotation] + + def member(call: ToolCall, value: Any, converted: Dict[str, Any]): + if not isinstance(value, str): + raise ToolError("invalid_params", f"Expected string parameter '{name}'") + if value not in names: + raise ToolError("invalid_params", f"Invalid enum value for parameter '{name}'") + return annotation[value] + + return {"type": "string", "enum": names}, "string", member + + origin = typing.get_origin(annotation) + arguments = typing.get_args(annotation) + if origin is Literal: + choices = list(arguments) + if not all(isinstance(choice, str) for choice in choices): + raise TypeError(f"Parameter '{name}': only string Literal choices are supported") + + def choice(call: ToolCall, value: Any, converted: Dict[str, Any]) -> str: + if not isinstance(value, str): + raise ToolError("invalid_params", f"Expected string parameter '{name}'") + if value not in choices: + raise ToolError("invalid_params", f"Invalid enum value for parameter '{name}'") + return value + + return {"type": "string", "enum": choices}, "string", choice + if origin in (list, List): + item_schema, item_kind, item_convert = _describe(arguments[0] if arguments else str, name, relative_to) + + def items(call: ToolCall, value: Any, converted: Dict[str, Any]) -> list: + if not isinstance(value, list): + raise ToolError("invalid_params", f"Expected {item_kind} array parameter '{name}'") + return [item_convert(call, item, converted) for item in value] + + return {"type": "array", "items": item_schema}, f"{item_kind} array", items + if origin is typing.Annotated: + base, *metadata = arguments + for item in metadata: + if isinstance(item, Schema): + return dict(item.schema), "JSON", lambda call, value, converted: value + limits = [item for item in metadata if isinstance(item, (Minimum, Maximum, ClampTo))] + if limits: + return _describe_limited_int(base, name, limits) + if any(item is NonEmpty or isinstance(item, NonEmpty) for item in metadata): + if base is not str: + raise TypeError(f"Parameter '{name}': NonEmpty only applies to str") + return ( + {"type": "string", "minLength": 1}, + "non-empty string", + expect("non-empty string", lambda value: isinstance(value, str) and value != ""), + ) + return _describe(base, name, relative_to) + raise TypeError(f"Parameter '{name}': unsupported annotation {annotation!r}") + + +def _describe_limited_int(base, name: str, limits: list) -> Tuple[dict, str, _Converter]: + if base is not int: + raise TypeError(f"Parameter '{name}': Minimum, Maximum and ClampTo only apply to int") + minimum: Optional[int] = None + maximum: Optional[int] = None + clamp = False + for limit in limits: + if isinstance(limit, Minimum): + minimum = limit.value + elif isinstance(limit, Maximum): + maximum = limit.value + else: + minimum, maximum, clamp = limit.minimum, limit.maximum, True + + schema: Dict[str, Any] = {"type": "integer"} + if minimum is not None: + schema["minimum"] = minimum + if maximum is not None: + schema["maximum"] = maximum + + def convert(call: ToolCall, value: Any, converted: Dict[str, Any]) -> int: + if not _is_integer(value): + raise ToolError("invalid_params", f"Expected integer parameter '{name}'") + if minimum is not None and value < minimum: + if clamp: + return minimum + raise ToolError("invalid_params", f"Invalid integer parameter '{name}': Must be at least {minimum}") + if maximum is not None and value > maximum: + if clamp: + return maximum + raise ToolError("invalid_params", f"Invalid integer parameter '{name}': Must be at most {maximum}") + return value + + return schema, "integer", convert + + +def _optional_inner(annotation): + """Returns T for Optional[T] or T | None, or None when the annotation is not optional.""" + origin = typing.get_origin(annotation) + if origin is Union or origin is _UnionType: + arguments = [argument for argument in typing.get_args(annotation) if argument is not type(None)] + if len(arguments) == 1 and len(typing.get_args(annotation)) == 2: + return arguments[0] + return None + + +def _parse_docstring(docstring: Optional[str]) -> Tuple[str, Dict[str, str]]: + """Returns the leading paragraph and the ``:param name:`` descriptions.""" + text = inspect.cleandoc(docstring or "") + # The description ends at the first blank line or ":field:" line. A field's continuation lines are + # indented and non-blank. + summary = re.split(r"\n\s*\n|\n(?=:)", text, maxsplit=1)[0] + description = " ".join(line.strip() for line in summary.splitlines()) + params: Dict[str, str] = {} + for match in re.finditer(r"^:param\s+(\w+):\s*(.*(?:\n[ \t]+\S.*)*)", text, re.MULTILINE): + params[match.group(1)] = " ".join(part.strip() for part in match.group(2).splitlines()) + return description, params + + +def _default_json(annotation, value: Any) -> Any: + """Returns a default value in the form the parameter's schema describes.""" + if annotation is Address: + return hex(value) + if isinstance(value, enum.Enum): + return value.name + origin = typing.get_origin(annotation) + arguments = typing.get_args(annotation) + if origin in (list, List): + return [_default_json(arguments[0] if arguments else str, item) for item in value] + if origin is typing.Annotated: + return _default_json(arguments[0], value) + return value + + +def _parameters(function: Callable, descriptions: Dict[str, str]) -> List[_Parameter]: + signature = inspect.signature(function) + hints = typing.get_type_hints(function, include_extras=True) + parameters = list(signature.parameters.values()) + if not parameters: + raise TypeError(f"MCP tool '{function.__name__}' must take the ToolCall as its first parameter") + + result: List[_Parameter] = [] + for parameter in parameters[1:]: + name = parameter.name + if parameter.kind in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD): + raise TypeError(f"MCP tool '{function.__name__}': *args and **kwargs are not supported") + if parameter.kind == inspect.Parameter.POSITIONAL_ONLY: + raise TypeError(f"MCP tool '{function.__name__}': parameter '{name}' cannot be positional-only") + if name not in hints: + raise TypeError(f"MCP tool '{function.__name__}': parameter '{name}' needs a type annotation") + if name not in descriptions: + raise TypeError(f"MCP tool '{function.__name__}': parameter '{name}' needs a ':param {name}:' description") + + annotation = hints[name] + inner = _optional_inner(annotation) + target = inner if inner is not None else annotation + relative_to = _relative_to(target) + if relative_to is not None and not any( + earlier.name == relative_to and earlier.kind in ("address expression", "unsigned integer", "integer") + for earlier in result + ): + raise TypeError( + f"MCP tool '{function.__name__}': parameter '{name}' is relative to '{relative_to}', which must be an " + "earlier integer parameter" + ) + schema, kind, convert = _describe(target, name, relative_to) + schema = dict(schema) + schema["description"] = descriptions[name] + if has_default := parameter.default is not inspect.Parameter.empty: + if parameter.default is not None: + schema["default"] = _default_json(target, parameter.default) + required = inner is None and not has_default + default = parameter.default if has_default else None + result.append(_Parameter(name, schema, kind, convert, required, default, relative_to)) + return result + + +def tool( + name: Optional[str] = None, + *, + title: str = "", + scope: McpToolScope = McpToolScope.BinaryViewScope, + read_only: bool = False, + destructive: bool = False, + idempotent: bool = False, + open_world: bool = False, + output_schema: Optional[dict] = None, +): + """ + Registers the decorated function as a tool. The function's first parameter receives the + :py:class:`ToolCall`. Every other parameter becomes a tool parameter, and needs a type annotation and + a ``:param name:`` line in the docstring. The docstring's leading paragraph becomes the tool's + description. + + Supported annotations are ``str``, ``int``, ``float``, ``bool``, :py:class:`Address`, + :py:class:`IntegerExpression`, Binary Ninja enums (by member name), ``Literal`` of strings, ``List[T]``, + ``Annotated[T, Schema({...})]``, ``Annotated[IntegerExpression, RelativeTo("name")]``, + ``Annotated[str, NonEmpty()]`` and ``int`` annotated with :py:class:`Minimum`, :py:class:`Maximum` or + :py:class:`ClampTo`. ``Optional[T]``, + ``T | None`` or a default value makes a parameter optional, and a null argument for one is treated as + absent. + + Use it with parentheses or without, as ``@mcp.tool()`` or ``@mcp.tool``. + + The tool's name defaults to the function's name. Arguments are checked before the function runs, and + a missing, invalid or undeclared argument produces an ``invalid_params`` error without calling it. + """ + if callable(name): + return tool()(name) + + def register(function: Callable) -> Callable: + tool_name = name or function.__name__ + if not _TOOL_NAME.match(tool_name): + raise ValueError(f"MCP tool name '{tool_name}' must be 1 to 64 letters, digits, '_' or '-'") + description, descriptions = _parse_docstring(function.__doc__) + if not description: + raise TypeError(f"MCP tool '{tool_name}' needs a docstring describing it") + parameters = _parameters(function, descriptions) + + input_schema: Dict[str, Any] = { + "type": "object", + "properties": {parameter.name: parameter.schema for parameter in parameters}, + } + required = [parameter.name for parameter in parameters if parameter.required] + if required: + input_schema["required"] = required + input_schema["additionalProperties"] = False + + declared = {parameter.name for parameter in parameters} + + def handler(call: ToolCall, arguments: Any) -> Any: + if not isinstance(arguments, dict): + raise ToolError("invalid_params", "Expected object arguments") + for argument in arguments: + if argument not in declared: + raise ToolError("invalid_params", f"Unexpected parameter '{argument}'") + values = {} + for parameter in parameters: + # Clients may send null for an optional argument they leave unset. + if parameter.name in arguments and (parameter.required or arguments[parameter.name] is not None): + values[parameter.name] = parameter.convert(call, arguments[parameter.name], values) + elif parameter.required: + raise ToolError("invalid_params", f"Expected {parameter.kind} parameter '{parameter.name}'") + else: + values[parameter.name] = parameter.default + return function(call, **values) + + annotations = 0 + if read_only: + annotations |= McpToolAnnotation.ReadOnlyHint + if destructive: + annotations |= McpToolAnnotation.DestructiveHint + if idempotent: + annotations |= McpToolAnnotation.IdempotentHint + if open_world: + annotations |= McpToolAnnotation.OpenWorldHint + + register_tool( + tool_name, + description, + input_schema, + handler, + title=title, + scope=scope, + annotations=annotations, + output_schema=output_schema, + ) + return function + + return register diff --git a/rust/Cargo.toml b/rust/Cargo.toml index 2889015d25..0ad1bdeae3 100644 --- a/rust/Cargo.toml +++ b/rust/Cargo.toml @@ -26,6 +26,8 @@ serde = "1.0" serde_derive = "1.0" # Parts of the collaboration and workflow APIs consume and produce JSON. serde_json = "1.0" +# Generates input schemas for MCP tools implementing mcp::TypedMcpTool. +schemars = { version = "1.2", optional = true } # Used for tracing compatible logs tracing = { version = "0.1", default-features = false, features = ["std"] } tracing-subscriber = { version = "0.3", default-features = false, features = ["std", "registry", "parking_lot"] } diff --git a/rust/binaryninjacore-sys/build.rs b/rust/binaryninjacore-sys/build.rs index 7947d0a971..41ff06c662 100644 --- a/rust/binaryninjacore-sys/build.rs +++ b/rust/binaryninjacore-sys/build.rs @@ -123,6 +123,7 @@ fn main() { .allowlist_var("BN_MINIMUM_CORE_ABI_VERSION") .allowlist_var("MAX_RELOCATION_SIZE") .allowlist_type("BNLinearSweepAnalysisCapability") + .allowlist_type("BNMcpToolAnnotation") .raw_line(format!( "pub const BN_CURRENT_UI_ABI_VERSION: u32 = {};", current_version @@ -135,6 +136,7 @@ fn main() { // Flag enums (BN_OPTIONS) must be newtypes, as combined bit values would be // undefined behavior for a fieldless Rust enum. .bitfield_enum("BNMetadataStoreFlag") + .bitfield_enum("BNMcpToolAnnotation") .generate() .expect("Unable to generate bindings") .write_to_file(PathBuf::from(out_dir).join("bindings.rs")) diff --git a/rust/src/lib.rs b/rust/src/lib.rs index 8e3f60b9ed..bf39e8a7e1 100644 --- a/rust/src/lib.rs +++ b/rust/src/lib.rs @@ -62,6 +62,7 @@ pub mod llvm; pub mod logger; pub mod low_level_il; pub mod main_thread; +pub mod mcp; pub mod medium_level_il; pub mod metadata; pub mod object_destructor; diff --git a/rust/src/mcp.rs b/rust/src/mcp.rs new file mode 100644 index 0000000000..5894d3a3ca --- /dev/null +++ b/rust/src/mcp.rs @@ -0,0 +1,20 @@ +// Copyright 2022-2026 Vector 35 Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//! Binary Ninja's Model Context Protocol (MCP) tools. +//! +//! [`tool`] is for writing and registering tools. [`server`] is for an MCP server that offers them. + +pub mod server; +pub mod tool; diff --git a/rust/src/mcp/server.rs b/rust/src/mcp/server.rs new file mode 100644 index 0000000000..5bdad48bc9 --- /dev/null +++ b/rust/src/mcp/server.rs @@ -0,0 +1,117 @@ +// Copyright 2022-2026 Vector 35 Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//! Offering MCP tools from an MCP server. + +use binaryninjacore_sys::*; +use serde_json::Value; +use std::borrow::Cow; +use std::ffi::{c_char, c_void, CStr}; +use std::panic::{self, AssertUnwindSafe}; +use std::ptr; + +use super::tool::{create_with, CoreMcpTool, McpTool}; +use crate::binary_view::BinaryView; +use crate::rc::*; +use crate::string::{BnString, IntoCStr}; + +/// Creates a tool without registering it. This can be used to expose a tool that is specific to one +/// MCP server, such as one that manages that server's active binary view. Otherwise the same as +/// [`register_mcp_tool`](super::tool::register_mcp_tool). +pub fn create_mcp_tool(tool: T) -> Option> { + create_with(tool, BNCreateMcpTool) +} + +/// Creates a [`TypedMcpTool`](super::tool::TypedMcpTool) without registering it. See +/// [`create_mcp_tool`]. +#[cfg(feature = "schemars")] +pub fn create_typed_mcp_tool(tool: T) -> Option> { + create_mcp_tool(super::tool::typed::TypedAdapter::new(tool)) +} + +/// The request state an MCP server supplies for one invocation of a tool. +pub trait McpToolCallHost: Sync { + fn binary_view(&self) -> Option>; + + fn is_cancelled(&self) -> bool { + false + } + + fn report_progress(&self, _progress: f64, _total: f64, _message: &str) {} +} + +/// Runs the tool for an MCP server and returns the MCP `CallToolResult`. +pub fn invoke_mcp_tool( + tool: &CoreMcpTool, + host: &H, + arguments: &Value, +) -> Value { + extern "C" fn cb_get_binary_view(ctxt: *mut c_void) -> *mut BNBinaryView { + let host = unsafe { &*(ctxt as *const H) }; + match panic::catch_unwind(AssertUnwindSafe(|| host.binary_view())) { + Ok(Some(view)) => unsafe { BNNewViewReference(view.handle) }, + Ok(None) => ptr::null_mut(), + Err(_) => { + tracing::error!("MCP host panicked supplying a binary view"); + ptr::null_mut() + } + } + } + + extern "C" fn cb_is_cancelled(ctxt: *mut c_void) -> bool { + let host = unsafe { &*(ctxt as *const H) }; + panic::catch_unwind(AssertUnwindSafe(|| host.is_cancelled())).unwrap_or_else(|_| { + tracing::error!("MCP host panicked reporting cancellation"); + false + }) + } + + extern "C" fn cb_report_progress( + ctxt: *mut c_void, + progress: f64, + total: f64, + message: *const c_char, + ) { + let host = unsafe { &*(ctxt as *const H) }; + let message = if message.is_null() { + Cow::Borrowed("") + } else { + unsafe { CStr::from_ptr(message) }.to_string_lossy() + }; + if panic::catch_unwind(AssertUnwindSafe(|| { + host.report_progress(progress, total, &message) + })) + .is_err() + { + tracing::error!("MCP host panicked receiving progress"); + } + } + + let callbacks = BNMcpToolCallCallbacks { + context: host as *const H as *mut c_void, + getBinaryView: Some(cb_get_binary_view::), + isCancelled: Some(cb_is_cancelled::), + reportProgress: Some(cb_report_progress::), + }; + let arguments = arguments.to_string().to_cstr(); + unsafe { + let call = BNCreateMcpToolCall(&callbacks); + let result = BNCreateMcpToolResult(); + BNInvokeMcpTool(tool.handle, call, arguments.as_ptr(), result); + let json = BnString::into_string(BNGetMcpToolResultJson(result)); + BNFreeMcpToolResult(result); + BNFreeMcpToolCall(call); + serde_json::from_str(&json).unwrap_or(Value::Null) + } +} diff --git a/rust/src/mcp/tool.rs b/rust/src/mcp/tool.rs new file mode 100644 index 0000000000..072549f8ca --- /dev/null +++ b/rust/src/mcp/tool.rs @@ -0,0 +1,864 @@ +// Copyright 2022-2026 Vector 35 Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//! Writing and registering MCP tools. +//! +//! A tool registered here is offered by every MCP server in the process, alongside the built-in +//! `bn_*` tools. Implement [`McpTool`] and pass it to [`register_mcp_tool`]. With the `schemars` +//! feature, implement `TypedMcpTool` instead and pass it to `register_typed_mcp_tool`. The input +//! schema and argument parsing then come from the arguments type. + +use binaryninjacore_sys::*; +use serde_json::{json, Value}; +use std::ffi::{c_char, c_void, CStr, CString}; +use std::fmt; +use std::panic::{self, AssertUnwindSafe}; +use std::ptr; + +use crate::binary_view::BinaryView; +use crate::rc::*; +use crate::string::{BnString, IntoCStr}; + +pub use binaryninjacore_sys::BNMcpToolScope as McpToolScope; + +bitflags::bitflags! { + /// Hints a client uses to decide how to present a tool and whether to ask before running it. + #[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)] + pub struct McpToolAnnotations: u32 { + const READ_ONLY = BNMcpToolAnnotation::ReadOnlyHint.0; + const DESTRUCTIVE = BNMcpToolAnnotation::DestructiveHint.0; + const IDEMPOTENT = BNMcpToolAnnotation::IdempotentHint.0; + const OPEN_WORLD = BNMcpToolAnnotation::OpenWorldHint.0; + } +} + +/// Everything about a tool except its input schema. +#[derive(Clone, Debug)] +pub struct McpToolInfo { + pub name: String, + pub title: String, + pub description: String, + pub scope: McpToolScope, + pub annotations: McpToolAnnotations, + pub output_schema: Option, +} + +impl McpToolInfo { + /// A BinaryView-scoped tool with no title, annotations or output schema. + pub fn new(name: impl Into, description: impl Into) -> Self { + Self { + name: name.into(), + title: String::new(), + description: description.into(), + scope: McpToolScope::BinaryViewScope, + annotations: McpToolAnnotations::empty(), + output_schema: None, + } + } + + pub fn with_title(mut self, title: impl Into) -> Self { + self.title = title.into(); + self + } + + pub fn with_scope(mut self, scope: McpToolScope) -> Self { + self.scope = scope; + self + } + + pub fn with_annotations(mut self, annotations: McpToolAnnotations) -> Self { + self.annotations = annotations; + self + } + + pub fn with_output_schema(mut self, schema: Value) -> Self { + self.output_schema = Some(schema); + self + } + + /// Completes the definition with a JSON Schema object for the tool's input. + pub fn with_input_schema(self, input_schema: Value) -> McpToolDefinition { + McpToolDefinition { + info: self, + input_schema, + } + } +} + +#[derive(Clone, Debug)] +pub struct McpToolDefinition { + pub info: McpToolInfo, + /// A JSON Schema object with `"type": "object"`. + pub input_schema: Value, +} + +/// An error result with a machine-readable code. +#[derive(Clone, Debug, PartialEq)] +pub struct McpToolError { + pub code: String, + pub message: String, + pub details: Option, +} + +impl McpToolError { + pub fn new(code: impl Into, message: impl Into) -> Self { + Self { + code: code.into(), + message: message.into(), + details: None, + } + } + + pub fn invalid_params(message: impl Into) -> Self { + Self::new("invalid_params", message) + } + + pub fn with_details(mut self, details: Value) -> Self { + self.details = Some(details); + self + } +} + +impl fmt::Display for McpToolError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}: {}", self.code, self.message) + } +} + +impl std::error::Error for McpToolError {} + +/// Arguments that do not match the type they are parsed into are invalid parameters. +impl From for McpToolError { + fn from(error: serde_json::Error) -> Self { + Self::invalid_params(error.to_string()) + } +} + +/// Tool output is arbitrary text, and a NUL would otherwise panic in `CString::new`. +fn c_text(text: &str) -> CString { + CString::new(text.replace('\0', "\u{FFFD}")).unwrap_or_default() +} + +#[derive(Clone, Debug, Default, PartialEq)] +pub struct McpToolResult { + text: Vec, + structured_content: Option, + error: Option, + warnings: Vec<(String, String)>, +} + +impl McpToolResult { + pub fn text(text: impl Into) -> Self { + Self { + text: vec![text.into()], + ..Self::default() + } + } + + /// `content` must be a JSON object. When no text is added, clients that ignore structured + /// content see it serialized as text. + pub fn structured(content: Value) -> Self { + Self { + structured_content: Some(content), + ..Self::default() + } + } + + pub fn error(error: McpToolError) -> Self { + Self { + error: Some(error), + ..Self::default() + } + } + + pub fn with_text(mut self, text: impl Into) -> Self { + self.text.push(text.into()); + self + } + + /// An advisory warning about the result, such as analysis that has not finished. Clients see + /// it in the structured content's reserved `warnings` member, and after any text. + pub fn with_warning(mut self, code: impl Into, message: impl Into) -> Self { + self.warnings.push((code.into(), message.into())); + self + } + + fn apply(&self, result: *mut BNMcpToolResult) { + unsafe { + if let Some(error) = &self.error { + let code = c_text(&error.code); + let message = c_text(&error.message); + let details = error + .details + .as_ref() + .map(|details| details.to_string().to_cstr()); + BNSetMcpToolResultError( + result, + code.as_ptr(), + message.as_ptr(), + details + .as_ref() + .map_or(ptr::null(), |details| details.as_ptr()), + ); + } else if let Some(content) = &self.structured_content { + let json = content.to_string().to_cstr(); + if !BNSetMcpToolResultStructuredContent(result, json.as_ptr()) { + let code = "internal_error".to_cstr(); + let message = + "The tool produced structured content that is not a JSON object".to_cstr(); + BNSetMcpToolResultError(result, code.as_ptr(), message.as_ptr(), ptr::null()); + return; + } + } + for text in &self.text { + let text = c_text(text); + BNAddMcpToolResultText(result, text.as_ptr()); + } + for (code, message) in &self.warnings { + let code = c_text(code); + let message = c_text(message); + BNAddMcpToolResultWarning(result, code.as_ptr(), message.as_ptr()); + } + } + } +} + +impl From for McpToolResult { + fn from(error: McpToolError) -> Self { + Self::error(error) + } +} + +/// One invocation of a tool. The binary view and cancellation state are only available while the +/// tool is running. +pub struct McpToolCall { + handle: *mut BNMcpToolCall, +} + +impl McpToolCall { + pub(crate) unsafe fn ref_from_raw(handle: *mut BNMcpToolCall) -> Ref { + debug_assert!(!handle.is_null()); + Ref::new(Self { handle }) + } + + /// The binary view the MCP session targets. Never `None` while a BinaryView-scoped tool runs. + pub fn binary_view(&self) -> Option> { + let view = unsafe { BNGetMcpToolCallBinaryView(self.handle) }; + (!view.is_null()).then(|| unsafe { BinaryView::ref_from_raw(view) }) + } + + /// The binary view, or the `no_active_binary_view` error to return with `?` when there is none. + pub fn require_binary_view(&self) -> Result, McpToolError> { + self.binary_view().ok_or_else(|| { + McpToolError::new( + "no_active_binary_view", + "No active Binary Ninja binary view is selected", + ) + }) + } + + pub fn is_cancelled(&self) -> bool { + unsafe { BNIsMcpToolCallCancelled(self.handle) } + } + + pub fn report_progress(&self, progress: f64, total: f64, message: &str) { + let message = c_text(message); + unsafe { BNReportMcpToolCallProgress(self.handle, progress, total, message.as_ptr()) } + } + + /// Evaluates an address expression string against the call's binary view, with `here` as the + /// value of `$here`. + pub fn parse_address(&self, value: &Value, here: u64) -> Result { + self.parse(BNParseMcpToolCallAddress, value, here) + } + + /// Evaluates an unsigned integer, or an expression string against the call's binary view with + /// `here` as the value of `$here`. + pub fn parse_integer(&self, value: &Value, here: u64) -> Result { + self.parse(BNParseMcpToolCallInteger, value, here) + } + + fn parse( + &self, + parse: unsafe extern "C" fn( + *mut BNMcpToolCall, + *const c_char, + *mut u64, + u64, + *mut *mut c_char, + ) -> bool, + value: &Value, + here: u64, + ) -> Result { + let json = value.to_string().to_cstr(); + let mut result = 0; + let mut error: *mut c_char = ptr::null_mut(); + if unsafe { parse(self.handle, json.as_ptr(), &mut result, here, &mut error) } { + Ok(result) + } else if error.is_null() { + Err(String::new()) + } else { + Err(unsafe { BnString::into_string(error) }) + } + } +} + +unsafe impl Send for McpToolCall {} +unsafe impl Sync for McpToolCall {} + +impl ToOwned for McpToolCall { + type Owned = Ref; + + fn to_owned(&self) -> Self::Owned { + unsafe { RefCountable::inc_ref(self) } + } +} + +unsafe impl RefCountable for McpToolCall { + unsafe fn inc_ref(handle: &Self) -> Ref { + Ref::new(Self { + handle: BNNewMcpToolCallReference(handle.handle), + }) + } + + unsafe fn dec_ref(handle: &Self) { + BNFreeMcpToolCall(handle.handle); + } +} + +/// A tool with a hand-written input schema that receives its arguments as JSON. +pub trait McpTool: 'static + Send + Sync { + fn definition(&self) -> McpToolDefinition; + + /// Called on an arbitrary thread. An error produces an error result. Under `panic = "unwind"` + /// a panic produces an `internal_error` result, and under `panic = "abort"` it ends the process. + fn invoke(&self, call: &McpToolCall, arguments: Value) -> Result; +} + +/// Registers a tool for the life of the process. Returns `None` when the definition is invalid or +/// its name is already registered, or when a string in it contains a NUL. +pub fn register_mcp_tool(tool: T) -> Option> { + create_with(tool, BNRegisterMcpTool) +} + +pub(super) fn create_with( + tool: T, + create: unsafe extern "C" fn( + *const BNMcpToolDefinition, + *const BNMcpToolCallbacks, + ) -> *mut BNMcpTool, +) -> Option> { + struct ToolContext { + name: String, + tool: T, + } + + extern "C" fn cb_invoke( + ctxt: *mut c_void, + call: *mut BNMcpToolCall, + arguments: *const c_char, + result: *mut BNMcpToolResult, + ) { + let context = unsafe { &*(ctxt as *const ToolContext) }; + let outcome = panic::catch_unwind(AssertUnwindSafe(|| { + let call = unsafe { McpToolCall::ref_from_raw(BNNewMcpToolCallReference(call)) }; + let arguments = unsafe { CStr::from_ptr(arguments) }.to_string_lossy(); + serde_json::from_str::(&arguments) + .map_err(|error| { + McpToolError::invalid_params(format!("Arguments are not valid JSON: {error}")) + }) + .and_then(|arguments| context.tool.invoke(&call, arguments)) + .unwrap_or_else(McpToolResult::error) + })); + let tool_result = outcome.unwrap_or_else(|payload| { + let message = payload + .downcast_ref::<&str>() + .copied() + .or_else(|| payload.downcast_ref::().map(String::as_str)) + .unwrap_or("unknown panic"); + tracing::error!("MCP tool '{}' panicked: {message}", context.name); + McpToolError::new("internal_error", format!("The tool panicked: {message}")).into() + }); + tool_result.apply(result); + } + + let definition = tool.definition(); + let name = CString::new(definition.info.name.as_str()).ok()?; + let title = CString::new(definition.info.title.as_str()).ok()?; + let description = CString::new(definition.info.description.as_str()).ok()?; + let input_schema = definition.input_schema.to_string().to_cstr(); + let output_schema = definition + .info + .output_schema + .as_ref() + .map(|schema| schema.to_string().to_cstr()); + let api_definition = BNMcpToolDefinition { + name: name.as_ptr(), + title: title.as_ptr(), + description: description.as_ptr(), + inputSchema: input_schema.as_ptr(), + outputSchema: output_schema + .as_ref() + .map_or(ptr::null(), |schema| schema.as_ptr()), + scope: definition.info.scope, + annotations: definition.info.annotations.bits(), + }; + + extern "C" fn cb_free_object(ctxt: *mut c_void) { + unsafe { drop(Box::from_raw(ctxt as *mut ToolContext)) }; + } + + let ctxt = Box::into_raw(Box::new(ToolContext { + name: definition.info.name.clone(), + tool, + })); + let callbacks = BNMcpToolCallbacks { + context: ctxt as *mut c_void, + invoke: Some(cb_invoke::), + freeObject: Some(cb_free_object::), + }; + let handle = unsafe { create(&api_definition, &callbacks) }; + if handle.is_null() { + unsafe { drop(Box::from_raw(ctxt)) }; + return None; + } + Some(unsafe { CoreMcpTool::ref_from_raw(handle) }) +} + +/// A tool in the tool registry. +#[derive(PartialEq, Eq, Hash)] +pub struct CoreMcpTool { + pub(super) handle: *mut BNMcpTool, +} + +impl CoreMcpTool { + pub(crate) unsafe fn from_raw(handle: *mut BNMcpTool) -> Self { + debug_assert!(!handle.is_null()); + Self { handle } + } + + pub(crate) unsafe fn ref_from_raw(handle: *mut BNMcpTool) -> Ref { + Ref::new(Self::from_raw(handle)) + } + + /// Every registered tool, sorted by name. + pub fn list() -> Array { + let mut count = 0; + let tools = unsafe { BNGetMcpToolList(&mut count) }; + unsafe { Array::new(tools, count, ()) } + } + + pub fn from_name(name: &str) -> Option> { + let name = name.to_cstr(); + let handle = unsafe { BNGetMcpToolByName(name.as_ptr()) }; + (!handle.is_null()).then(|| unsafe { Self::ref_from_raw(handle) }) + } + + pub fn name(&self) -> String { + unsafe { BnString::into_string(BNGetMcpToolName(self.handle)) } + } + + pub fn title(&self) -> String { + unsafe { BnString::into_string(BNGetMcpToolTitle(self.handle)) } + } + + pub fn description(&self) -> String { + unsafe { BnString::into_string(BNGetMcpToolDescription(self.handle)) } + } + + pub fn input_schema(&self) -> Value { + let schema = unsafe { BnString::into_string(BNGetMcpToolInputSchema(self.handle)) }; + serde_json::from_str(&schema).unwrap_or(Value::Null) + } + + pub fn output_schema(&self) -> Option { + let schema = unsafe { BnString::into_string(BNGetMcpToolOutputSchema(self.handle)) }; + serde_json::from_str(&schema).ok() + } + + pub fn scope(&self) -> McpToolScope { + unsafe { BNGetMcpToolScope(self.handle) } + } + + pub fn annotations(&self) -> McpToolAnnotations { + McpToolAnnotations::from_bits_retain(unsafe { BNGetMcpToolAnnotations(self.handle) }) + } +} + +impl fmt::Debug for CoreMcpTool { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("CoreMcpTool") + .field("name", &self.name()) + .finish() + } +} + +unsafe impl Send for CoreMcpTool {} +unsafe impl Sync for CoreMcpTool {} + +impl ToOwned for CoreMcpTool { + type Owned = Ref; + + fn to_owned(&self) -> Self::Owned { + unsafe { RefCountable::inc_ref(self) } + } +} + +unsafe impl RefCountable for CoreMcpTool { + unsafe fn inc_ref(handle: &Self) -> Ref { + Self::ref_from_raw(BNNewMcpToolReference(handle.handle)) + } + + unsafe fn dec_ref(handle: &Self) { + BNFreeMcpTool(handle.handle); + } +} + +impl CoreArrayProvider for CoreMcpTool { + type Raw = *mut BNMcpTool; + type Context = (); + type Wrapped<'a> = Guard<'a, CoreMcpTool>; +} + +unsafe impl CoreArrayProviderInner for CoreMcpTool { + unsafe fn free(raw: *mut Self::Raw, count: usize, _context: &Self::Context) { + BNFreeMcpToolList(raw, count); + } + + unsafe fn wrap_raw<'a>(raw: &'a Self::Raw, context: &'a Self::Context) -> Self::Wrapped<'a> { + Guard::new(CoreMcpTool::from_raw(*raw), context) + } +} + +/// An address expression string, evaluated against the call's binary view with +/// [`McpAddress::resolve`]. +#[derive(Clone, Debug, PartialEq, Eq, serde_derive::Deserialize)] +#[serde(transparent)] +pub struct McpAddress(pub String); + +impl McpAddress { + /// An invalid expression is an `invalid_params` error. + pub fn resolve(&self, call: &McpToolCall) -> Result { + call.parse_address(&Value::String(self.0.clone()), 0) + .map_err(McpToolError::invalid_params) + } +} + +/// A string argument that rejects an empty string. +#[derive(Clone, Debug, PartialEq, Eq, Hash)] +pub struct McpNonEmptyString(pub String); + +impl<'de> serde::Deserialize<'de> for McpNonEmptyString { + fn deserialize>(deserializer: D) -> Result { + let value = ::deserialize(deserializer)?; + if value.is_empty() { + return Err(serde::de::Error::invalid_value( + serde::de::Unexpected::Str(&value), + &"a non-empty string", + )); + } + Ok(McpNonEmptyString(value)) + } +} + +/// An unsigned integer, or an expression string evaluated against the call's binary view with +/// [`McpIntegerExpression::resolve`]. +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum McpIntegerExpression { + Integer(u64), + Expression(String), +} + +impl<'de> serde::Deserialize<'de> for McpIntegerExpression { + fn deserialize>(deserializer: D) -> Result { + struct Visitor; + + impl serde::de::Visitor<'_> for Visitor { + type Value = McpIntegerExpression; + + fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result { + formatter.write_str("an unsigned integer or an integer expression string") + } + + fn visit_u64(self, value: u64) -> Result { + Ok(McpIntegerExpression::Integer(value)) + } + + fn visit_i64(self, value: i64) -> Result { + u64::try_from(value) + .map(McpIntegerExpression::Integer) + .map_err(|_| E::invalid_value(serde::de::Unexpected::Signed(value), &self)) + } + + fn visit_str(self, value: &str) -> Result { + Ok(McpIntegerExpression::Expression(value.to_owned())) + } + } + + deserializer.deserialize_any(Visitor) + } +} + +impl McpIntegerExpression { + /// An invalid expression is an `invalid_params` error. + pub fn resolve(&self, call: &McpToolCall) -> Result { + self.resolve_relative_to(call, 0) + } + + /// Resolves an expression with `here` as the value of `$here`, such as a length measured from + /// an address argument. + pub fn resolve_relative_to(&self, call: &McpToolCall, here: u64) -> Result { + match self { + McpIntegerExpression::Integer(value) => Ok(*value), + McpIntegerExpression::Expression(expression) => call + .parse_integer(&Value::String(expression.clone()), here) + .map_err(McpToolError::invalid_params), + } + } +} + +#[cfg(feature = "schemars")] +pub(super) mod typed { + use super::*; + use schemars::{json_schema, JsonSchema, Schema, SchemaGenerator}; + use serde::Deserialize; + use std::borrow::Cow; + use std::sync::OnceLock; + + impl JsonSchema for McpAddress { + fn schema_name() -> Cow<'static, str> { + "McpAddress".into() + } + + fn inline_schema() -> bool { + true + } + + fn json_schema(_: &mut SchemaGenerator) -> Schema { + json_schema!({ "type": "string" }) + } + } + + impl JsonSchema for McpNonEmptyString { + fn schema_name() -> Cow<'static, str> { + "McpNonEmptyString".into() + } + + fn inline_schema() -> bool { + true + } + + fn json_schema(_: &mut SchemaGenerator) -> Schema { + json_schema!({ "type": "string", "minLength": 1 }) + } + } + + impl JsonSchema for McpIntegerExpression { + fn schema_name() -> Cow<'static, str> { + "McpIntegerExpression".into() + } + + fn inline_schema() -> bool { + true + } + + fn json_schema(_: &mut SchemaGenerator) -> Schema { + json_schema!({ "type": ["integer", "string"], "minimum": 0 }) + } + } + + /// A tool whose input schema and argument parsing come from its arguments type. Doc comments on + /// the arguments type's fields become parameter descriptions, and `Option` fields and + /// `#[serde(default)]` make parameters optional, with a null argument treated as absent. + /// Undeclared arguments produce an `invalid_params` result, unless the arguments type accepts + /// additional properties. + /// + /// The input schema must be an object, so `Args` is a struct with named fields, or with none for + /// a tool that takes no arguments. A unit struct or `()` describes `null` and is rejected. + pub trait TypedMcpTool: 'static + Send + Sync { + type Args: serde::de::DeserializeOwned + JsonSchema; + + fn info(&self) -> McpToolInfo; + + /// Called on an arbitrary thread with arguments that already match `Args`. A panic is + /// handled as it is for [`McpTool::invoke`]. + fn invoke( + &self, + call: &McpToolCall, + args: Self::Args, + ) -> Result; + } + + struct Parameters { + // The only argument names `invoke` accepts, or `None` when it accepts any name. + accepted: Option>, + // The parameters whose null arguments are treated as absent. + optional: Vec, + } + + pub(in crate::mcp) struct TypedAdapter { + tool: T, + parameters: OnceLock, + } + + impl TypedAdapter { + pub(in crate::mcp) fn new(tool: T) -> Self { + Self { + tool, + parameters: OnceLock::new(), + } + } + + fn input_schema() -> Value { + let settings = schemars::generate::SchemaSettings::draft2020_12().with(|settings| { + settings.inline_subschemas = true; + }); + let mut schema = settings + .into_generator() + .into_root_schema_for::() + .to_value(); + if let Some(object) = schema.as_object_mut() { + object.remove("$schema"); + object.remove("title"); + if object.get("type") == Some(&Value::String("object".into())) + && !object.contains_key("additionalProperties") + { + object.insert("additionalProperties".into(), Value::Bool(false)); + } + } + schema + } + + fn parameters(schema: &Value) -> Parameters { + let no_properties = serde_json::Map::new(); + let properties = schema + .get("properties") + .and_then(Value::as_object) + .unwrap_or(&no_properties); + let required = schema.get("required").and_then(Value::as_array); + let is_required = |name: &String| { + required.is_some_and(|required| required.iter().any(|entry| entry == name)) + }; + let accepted = (schema.get("additionalProperties") == Some(&Value::Bool(false))) + .then(|| properties.keys().cloned().collect()); + let optional = properties + .keys() + .filter(|name| !is_required(name)) + .cloned() + .collect(); + Parameters { accepted, optional } + } + } + + /// Deserializes the arguments object as `Args`, naming the parameter in each error. + struct Arguments<'de>(&'de serde_json::Map); + + impl<'de> serde::Deserializer<'de> for Arguments<'de> { + type Error = serde_json::Error; + + fn deserialize_any>( + self, + visitor: V, + ) -> Result { + visitor.visit_map(ArgumentsMap { + entries: self.0.iter(), + value: None, + }) + } + + serde::forward_to_deserialize_any! { + bool i8 i16 i32 i64 i128 u8 u16 u32 u64 u128 f32 f64 char str string bytes byte_buf + option unit unit_struct newtype_struct seq tuple tuple_struct map struct enum identifier + ignored_any + } + } + + struct ArgumentsMap<'de> { + entries: serde_json::map::Iter<'de>, + value: Option<(&'de String, &'de Value)>, + } + + impl<'de> serde::de::MapAccess<'de> for ArgumentsMap<'de> { + type Error = serde_json::Error; + + fn next_key_seed>( + &mut self, + seed: K, + ) -> Result, Self::Error> { + use serde::de::IntoDeserializer; + + let Some((name, value)) = self.entries.next() else { + return Ok(None); + }; + self.value = Some((name, value)); + seed.deserialize(name.as_str().into_deserializer()) + .map(Some) + } + + fn next_value_seed>( + &mut self, + seed: S, + ) -> Result { + let (name, value) = self + .value + .take() + .expect("next_value_seed called before next_key_seed"); + seed.deserialize(value).map_err(|error| { + serde::de::Error::custom(format!("Invalid parameter '{name}': {error}")) + }) + } + } + + impl McpTool for TypedAdapter { + fn definition(&self) -> McpToolDefinition { + self.tool.info().with_input_schema(Self::input_schema()) + } + + fn invoke( + &self, + call: &McpToolCall, + mut arguments: Value, + ) -> Result { + let parameters = self + .parameters + .get_or_init(|| Self::parameters(&Self::input_schema())); + let Some(arguments) = arguments.as_object_mut() else { + return Err(McpToolError::invalid_params("Expected object arguments")); + }; + if let Some(accepted) = ¶meters.accepted { + if let Some(name) = arguments.keys().find(|name| !accepted.contains(name)) { + return Err(McpToolError::invalid_params(format!( + "Unexpected parameter '{name}'" + ))); + } + } + arguments.retain(|name, value| !value.is_null() || !parameters.optional.contains(name)); + let args = T::Args::deserialize(Arguments(arguments))?; + self.tool.invoke(call, args) + } + } + + /// Registers a [`TypedMcpTool`] for the life of the process. See [`register_mcp_tool`]. + pub fn register_typed_mcp_tool(tool: T) -> Option> { + register_mcp_tool(TypedAdapter::new(tool)) + } +} + +#[cfg(feature = "schemars")] +pub use typed::{register_typed_mcp_tool, TypedMcpTool}; + +/// A JSON Schema object for a tool that takes no arguments. +pub fn empty_input_schema() -> Value { + json!({ "type": "object", "additionalProperties": false }) +} diff --git a/rust/tests/mcp.rs b/rust/tests/mcp.rs new file mode 100644 index 0000000000..72549ef66b --- /dev/null +++ b/rust/tests/mcp.rs @@ -0,0 +1,657 @@ +use binaryninja::binary_view::BinaryView; +use binaryninja::file_metadata::FileMetadata; +use binaryninja::headless::Session; +use binaryninja::mcp::server::{create_mcp_tool, invoke_mcp_tool, McpToolCallHost}; +use binaryninja::mcp::tool::{ + empty_input_schema, register_mcp_tool, CoreMcpTool, McpIntegerExpression, McpTool, + McpToolAnnotations, McpToolCall, McpToolDefinition, McpToolError, McpToolInfo, McpToolResult, + McpToolScope, +}; +use binaryninja::rc::Ref; +use serde_json::{json, Value}; +use std::sync::Mutex; + +// The registry is process-global and has no unregistration, so every test registers its tools under +// its own names. + +/// Reports whether it was given a binary view. +struct ViewTool { + name: &'static str, + scope: McpToolScope, +} + +impl McpTool for ViewTool { + fn definition(&self) -> McpToolDefinition { + McpToolInfo::new(self.name, "Report whether the call has a binary view.") + .with_scope(self.scope) + .with_annotations(McpToolAnnotations::READ_ONLY | McpToolAnnotations::IDEMPOTENT) + .with_input_schema(empty_input_schema()) + } + + fn invoke(&self, call: &McpToolCall, _arguments: Value) -> Result { + Ok(McpToolResult::structured( + json!({ "hasView": call.binary_view().is_some() }), + )) + } +} + +/// Evaluates its `size` argument, which may be an integer or an expression. +struct SizeTool; + +impl McpTool for SizeTool { + fn definition(&self) -> McpToolDefinition { + McpToolInfo::new("rust_mcp_size", "Evaluate a size.").with_input_schema(json!({ + "type": "object", + "properties": { "size": { "type": ["integer", "string"] } }, + "required": ["size"], + })) + } + + fn invoke(&self, call: &McpToolCall, arguments: Value) -> Result { + let size: McpIntegerExpression = serde_json::from_value(arguments["size"].clone())?; + Ok(McpToolResult::structured( + json!({ "size": size.resolve(call)? }), + )) + } +} + +struct FailingTool; + +impl McpTool for FailingTool { + fn definition(&self) -> McpToolDefinition { + McpToolInfo::new("rust_mcp_failing", "Always fail.") + .with_scope(McpToolScope::GlobalScope) + .with_input_schema(empty_input_schema()) + } + + fn invoke( + &self, + _call: &McpToolCall, + _arguments: Value, + ) -> Result { + Err(McpToolError::new("requested_failure", "Failed on request") + .with_details(json!({ "why": "asked" }))) + } +} + +struct PanickingTool; + +impl McpTool for PanickingTool { + fn definition(&self) -> McpToolDefinition { + McpToolInfo::new("rust_mcp_panicking", "Always panic.") + .with_scope(McpToolScope::GlobalScope) + .with_input_schema(empty_input_schema()) + } + + fn invoke( + &self, + _call: &McpToolCall, + _arguments: Value, + ) -> Result { + panic!("requested panic") + } +} + +/// Reports progress, then whether the call was cancelled. +struct ProgressTool; + +impl McpTool for ProgressTool { + fn definition(&self) -> McpToolDefinition { + McpToolInfo::new("rust_mcp_progress", "Report progress and cancellation.") + .with_scope(McpToolScope::GlobalScope) + .with_input_schema(empty_input_schema()) + } + + fn invoke(&self, call: &McpToolCall, _arguments: Value) -> Result { + call.report_progress(1.0, 2.0, "halfway"); + Ok(McpToolResult::structured( + json!({ "cancelled": call.is_cancelled() }), + )) + } +} + +struct ViewHost(Option>); + +impl McpToolCallHost for ViewHost { + fn binary_view(&self) -> Option> { + self.0.clone() + } +} + +struct RecordingHost { + cancelled: bool, + progress: Mutex>, +} + +impl McpToolCallHost for RecordingHost { + fn binary_view(&self) -> Option> { + None + } + + fn is_cancelled(&self) -> bool { + self.cancelled + } + + fn report_progress(&self, progress: f64, total: f64, message: &str) { + self.progress + .lock() + .unwrap() + .push((progress, total, message.to_string())); + } +} + +fn raw_view() -> Ref { + BinaryView::from_data(&FileMetadata::new(), &[0u8; 0x100]) +} + +fn invoke(tool: &CoreMcpTool, arguments: &Value, view: Option<&BinaryView>) -> Value { + invoke_mcp_tool(tool, &ViewHost(view.map(|view| view.to_owned())), arguments) +} + +fn error_code(result: &Value) -> &str { + assert_eq!(result["isError"], true, "{result}"); + result["structuredContent"]["errorCode"] + .as_str() + .unwrap_or_default() +} + +#[test] +fn registers_lists_and_describes_tools() { + let _session = Session::new().expect("Failed to initialize session"); + let tool = register_mcp_tool(ViewTool { + name: "rust_mcp_describe", + scope: McpToolScope::GlobalScope, + }) + .expect("registration failed"); + + assert_eq!(tool.name(), "rust_mcp_describe"); + assert_eq!( + tool.description(), + "Report whether the call has a binary view." + ); + assert_eq!(tool.scope(), McpToolScope::GlobalScope); + assert_eq!( + tool.annotations(), + McpToolAnnotations::READ_ONLY | McpToolAnnotations::IDEMPOTENT + ); + assert_eq!(tool.input_schema(), empty_input_schema()); + assert_eq!(tool.output_schema(), None); + assert_eq!(CoreMcpTool::from_name("rust_mcp_describe"), Some(tool)); + + let names: Vec = CoreMcpTool::list().iter().map(|tool| tool.name()).collect(); + assert!(names.iter().any(|name| name == "rust_mcp_describe")); + assert!(names.windows(2).all(|pair| pair[0] < pair[1])); +} + +#[test] +fn hosts_supply_cancellation_and_receive_progress() { + let _session = Session::new().expect("Failed to initialize session"); + let tool = create_mcp_tool(ProgressTool).expect("creation failed"); + + for cancelled in [false, true] { + let host = RecordingHost { + cancelled, + progress: Mutex::new(Vec::new()), + }; + let result = invoke_mcp_tool(&tool, &host, &json!({})); + assert_eq!(result["structuredContent"]["cancelled"], cancelled); + assert_eq!( + *host.progress.lock().unwrap(), + vec![(1.0, 2.0, "halfway".to_string())] + ); + } +} + +#[test] +fn created_tools_are_not_registered() { + let _session = Session::new().expect("Failed to initialize session"); + let tool = create_mcp_tool(ViewTool { + name: "rust_mcp_created", + scope: McpToolScope::GlobalScope, + }) + .expect("creation failed"); + + assert_eq!(tool.name(), "rust_mcp_created"); + assert_eq!(CoreMcpTool::from_name("rust_mcp_created"), None); + assert!(CoreMcpTool::list() + .iter() + .all(|tool| tool.name() != "rust_mcp_created")); + assert_eq!( + invoke(&tool, &json!({}), None)["structuredContent"]["hasView"], + false + ); +} + +#[test] +fn rejects_duplicate_and_invalid_names() { + let _session = Session::new().expect("Failed to initialize session"); + let scope = McpToolScope::GlobalScope; + assert!(register_mcp_tool(ViewTool { + name: "rust_mcp_duplicate", + scope + }) + .is_some()); + assert!(register_mcp_tool(ViewTool { + name: "rust_mcp_duplicate", + scope + }) + .is_none()); + assert!(register_mcp_tool(ViewTool { + name: "rust mcp", + scope + }) + .is_none()); +} + +#[test] +fn errors_and_panics_become_error_results() { + let _session = Session::new().expect("Failed to initialize session"); + let failing = register_mcp_tool(FailingTool).expect("registration failed"); + let panicking = register_mcp_tool(PanickingTool).expect("registration failed"); + + let result = invoke(&failing, &json!({}), None); + assert_eq!(error_code(&result), "requested_failure"); + assert_eq!(result["structuredContent"]["details"]["why"], "asked"); + + let result = invoke(&panicking, &json!({}), None); + assert_eq!(error_code(&result), "internal_error"); + assert_eq!( + result["structuredContent"]["errorMessage"], + "The tool panicked: requested panic" + ); +} + +#[test] +fn binary_view_scope_requires_a_view() { + let _session = Session::new().expect("Failed to initialize session"); + let tool = register_mcp_tool(ViewTool { + name: "rust_mcp_view", + scope: McpToolScope::BinaryViewScope, + }) + .expect("registration failed"); + + assert_eq!( + error_code(&invoke(&tool, &json!({}), None)), + "no_active_binary_view" + ); + + let view = raw_view(); + let result = invoke(&tool, &json!({}), Some(&view)); + assert_eq!(result["structuredContent"]["hasView"], true); +} + +#[test] +fn integer_arguments_accept_integers_and_expressions() { + let _session = Session::new().expect("Failed to initialize session"); + let tool = register_mcp_tool(SizeTool).expect("registration failed"); + let view = raw_view(); + + assert_eq!( + invoke(&tool, &json!({ "size": 4 }), Some(&view))["structuredContent"]["size"], + 4 + ); + assert_eq!( + invoke(&tool, &json!({ "size": "0x10 + 2" }), Some(&view))["structuredContent"]["size"], + 0x12 + ); + assert_eq!( + error_code(&invoke(&tool, &json!({ "size": true }), Some(&view))), + "invalid_params" + ); +} + +#[cfg(feature = "schemars")] +mod typed { + use super::*; + use binaryninja::mcp::server::create_typed_mcp_tool; + use binaryninja::mcp::tool::{ + register_typed_mcp_tool, McpAddress, McpNonEmptyString, TypedMcpTool, + }; + use schemars::JsonSchema; + use serde_derive::Deserialize; + + #[derive(Deserialize, JsonSchema)] + #[serde(rename_all = "lowercase")] + enum CommentKind { + Regular, + Repeatable, + } + + #[derive(Deserialize, JsonSchema)] + #[serde(deny_unknown_fields)] + struct CommentArgs { + /// Address expression of the comment. + address: McpAddress, + /// Comment text. + text: String, + /// Kind of comment. + kind: CommentKind, + /// Bytes the comment covers. + length: Option, + /// Optional repeat count. + #[serde(default)] + count: Option, + } + + struct CommentTool; + + impl TypedMcpTool for CommentTool { + type Args = CommentArgs; + + fn info(&self) -> McpToolInfo { + McpToolInfo::new("rust_mcp_typed", "Resolve a comment address.") + } + + fn invoke( + &self, + call: &McpToolCall, + args: CommentArgs, + ) -> Result { + let address = args.address.resolve(call)?; + let length = args.length.map(|length| length.resolve(call)).transpose()?; + let repeatable = matches!(args.kind, CommentKind::Repeatable); + Ok(McpToolResult::structured(json!({ + "address": address, + "text": args.text, + "repeatable": repeatable, + "length": length, + "count": args.count, + }))) + } + } + + #[test] + fn typed_tools_generate_schemas_and_parse_arguments() { + let _session = Session::new().expect("Failed to initialize session"); + let tool = register_typed_mcp_tool(CommentTool).expect("registration failed"); + let created = create_typed_mcp_tool(CommentTool).expect("creation failed"); + assert_eq!(created.input_schema(), tool.input_schema()); + + assert_eq!( + tool.input_schema(), + json!({ + "type": "object", + "properties": { + "address": { "type": "string", "description": "Address expression of the comment." }, + "text": { "type": "string", "description": "Comment text." }, + "kind": { "type": "string", "enum": ["regular", "repeatable"], "description": "Kind of comment." }, + "length": { + "type": ["integer", "string", "null"], + "minimum": 0, + "description": "Bytes the comment covers.", + }, + "count": { + "type": ["integer", "null"], + "format": "uint32", + "minimum": 0, + "default": null, + "description": "Optional repeat count.", + }, + }, + "required": ["address", "text", "kind"], + "additionalProperties": false, + }) + ); + + let view = raw_view(); + let result = invoke( + &tool, + &json!({ "address": "0x20", "text": "hi", "kind": "repeatable", "length": "0x10 + 2", "count": 3 }), + Some(&view), + ); + assert_eq!( + result["structuredContent"], + json!({ "address": 0x20, "text": "hi", "repeatable": true, "length": 0x12, "count": 3 }) + ); + + let result = invoke( + &tool, + &json!({ "address": "0x20", "text": "hi", "kind": "regular" }), + Some(&view), + ); + assert_eq!(result["structuredContent"]["length"], Value::Null); + assert_eq!(result["structuredContent"]["count"], Value::Null); + + for arguments in [ + json!({ "address": "0x20", "text": "hi", "kind": "regular", "other": 1 }), + json!({ "address": "0x20", "text": "hi" }), + json!({ "address": "0x20", "text": "hi", "kind": "block" }), + json!({ "address": "0x20", "text": 5, "kind": "regular" }), + ] { + let result = invoke(&tool, &arguments, Some(&view)); + assert_eq!(error_code(&result), "invalid_params", "{arguments}"); + } + + let result = invoke( + &tool, + &json!({ "address": "not_a_symbol", "text": "hi", "kind": "regular" }), + Some(&view), + ); + assert_eq!(error_code(&result), "invalid_params"); + } + + #[derive(Deserialize, JsonSchema)] + struct LenientArgs { + /// Optional value. + value: Option, + } + + struct LenientTool; + + impl TypedMcpTool for LenientTool { + type Args = LenientArgs; + + fn info(&self) -> McpToolInfo { + McpToolInfo::new("rust_mcp_typed_lenient", "Echo a value.") + .with_scope(McpToolScope::GlobalScope) + } + + fn invoke( + &self, + _call: &McpToolCall, + args: LenientArgs, + ) -> Result { + Ok(McpToolResult::structured(json!({ "value": args.value }))) + } + } + + #[test] + fn typed_tools_reject_undeclared_arguments() { + let _session = Session::new().expect("Failed to initialize session"); + let tool = register_typed_mcp_tool(LenientTool).expect("registration failed"); + assert_eq!(tool.input_schema()["additionalProperties"], false); + + assert_eq!( + invoke(&tool, &json!({ "value": 3 }), None)["structuredContent"]["value"], + 3 + ); + let result = invoke(&tool, &json!({ "value": 3, "other": 1 }), None); + assert_eq!(error_code(&result), "invalid_params"); + assert_eq!( + result["structuredContent"]["errorMessage"], + "Unexpected parameter 'other'" + ); + } + + #[test] + fn typed_tools_name_the_invalid_parameter() { + let _session = Session::new().expect("Failed to initialize session"); + let tool = create_typed_mcp_tool(CommentTool).expect("creation failed"); + let view = raw_view(); + + for (arguments, message) in [ + ( + json!({ "address": "0x20", "text": 5, "kind": "regular" }), + "Invalid parameter 'text': invalid type: integer `5`, expected a string", + ), + ( + json!({ "address": "0x20", "text": "hi", "kind": "regular", "length": -1 }), + "Invalid parameter 'length': invalid value: integer `-1`, expected an unsigned integer or an integer expression string", + ), + ( + json!({ "address": "0x20", "text": "hi", "kind": "regular", "length": 1.5 }), + "Invalid parameter 'length': invalid type: floating point `1.5`, expected an unsigned integer or an integer expression string", + ), + (json!({ "address": "0x20", "text": "hi" }), "missing field `kind`"), + ] { + let result = invoke(&tool, &arguments, Some(&view)); + assert_eq!(error_code(&result), "invalid_params", "{arguments}"); + assert_eq!(result["structuredContent"]["errorMessage"], message, "{arguments}"); + } + } + + #[derive(Deserialize, JsonSchema)] + struct LabelArgs { + /// Some text. + text: McpNonEmptyString, + /// A label. + label: Option, + } + + struct LabelTool; + + impl TypedMcpTool for LabelTool { + type Args = LabelArgs; + + fn info(&self) -> McpToolInfo { + McpToolInfo::new("rust_mcp_typed_non_empty", "Echo non-empty strings.") + .with_scope(McpToolScope::GlobalScope) + } + + fn invoke( + &self, + _call: &McpToolCall, + args: LabelArgs, + ) -> Result { + Ok(McpToolResult::structured( + json!({ "text": args.text.0, "label": args.label.map(|label| label.0) }), + )) + } + } + + #[test] + fn typed_tools_reject_empty_non_empty_strings() { + let _session = Session::new().expect("Failed to initialize session"); + let tool = create_typed_mcp_tool(LabelTool).expect("creation failed"); + let schema = tool.input_schema(); + assert_eq!( + schema["properties"]["text"], + json!({ "type": "string", "minLength": 1, "description": "Some text." }) + ); + assert_eq!(schema["properties"]["label"]["minLength"], 1); + + assert_eq!( + invoke(&tool, &json!({ "text": "a", "label": "b" }), None)["structuredContent"], + json!({ "text": "a", "label": "b" }) + ); + + for (arguments, message) in [ + ( + json!({ "text": "" }), + "Invalid parameter 'text': invalid value: string \"\", expected a non-empty string", + ), + ( + json!({ "text": "a", "label": "" }), + "Invalid parameter 'label': invalid value: string \"\", expected a non-empty string", + ), + (json!({}), "missing field `text`"), + ] { + let result = invoke(&tool, &arguments, None); + assert_eq!(error_code(&result), "invalid_params", "{arguments}"); + assert_eq!(result["structuredContent"]["errorMessage"], message, "{arguments}"); + } + } + + #[derive(Deserialize, JsonSchema)] + struct NoArgs {} + + struct NoArgsTool; + + impl TypedMcpTool for NoArgsTool { + type Args = NoArgs; + + fn info(&self) -> McpToolInfo { + McpToolInfo::new("rust_mcp_typed_no_args", "Take no arguments.") + .with_scope(McpToolScope::GlobalScope) + } + + fn invoke( + &self, + _call: &McpToolCall, + _args: NoArgs, + ) -> Result { + Ok(McpToolResult::text("ok")) + } + } + + #[test] + fn typed_tools_without_arguments_reject_every_argument() { + let _session = Session::new().expect("Failed to initialize session"); + let tool = create_typed_mcp_tool(NoArgsTool).expect("creation failed"); + assert_eq!(tool.input_schema()["type"], "object"); + assert_eq!(tool.input_schema()["additionalProperties"], false); + + assert_eq!(invoke(&tool, &json!({}), None)["content"][0]["text"], "ok"); + let result = invoke(&tool, &json!({ "other": 1 }), None); + assert_eq!(error_code(&result), "invalid_params"); + assert_eq!( + result["structuredContent"]["errorMessage"], + "Unexpected parameter 'other'" + ); + } + + #[derive(Deserialize, JsonSchema)] + struct NullableArgs { + /// Any JSON value. + value: Value, + /// An optional label. + label: Option, + /// A count. + #[serde(default)] + count: u32, + } + + struct NullableTool; + + impl TypedMcpTool for NullableTool { + type Args = NullableArgs; + + fn info(&self) -> McpToolInfo { + McpToolInfo::new("rust_mcp_typed_nullable", "Echo a JSON value.") + .with_scope(McpToolScope::GlobalScope) + } + + fn invoke( + &self, + _call: &McpToolCall, + args: NullableArgs, + ) -> Result { + Ok(McpToolResult::structured( + json!({ "value": args.value, "label": args.label, "count": args.count }), + )) + } + } + + #[test] + fn typed_tools_pass_null_to_required_parameters() { + let _session = Session::new().expect("Failed to initialize session"); + let tool = create_typed_mcp_tool(NullableTool).expect("creation failed"); + assert_eq!(tool.input_schema()["required"], json!(["value"])); + + assert_eq!( + invoke( + &tool, + &json!({ "value": null, "label": null, "count": null }), + None + )["structuredContent"], + json!({ "value": null, "label": null, "count": 0 }) + ); + + let result = invoke(&tool, &json!({}), None); + assert_eq!(error_code(&result), "invalid_params"); + assert_eq!( + result["structuredContent"]["errorMessage"], + "missing field `value`" + ); + } +} diff --git a/zensical.toml b/zensical.toml index d733504cb4..869accc1c7 100644 --- a/zensical.toml +++ b/zensical.toml @@ -67,6 +67,7 @@ nav = [ { "Cookbook" = "dev/cookbook.md" }, { "Writing Plugins" = "dev/plugins.md" }, { "Container Transforms" = "dev/containertransforms.md" }, + { "Writing MCP Tools" = "dev/mcp-tools.md" }, { "Automation" = "dev/batch.md" }, { "Architecture / Platform" = [ { "Part 1: Disassembly" = "dev/archplatform-disassembly.md" },