From 08b73f300115299c7a0bba5b2ecf907b0f21a604 Mon Sep 17 00:00:00 2001 From: Kalvin Chau Date: Wed, 19 Feb 2025 15:04:26 -0800 Subject: [PATCH] feat: add prompts support to mcp-client, ahere to MCP spec for prompts - add new endpoints `list_prompts` and `get_prompt` in the MCP client - update prompt model in mcp-core to make `description` and `arguments` optional, following MCP spec --- crates/mcp-client/src/client.rs | 48 ++++++++++++++++++++++++++++++--- crates/mcp-core/src/prompt.rs | 28 ++++++++++++------- crates/mcp-server/src/router.rs | 27 ++++++++++--------- 3 files changed, 78 insertions(+), 25 deletions(-) diff --git a/crates/mcp-client/src/client.rs b/crates/mcp-client/src/client.rs index 0a00e8c..0d722e5 100644 --- a/crates/mcp-client/src/client.rs +++ b/crates/mcp-client/src/client.rs @@ -1,7 +1,7 @@ use mcp_core::protocol::{ - CallToolResult, Implementation, InitializeResult, JsonRpcError, JsonRpcMessage, - JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, ListResourcesResult, ListToolsResult, - ReadResourceResult, ServerCapabilities, METHOD_NOT_FOUND, + CallToolResult, GetPromptResult, Implementation, InitializeResult, JsonRpcError, + JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, ListPromptsResult, + ListResourcesResult, ListToolsResult, ReadResourceResult, ServerCapabilities, METHOD_NOT_FOUND, }; use serde::{Deserialize, Serialize}; use serde_json::Value; @@ -93,6 +93,10 @@ pub trait McpClientTrait: Send + Sync { async fn list_tools(&self, next_cursor: Option) -> Result; async fn call_tool(&self, name: &str, arguments: Value) -> Result; + + async fn list_prompts(&self, next_cursor: Option) -> Result; + + async fn get_prompt(&self, name: &str, arguments: Value) -> Result; } /// The MCP client is the interface for MCP operations. @@ -346,4 +350,42 @@ where // https://modelcontextprotocol.io/docs/concepts/tools#error-handling-2 self.send_request("tools/call", params).await } + + async fn list_prompts(&self, next_cursor: Option) -> Result { + if !self.completed_initialization() { + return Err(Error::NotInitialized); + } + + // If prompts is not supported, return an error + if self.server_capabilities.as_ref().unwrap().prompts.is_none() { + return Err(Error::RpcError { + code: METHOD_NOT_FOUND, + message: "Server does not support 'prompts' capability".to_string(), + }); + } + + let payload = next_cursor + .map(|cursor| serde_json::json!({"cursor": cursor})) + .unwrap_or_else(|| serde_json::json!({})); + + self.send_request("prompts/list", payload).await + } + + async fn get_prompt(&self, name: &str, arguments: Value) -> Result { + if !self.completed_initialization() { + return Err(Error::NotInitialized); + } + + // If prompts is not supported, return an error + if self.server_capabilities.as_ref().unwrap().prompts.is_none() { + return Err(Error::RpcError { + code: METHOD_NOT_FOUND, + message: "Server does not support 'prompts' capability".to_string(), + }); + } + + let params = serde_json::json!({ "name": name, "arguments": arguments }); + + self.send_request("prompts/get", params).await + } } diff --git a/crates/mcp-core/src/prompt.rs b/crates/mcp-core/src/prompt.rs index 7b814fd..4a0106e 100644 --- a/crates/mcp-core/src/prompt.rs +++ b/crates/mcp-core/src/prompt.rs @@ -10,22 +10,28 @@ use serde::{Deserialize, Serialize}; pub struct Prompt { /// The name of the prompt pub name: String, - /// A description of what the prompt does - pub description: String, - /// The arguments that can be passed to customize the prompt - pub arguments: Vec, + /// Optional description of what the prompt does + #[serde(skip_serializing_if = "Option::is_none")] + pub description: Option, + /// Optional arguments that can be passed to customize the prompt + #[serde(skip_serializing_if = "Option::is_none")] + pub arguments: Option>, } impl Prompt { /// Create a new prompt with the given name, description and arguments - pub fn new(name: N, description: D, arguments: Vec) -> Self + pub fn new( + name: N, + description: Option, + arguments: Option>, + ) -> Self where N: Into, D: Into, { Prompt { name: name.into(), - description: description.into(), + description: description.map(Into::into), arguments, } } @@ -37,9 +43,11 @@ pub struct PromptArgument { /// The name of the argument pub name: String, /// A description of what the argument is used for - pub description: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub description: Option, /// Whether this argument is required - pub required: bool, + #[serde(skip_serializing_if = "Option::is_none")] + pub required: Option, } /// Represents the role of a message sender in a prompt conversation @@ -151,6 +159,6 @@ pub struct PromptTemplate { #[derive(Debug, Serialize, Deserialize)] pub struct PromptArgumentTemplate { pub name: String, - pub description: String, - pub required: bool, + pub description: Option, + pub required: Option, } diff --git a/crates/mcp-server/src/router.rs b/crates/mcp-server/src/router.rs index d291831..0060ffd 100644 --- a/crates/mcp-server/src/router.rs +++ b/crates/mcp-server/src/router.rs @@ -305,18 +305,21 @@ pub trait Router: Send + Sync + 'static { }; // Validate required arguments - for arg in &prompt.arguments { - if arg.required - && (!arguments.contains_key(&arg.name) - || arguments - .get(&arg.name) - .and_then(Value::as_str) - .is_none_or(str::is_empty)) - { - return Err(RouterError::InvalidParams(format!( - "Missing required argument: '{}'", - arg.name - ))); + if let Some(args) = &prompt.arguments { + for arg in args { + if arg.required.is_some() + && arg.required.unwrap() + && (!arguments.contains_key(&arg.name) + || arguments + .get(&arg.name) + .and_then(Value::as_str) + .is_none_or(str::is_empty)) + { + return Err(RouterError::InvalidParams(format!( + "Missing required argument: '{}'", + arg.name + ))); + } } }