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);