feat(cli): add mcp prompt support via slash commands (#1323)

This commit is contained in:
Kalvin C
2025-02-27 15:47:29 -08:00
committed by GitHub
parent 5bf05d545e
commit d0ca46983e
24 changed files with 958 additions and 82 deletions
+11
View File
@@ -1,3 +1,5 @@
use std::collections::HashMap;
use anyhow::Result;
use async_trait::async_trait;
use futures::stream::BoxStream;
@@ -6,6 +8,8 @@ use serde_json::Value;
use super::extension::{ExtensionConfig, ExtensionResult};
use crate::message::Message;
use crate::providers::base::ProviderUsage;
use mcp_core::prompt::Prompt;
use mcp_core::protocol::GetPromptResult;
/// Core trait defining the behavior of an Agent
#[async_trait]
@@ -37,4 +41,11 @@ pub trait Agent: Send + Sync {
/// Override the system prompt with custom text
async fn override_system_prompt(&mut self, template: String);
/// Lists all prompts from all extensions
async fn list_extension_prompts(&self) -> HashMap<String, Vec<Prompt>>;
/// Get a prompt result with the given name and arguments
/// Returns the prompt text that would be used as user input
async fn get_prompt(&self, name: &str, arguments: Value) -> Result<GetPromptResult>;
}
+101 -2
View File
@@ -1,6 +1,8 @@
use anyhow::Result;
use chrono::{DateTime, TimeZone, Utc};
use futures::stream::{FuturesUnordered, StreamExt};
use mcp_client::McpService;
use mcp_core::protocol::GetPromptResult;
use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use std::sync::LazyLock;
@@ -13,7 +15,7 @@ use crate::prompt_template::{load_prompt, load_prompt_file};
use crate::providers::base::{Provider, ProviderUsage};
use mcp_client::client::{ClientCapabilities, ClientInfo, McpClient, McpClientTrait};
use mcp_client::transport::{SseTransport, StdioTransport, Transport};
use mcp_core::{Content, Tool, ToolCall, ToolError, ToolResult};
use mcp_core::{prompt::Prompt, Content, Tool, ToolCall, ToolError, ToolResult};
use serde_json::Value;
// By default, we set it to Jan 1, 2020 if the resource does not have a timestamp
@@ -544,6 +546,87 @@ impl Capabilities {
result
}
pub async fn list_prompts_from_extension(
&self,
extension_name: &str,
) -> Result<Vec<Prompt>, ToolError> {
let client = self.clients.get(extension_name).ok_or_else(|| {
ToolError::InvalidParameters(format!("Extension {} is not valid", extension_name))
})?;
let client_guard = client.lock().await;
client_guard
.list_prompts(None)
.await
.map_err(|e| {
ToolError::ExecutionError(format!(
"Unable to list prompts for {}, {:?}",
extension_name, e
))
})
.map(|lp| lp.prompts)
}
pub async fn list_prompts(&self) -> Result<HashMap<String, Vec<Prompt>>, ToolError> {
let mut futures = FuturesUnordered::new();
for extension_name in self.clients.keys() {
futures.push(async move {
(
extension_name,
self.list_prompts_from_extension(extension_name).await,
)
});
}
let mut all_prompts = HashMap::new();
let mut errors = Vec::new();
// Process results as they complete
while let Some(result) = futures.next().await {
let (name, prompts) = result;
match prompts {
Ok(content) => {
all_prompts.insert(name.to_string(), content);
}
Err(tool_error) => {
errors.push(tool_error);
}
}
}
// Log any errors that occurred
if !errors.is_empty() {
tracing::error!(
errors = ?errors
.into_iter()
.map(|e| format!("{:?}", e))
.collect::<Vec<_>>(),
"errors from listing prompts"
);
}
Ok(all_prompts)
}
pub async fn get_prompt(
&self,
extension_name: &str,
name: &str,
arguments: Value,
) -> Result<GetPromptResult> {
let client = self
.clients
.get(extension_name)
.ok_or_else(|| anyhow::anyhow!("Extension {} not found", extension_name))?;
let client_guard = client.lock().await;
client_guard
.get_prompt(name, arguments)
.await
.map_err(|e| anyhow::anyhow!("Failed to get prompt: {}", e))
}
}
#[cfg(test)]
@@ -556,7 +639,8 @@ mod tests {
use mcp_client::client::Error;
use mcp_client::client::McpClientTrait;
use mcp_core::protocol::{
CallToolResult, InitializeResult, ListResourcesResult, ListToolsResult, ReadResourceResult,
CallToolResult, GetPromptResult, InitializeResult, ListPromptsResult, ListResourcesResult,
ListToolsResult, ReadResourceResult,
};
use serde_json::json;
@@ -625,6 +709,21 @@ mod tests {
_ => Err(Error::NotInitialized),
}
}
async fn list_prompts(
&self,
_next_cursor: Option<String>,
) -> Result<ListPromptsResult, Error> {
Err(Error::NotInitialized)
}
async fn get_prompt(
&self,
_name: &str,
_arguments: Value,
) -> Result<GetPromptResult, Error> {
Err(Error::NotInitialized)
}
}
#[test]
+35
View File
@@ -2,6 +2,7 @@
/// It makes no attempt to handle context limits, and cannot read resources
use async_trait::async_trait;
use futures::stream::BoxStream;
use std::collections::HashMap;
use tokio::sync::Mutex;
use tracing::{debug, instrument};
@@ -13,7 +14,10 @@ use crate::providers::base::Provider;
use crate::providers::base::ProviderUsage;
use crate::register_agent;
use crate::token_counter::TokenCounter;
use anyhow::{anyhow, Result};
use indoc::indoc;
use mcp_core::prompt::Prompt;
use mcp_core::protocol::GetPromptResult;
use mcp_core::tool::Tool;
use serde_json::{json, Value};
@@ -198,6 +202,37 @@ impl Agent for ReferenceAgent {
let mut capabilities = self.capabilities.lock().await;
capabilities.set_system_prompt_override(template);
}
async fn list_extension_prompts(&self) -> HashMap<String, Vec<Prompt>> {
let capabilities = self.capabilities.lock().await;
capabilities
.list_prompts()
.await
.expect("Failed to list prompts")
}
async fn get_prompt(&self, name: &str, arguments: Value) -> Result<GetPromptResult> {
let capabilities = self.capabilities.lock().await;
// First find which extension has this prompt
let prompts = capabilities
.list_prompts()
.await
.map_err(|e| anyhow!("Failed to list prompts: {}", e))?;
if let Some(extension) = prompts
.iter()
.find(|(_, prompt_list)| prompt_list.iter().any(|p| p.name == name))
.map(|(extension, _)| extension)
{
return capabilities
.get_prompt(extension, name, arguments)
.await
.map_err(|e| anyhow!("Failed to get prompt: {}", e));
}
Err(anyhow!("Prompt '{}' not found", name))
}
}
register_agent!("reference", ReferenceAgent);
+35
View File
@@ -2,6 +2,7 @@
/// It makes no attempt to handle context limits, and cannot read resources
use async_trait::async_trait;
use futures::stream::BoxStream;
use std::collections::HashMap;
use tokio::sync::mpsc;
use tokio::sync::Mutex;
use tracing::{debug, error, instrument, warn};
@@ -19,7 +20,10 @@ use crate::providers::errors::ProviderError;
use crate::register_agent;
use crate::token_counter::TokenCounter;
use crate::truncate::{truncate_messages, OldestFirstTruncation};
use anyhow::{anyhow, Result};
use indoc::indoc;
use mcp_core::prompt::Prompt;
use mcp_core::protocol::GetPromptResult;
use mcp_core::{tool::Tool, Content};
use serde_json::{json, Value};
@@ -398,6 +402,37 @@ impl Agent for TruncateAgent {
let mut capabilities = self.capabilities.lock().await;
capabilities.set_system_prompt_override(template);
}
async fn list_extension_prompts(&self) -> HashMap<String, Vec<Prompt>> {
let capabilities = self.capabilities.lock().await;
capabilities
.list_prompts()
.await
.expect("Failed to list prompts")
}
async fn get_prompt(&self, name: &str, arguments: Value) -> Result<GetPromptResult> {
let capabilities = self.capabilities.lock().await;
// First find which extension has this prompt
let prompts = capabilities
.list_prompts()
.await
.map_err(|e| anyhow!("Failed to list prompts: {}", e))?;
if let Some(extension) = prompts
.iter()
.find(|(_, prompt_list)| prompt_list.iter().any(|p| p.name == name))
.map(|(extension, _)| extension)
{
return capabilities
.get_prompt(extension, name, arguments)
.await
.map_err(|e| anyhow!("Failed to get prompt: {}", e));
}
Err(anyhow!("Prompt '{}' not found", name))
}
}
register_agent!("truncate", TruncateAgent);
+190
View File
@@ -10,6 +10,8 @@ use std::collections::HashSet;
use chrono::Utc;
use mcp_core::content::{Content, ImageContent, TextContent};
use mcp_core::handler::ToolResult;
use mcp_core::prompt::{PromptMessage, PromptMessageContent, PromptMessageRole};
use mcp_core::resource::ResourceContents;
use mcp_core::role::Role;
use mcp_core::tool::ToolCall;
use serde_json::Value;
@@ -156,6 +158,37 @@ impl From<Content> for MessageContent {
}
}
impl From<PromptMessage> for Message {
fn from(prompt_message: PromptMessage) -> Self {
// Create a new message with the appropriate role
let message = match prompt_message.role {
PromptMessageRole::User => Message::user(),
PromptMessageRole::Assistant => Message::assistant(),
};
// Convert and add the content
let content = match prompt_message.content {
PromptMessageContent::Text { text } => MessageContent::text(text),
PromptMessageContent::Image { image } => {
MessageContent::image(image.data, image.mime_type)
}
PromptMessageContent::Resource { resource } => {
// For resources, convert to text content with the resource text
match resource.resource {
ResourceContents::TextResourceContents { text, .. } => {
MessageContent::text(text)
}
ResourceContents::BlobResourceContents { blob, .. } => {
MessageContent::text(format!("[Binary content: {}]", blob))
}
}
}
};
message.with_content(content)
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
/// A message to or from an LLM
#[serde(rename_all = "camelCase")]
@@ -305,7 +338,10 @@ impl Message {
#[cfg(test)]
mod tests {
use super::*;
use mcp_core::content::EmbeddedResource;
use mcp_core::handler::ToolError;
use mcp_core::prompt::PromptMessageContent;
use mcp_core::resource::ResourceContents;
use serde_json::{json, Value};
#[test]
@@ -420,4 +456,158 @@ mod tests {
panic!("Expected ToolRequest content");
}
}
#[test]
fn test_from_prompt_message_text() {
let prompt_content = PromptMessageContent::Text {
text: "Hello, world!".to_string(),
};
let prompt_message = PromptMessage {
role: PromptMessageRole::User,
content: prompt_content,
};
let message = Message::from(prompt_message);
if let MessageContent::Text(text_content) = &message.content[0] {
assert_eq!(text_content.text, "Hello, world!");
} else {
panic!("Expected MessageContent::Text");
}
}
#[test]
fn test_from_prompt_message_image() {
let prompt_content = PromptMessageContent::Image {
image: ImageContent {
data: "base64data".to_string(),
mime_type: "image/jpeg".to_string(),
annotations: None,
},
};
let prompt_message = PromptMessage {
role: PromptMessageRole::User,
content: prompt_content,
};
let message = Message::from(prompt_message);
if let MessageContent::Image(image_content) = &message.content[0] {
assert_eq!(image_content.data, "base64data");
assert_eq!(image_content.mime_type, "image/jpeg");
} else {
panic!("Expected MessageContent::Image");
}
}
#[test]
fn test_from_prompt_message_text_resource() {
let resource = ResourceContents::TextResourceContents {
uri: "file:///test.txt".to_string(),
mime_type: Some("text/plain".to_string()),
text: "Resource content".to_string(),
};
let prompt_content = PromptMessageContent::Resource {
resource: EmbeddedResource {
resource,
annotations: None,
},
};
let prompt_message = PromptMessage {
role: PromptMessageRole::User,
content: prompt_content,
};
let message = Message::from(prompt_message);
if let MessageContent::Text(text_content) = &message.content[0] {
assert_eq!(text_content.text, "Resource content");
} else {
panic!("Expected MessageContent::Text");
}
}
#[test]
fn test_from_prompt_message_blob_resource() {
let resource = ResourceContents::BlobResourceContents {
uri: "file:///test.bin".to_string(),
mime_type: Some("application/octet-stream".to_string()),
blob: "binary_data".to_string(),
};
let prompt_content = PromptMessageContent::Resource {
resource: EmbeddedResource {
resource,
annotations: None,
},
};
let prompt_message = PromptMessage {
role: PromptMessageRole::User,
content: prompt_content,
};
let message = Message::from(prompt_message);
if let MessageContent::Text(text_content) = &message.content[0] {
assert_eq!(text_content.text, "[Binary content: binary_data]");
} else {
panic!("Expected MessageContent::Text");
}
}
#[test]
fn test_from_prompt_message() {
// Test user message conversion
let prompt_message = PromptMessage {
role: PromptMessageRole::User,
content: PromptMessageContent::Text {
text: "Hello, world!".to_string(),
},
};
let message = Message::from(prompt_message);
assert_eq!(message.role, Role::User);
assert_eq!(message.content.len(), 1);
assert_eq!(message.as_concat_text(), "Hello, world!");
// Test assistant message conversion
let prompt_message = PromptMessage {
role: PromptMessageRole::Assistant,
content: PromptMessageContent::Text {
text: "I can help with that.".to_string(),
},
};
let message = Message::from(prompt_message);
assert_eq!(message.role, Role::Assistant);
assert_eq!(message.content.len(), 1);
assert_eq!(message.as_concat_text(), "I can help with that.");
}
#[test]
fn test_message_with_text() {
let message = Message::user().with_text("Hello");
assert_eq!(message.as_concat_text(), "Hello");
}
#[test]
fn test_message_with_tool_request() {
let tool_call = Ok(ToolCall {
name: "test_tool".to_string(),
arguments: serde_json::json!({}),
});
let message = Message::assistant().with_tool_request("req1", tool_call);
assert!(message.is_tool_call());
assert!(!message.is_tool_response());
let ids = message.get_tool_ids();
assert_eq!(ids.len(), 1);
assert!(ids.contains("req1"));
}
}