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);
|
||||
|
||||
Reference in New Issue
Block a user