feat(cli): add mcp prompt support via slash commands (#1323)
This commit is contained in:
@@ -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>;
|
||||
}
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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"));
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user