From 61a3cfb662dbd4af7a55b32a3b1fbdc8b020a697 Mon Sep 17 00:00:00 2001 From: Alex Clarke Date: Mon, 24 Aug 2026 18:53:58 -0600 Subject: [PATCH] feat(repl): add .prompt command with live staged tab-completion MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Implements plans/mcp-resources-prompts-design.md §5.1/§5.4 (T6): - .prompt [key=value ...] fetches an MCP prompt and submits the result as chat input via Input::from_str + ask(), never through REPL line parsing; GetPromptResult messages are flattened into one user-role block with unconditional [user]/[assistant] labels - missing required prompt arguments are collected interactively - .list prompts renders server/name/description/args via the unified catalog (CatalogItem gains an arguments field), degrading per server - staged live tab-completion: enabled+running+prompts-capable servers (no RPC), then live prompt names, then key= argument suggestions with (required) markers; 2s timeout per RPC, all errors degrade to silent empty suggestions, ctx read guard dropped before blocking - the enabled-server alias expansion is factored into a shared helper used by both tool-scope rebuild and completion - BREAKING: the former .prompt temp-role builtin is renamed to .temp-role (behavior preserved); .prompt now belongs to MCP prompts, and a user macro named prompt or temp-role is shadowed --- src/config/macro_policy.rs | 28 ++ src/config/mod.rs | 1 + src/config/request_context.rs | 134 ++++++++-- src/config/tool_scope.rs | 487 +++++++++++++++++++++++++++++++++- src/mcp/mod.rs | 4 +- src/repl/completer.rs | 241 ++++++++++++++++- src/repl/mod.rs | 157 ++++++++++- 7 files changed, 1008 insertions(+), 44 deletions(-) diff --git a/src/config/macro_policy.rs b/src/config/macro_policy.rs index 4c08ed7..835a5bb 100644 --- a/src/config/macro_policy.rs +++ b/src/config/macro_policy.rs @@ -565,6 +565,34 @@ mod tests { assert_eq!(state_of(&policy, "a"), &MacroState::Enabled); } + #[test] + fn prompt_macro_is_shadowed_by_builtin() { + let policy = MacroPolicy::effective_with( + globals(&["prompt"]), + None, + None, + None, + None, + &crate::repl::builtin_command_names(), + ); + + assert_eq!(state_of(&policy, "prompt"), &MacroState::ShadowedBuiltin); + } + + #[test] + fn temp_role_macro_is_shadowed_by_builtin() { + let policy = MacroPolicy::effective_with( + globals(&["temp-role"]), + None, + None, + None, + None, + &crate::repl::builtin_command_names(), + ); + + assert_eq!(state_of(&policy, "temp-role"), &MacroState::ShadowedBuiltin); + } + #[test] fn locked_wins_over_shadowed_builtin() { let l = list(&["a"]); diff --git a/src/config/mod.rs b/src/config/mod.rs index 78cbd73..a1f190f 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -55,6 +55,7 @@ pub use self::skill_policy::SkillPolicy; pub use self::skill_registry::SkillRegistry; #[cfg(test)] pub(crate) use self::tool_scope::test_fixtures; +pub use self::tool_scope::{McpPromptCompletion, flatten_prompt_messages}; pub use self::update::run_self_update; use crate::client::{ self, ClientConfig, MessageContentToolCalls, Model, ModelType, OPENAI_COMPATIBLE_PROVIDERS, diff --git a/src/config/request_context.rs b/src/config/request_context.rs index 965df3f..b1ae09d 100644 --- a/src/config/request_context.rs +++ b/src/config/request_context.rs @@ -5,7 +5,7 @@ use super::skill::{SKILL_SCAFFOLD, Skill}; use super::skill_policy::SkillPolicy; use super::skill_registry::SkillRegistry; use super::todo::TodoList; -use super::tool_scope::{McpRuntime, ToolScope}; +use super::tool_scope::{McpPromptCompletion, McpRuntime, ToolScope, format_prompt_arguments}; use super::{ AGENTS_DIR_NAME, Agent, AgentVariables, AppConfig, AppState, AssetCategory, CREATE_TITLE_ROLE, Input, InstallFilter, LEFT_PROMPT, LastMessage, MESSAGES_FILE_NAME, MacroAllowlistLevel, @@ -23,8 +23,8 @@ use crate::function::{ user_interaction::USER_FUNCTION_PREFIX, }; use crate::mcp::{ - MCP_SEARCH_META_FUNCTION_NAME_PREFIX, McpAuthReason, McpAuthRequired, is_auth_required_error, - is_mcp_meta_function, mcp_meta_function_names, + CatalogItem, MCP_SEARCH_META_FUNCTION_NAME_PREFIX, McpAuthReason, McpAuthRequired, + McpServersConfig, is_auth_required_error, is_mcp_meta_function, mcp_meta_function_names, }; use crate::rag::Rag; use crate::supervisor::Supervisor; @@ -77,6 +77,29 @@ fn installed_bundle_names() -> Vec { } } +pub(crate) fn expand_enabled_mcp_server_ids( + app: &AppConfig, + mcp_config: &McpServersConfig, + enabled_mcp_servers: &[String], +) -> Vec { + if enabled_mcp_servers.iter().any(|s| s.trim() == "all") { + return mcp_config.mcp_servers.keys().cloned().collect(); + } + let mut ids = Vec::new(); + for item in enabled_mcp_servers.iter().map(|s| s.trim()) { + if mcp_config.mcp_servers.contains_key(item) { + ids.push(item.to_string()); + } else if let Some(mapped) = app.mapping_mcp_servers.get(item) { + for mapped_id in mapped.split(',').map(|s| s.trim()) { + if mcp_config.mcp_servers.contains_key(mapped_id) { + ids.push(mapped_id.to_string()); + } + } + } + } + ids +} + pub struct AutoContinueConfig { pub enabled: bool, pub max_continues: usize, @@ -139,6 +162,20 @@ pub(crate) fn asset_table(header: &[&str]) -> Table { table } +fn prompt_asset_rows(items: &[CatalogItem]) -> Vec<[String; 4]> { + items + .iter() + .map(|item| { + [ + item.server.clone(), + item.name.clone(), + item.description.clone(), + format_prompt_arguments(item.arguments.as_deref().unwrap_or_default()), + ] + }) + .collect() +} + fn complete_skills_with_descriptions(names: Vec) -> Vec<(String, Option)> { names .into_iter() @@ -2779,7 +2816,7 @@ impl RequestContext { } "bundles" => bundles::list_installed_bundles(), _ => bail!( - "Unknown kind '{kind}'. Valid kinds: roles, sessions, agents, rags, macros, skills, tools, mcp-servers, bundles" + "Unknown kind '{kind}'. Valid kinds: roles, sessions, agents, rags, macros, skills, prompts, tools, mcp-servers, bundles" ), } } @@ -3663,6 +3700,40 @@ impl RequestContext { fuzzy_filter(values, |v| v.0.as_str(), filter) } + pub fn mcp_prompt_completion(&self, args: &[&str]) -> McpPromptCompletion { + let app = self.app.config.as_ref(); + let enabled_ids = match &self.app.mcp_config { + Some(mcp_config) => { + let mut servers = self + .enabled_mcp_servers_for_current_scope(app, true) + .unwrap_or_default(); + servers.extend(self.skill_registry.loaded_mcp_servers()); + expand_enabled_mcp_server_ids(app, mcp_config, &servers) + } + None => vec![], + }; + self.tool_scope + .mcp_runtime + .prompt_completion(&enabled_ids, args) + } + + pub async fn list_prompt_assets(&self) -> Result<()> { + let items = self.tool_scope.mcp_runtime.prompt_catalog().await; + if items.is_empty() { + println!("No prompts found."); + return Ok(()); + } + + let mut table = asset_table(&["server", "name", "description", "args"]); + for row in prompt_asset_rows(&items) { + table.add_row(row.to_vec()); + } + + println!("Prompts:"); + println!("{table}"); + Ok(()) + } + async fn rebuild_tool_scope( &mut self, app: &AppConfig, @@ -3706,24 +3777,7 @@ impl RequestContext { && let Some(mcp_config) = &self.app.mcp_config { let server_ids: Vec = match &enabled_mcp_servers { - Some(servers) if servers.iter().any(|s| s.trim() == "all") => { - mcp_config.mcp_servers.keys().cloned().collect() - } - Some(servers) => { - let mut ids = Vec::new(); - for item in servers.iter().map(|s| s.trim()) { - if mcp_config.mcp_servers.contains_key(item) { - ids.push(item.to_string()); - } else if let Some(mapped) = app.mapping_mcp_servers.get(item) { - for mapped_id in mapped.split(',').map(|s| s.trim()) { - if mcp_config.mcp_servers.contains_key(mapped_id) { - ids.push(mapped_id.to_string()); - } - } - } - } - ids - } + Some(servers) => expand_enabled_mcp_server_ids(app, mcp_config, servers), None => vec![], }; @@ -4632,6 +4686,7 @@ mod tests { use crate::utils; use crate::utils::get_env_name; use crate::vault::Vault; + use rmcp::model::PromptArgument; use serde_json::json; use serial_test::serial; use std::env; @@ -4801,6 +4856,41 @@ mod tests { ); } + #[test] + fn prompt_asset_rows_assembles_columns() { + let items = vec![ + CatalogItem { + name: "summarize".to_string(), + server: "docs".to_string(), + description: "Summarize a document".to_string(), + arguments: Some(vec![ + PromptArgument::new("path").with_required(true), + PromptArgument::new("style"), + ]), + ..Default::default() + }, + CatalogItem { + name: "greet".to_string(), + server: "misc".to_string(), + ..Default::default() + }, + ]; + + let rows = prompt_asset_rows(&items); + + assert_eq!(rows.len(), 2); + assert_eq!( + rows[0], + [ + "docs", + "summarize", + "Summarize a document", + "path (required), style" + ] + ); + assert_eq!(rows[1], ["misc", "greet", "", ""]); + } + #[test] fn extract_role_returns_standalone_role() { let mut ctx = create_test_ctx(); diff --git a/src/config/tool_scope.rs b/src/config/tool_scope.rs index eaac7aa..06402cc 100644 --- a/src/config/tool_scope.rs +++ b/src/config/tool_scope.rs @@ -4,8 +4,9 @@ use crate::mcp::{CatalogItem, CatalogItemKind, ConnectedServer, McpRegistry, Mcp use anyhow::{Context, Result, anyhow}; use bm25::{Document, Language, SearchEngineBuilder}; use rmcp::model::{ - CallToolRequestParams, CallToolResult, Prompt, ReadResourceRequestParams, ReadResourceResult, - Resource, ResourceTemplate, Tool, + CallToolRequestParams, CallToolResult, ContentBlock, GetPromptRequestParams, GetPromptResult, + Prompt, PromptArgument, PromptMessage, ReadResourceRequestParams, ReadResourceResult, Resource, + ResourceTemplate, Role, Tool, }; use serde_json::{Value, json}; use std::collections::HashMap; @@ -28,6 +29,18 @@ impl Default for ToolScope { } } +pub enum McpPromptCompletion { + Ready(Vec<(String, Option)>), + PromptNames { + server: Arc, + }, + ArgumentKeys { + server: Arc, + prompt: String, + typed_keys: Vec, + }, +} + #[derive(Default, Clone)] pub struct McpRuntime { pub servers: HashMap>, @@ -291,6 +304,130 @@ impl McpRuntime { .await .map_err(Into::into) } + + pub async fn list_prompts(&self, server: &str) -> Result> { + let server_handle = self + .get(server) + .cloned() + .with_context(|| format!("Prompt MCP server does not exist: {server}"))?; + + server_handle.list_all_prompts().await.map_err(Into::into) + } + + pub async fn prompt( + &self, + server: &str, + name: &str, + arguments: HashMap, + ) -> Result { + let server_handle = self + .get(server) + .cloned() + .with_context(|| format!("Prompt MCP server does not exist: {server}"))?; + + let mut request = GetPromptRequestParams::new(name.to_owned()); + if !arguments.is_empty() { + request.arguments = Some( + arguments + .into_iter() + .map(|(key, value)| (key, Value::String(value))) + .collect(), + ); + } + + server_handle.get_prompt(request).await.map_err(Into::into) + } + + pub async fn prompt_catalog(&self) -> Vec { + let mut items = Vec::new(); + for features in self.server_features() { + if !features.prompts { + continue; + } + let Some(server_handle) = self.get(&features.name) else { + continue; + }; + match server_handle.list_all_prompts().await { + Ok(prompts) => items.extend( + prompts + .into_iter() + .map(|prompt| prompt_catalog_item(&features.name, prompt)), + ), + Err(e) => warn!( + "Failed to list prompts on MCP server {}: {e}", + features.name + ), + } + } + items + } + + pub fn prompt_completion( + &self, + enabled_servers: &[String], + args: &[&str], + ) -> McpPromptCompletion { + match args { + [] | [_] => McpPromptCompletion::Ready( + self.server_features() + .into_iter() + .filter(|features| features.prompts && enabled_servers.contains(&features.name)) + .map(|features| (features.name, None)) + .collect(), + ), + [server, _] => match self.get(server) { + Some(handle) => McpPromptCompletion::PromptNames { + server: Arc::clone(handle), + }, + None => McpPromptCompletion::Ready(vec![]), + }, + [server, prompt, rest @ ..] => match self.get(server) { + Some(handle) => McpPromptCompletion::ArgumentKeys { + server: Arc::clone(handle), + prompt: (*prompt).to_string(), + typed_keys: rest + .iter() + .filter_map(|arg| arg.split_once('=').map(|(key, _)| key.to_string())) + .collect(), + }, + None => McpPromptCompletion::Ready(vec![]), + }, + } + } +} + +pub fn flatten_prompt_messages(messages: &[PromptMessage]) -> String { + messages + .iter() + .map(|message| { + let label = match message.role { + Role::User => "[user]", + Role::Assistant => "[assistant]", + }; + let content = match &message.content { + ContentBlock::Text(text) => text.text.clone(), + ContentBlock::Image(_) => "[image content omitted]".to_string(), + ContentBlock::Audio(_) => "[audio content omitted]".to_string(), + _ => "[resource content omitted]".to_string(), + }; + format!("{label}\n{content}") + }) + .collect::>() + .join("\n\n") +} + +pub fn format_prompt_arguments(arguments: &[PromptArgument]) -> String { + arguments + .iter() + .map(|arg| { + if arg.required == Some(true) { + format!("{} (required)", arg.name) + } else { + arg.name.clone() + } + }) + .collect::>() + .join(", ") } fn catalog_key(item: &CatalogItem) -> String { @@ -326,6 +463,7 @@ fn resource_catalog_item(server: &str, resource: Resource) -> CatalogItem { uri: Some(resource.uri), mime_type: resource.mime_type, size: resource.size, + arguments: None, } } @@ -338,6 +476,7 @@ fn resource_template_catalog_item(server: &str, template: ResourceTemplate) -> C uri: Some(template.uri_template), mime_type: template.mime_type, size: None, + arguments: None, } } @@ -347,6 +486,7 @@ fn prompt_catalog_item(server: &str, prompt: Prompt) -> CatalogItem { name: prompt.name, server: server.to_string(), description: prompt.description.unwrap_or_default(), + arguments: prompt.arguments, ..Default::default() } } @@ -370,14 +510,15 @@ pub(crate) mod test_fixtures { use base64::Engine; use base64::engine::general_purpose::STANDARD; use rmcp::model::{ - ErrorData, ListPromptsResult, ListResourceTemplatesResult, ListResourcesResult, - ListToolsResult, PaginatedRequestParams, PromptArgument, PromptsCapability, + ErrorData, GetPromptResponse, ListPromptsResult, ListResourceTemplatesResult, + ListResourcesResult, ListToolsResult, PaginatedRequestParams, PromptsCapability, ReadResourceResponse, ResourceContents, ResourcesCapability, ServerCapabilities, ServerInfo, }; use rmcp::service::{RequestContext, RunningService}; use rmcp::{RoleServer, ServerHandler, ServiceExt}; use std::sync::atomic::{AtomicUsize, Ordering}; + use std::time::Duration; pub(crate) const FIXTURE_LOG_URI: &str = "file:///app.log"; pub(crate) const FIXTURE_LOG_TEXT: &str = "début of the log\n\ @@ -397,8 +538,12 @@ pub(crate) mod test_fixtures { pub(crate) resources_capability: bool, pub(crate) prompts_capability: bool, pub(crate) fail_resource_listings: bool, + pub(crate) fail_prompt_listings: bool, + pub(crate) fail_get_prompt: bool, + pub(crate) prompt_delay: Option, pub(crate) list_resources_calls: Arc, pub(crate) list_prompts_calls: Arc, + pub(crate) get_prompt_calls: Arc, } impl Default for FixtureServer { @@ -408,8 +553,12 @@ pub(crate) mod test_fixtures { resources_capability: false, prompts_capability: false, fail_resource_listings: false, + fail_prompt_listings: false, + fail_get_prompt: false, + prompt_delay: None, list_resources_calls: Arc::default(), list_prompts_calls: Arc::default(), + get_prompt_calls: Arc::default(), } } } @@ -518,6 +667,12 @@ pub(crate) mod test_fixtures { _context: RequestContext, ) -> Result { self.list_prompts_calls.fetch_add(1, Ordering::SeqCst); + if let Some(delay) = self.prompt_delay { + tokio::time::sleep(delay).await; + } + if self.fail_prompt_listings { + return Err(ErrorData::internal_error("prompt listing exploded", None)); + } Ok(ListPromptsResult::with_all_items(vec![Prompt::new( "summarize", Some("Summarize a document"), @@ -529,22 +684,71 @@ pub(crate) mod test_fixtures { ]), )])) } + + async fn get_prompt( + &self, + request: GetPromptRequestParams, + _context: RequestContext, + ) -> Result { + self.get_prompt_calls.fetch_add(1, Ordering::SeqCst); + if let Some(delay) = self.prompt_delay { + tokio::time::sleep(delay).await; + } + if self.fail_get_prompt { + return Err(ErrorData::internal_error("get_prompt exploded", None)); + } + let messages = match request.name.as_str() { + "summarize" => { + let path = request + .arguments + .as_ref() + .and_then(|args| args.get("path")) + .and_then(|value| value.as_str()) + .unwrap_or_default(); + vec![ + PromptMessage::new_text(Role::User, format!("Summarize {path}")), + PromptMessage::new_text(Role::Assistant, "In which style?"), + PromptMessage::new_text(Role::User, "Concise."), + ] + } + "hostile" => vec![ + PromptMessage::new_text(Role::User, "!rm -rf /"), + PromptMessage::new_text(Role::User, ".session hijack"), + ], + other => { + return Err(ErrorData::invalid_params( + format!("Unknown prompt: {other}"), + None, + )); + } + }; + Ok(GetPromptResponse::Complete(GetPromptResult::new(messages))) + } + } + + pub(crate) async fn add_fixture_server( + runtime: &mut McpRuntime, + name: &str, + fixture: FixtureServer, + ) -> RunningService { + let (client_io, server_io) = tokio::io::duplex(4096); + let (server, client) = tokio::join!(fixture.serve(server_io), ().serve(client_io)); + runtime.insert(name.to_string(), Arc::new(client.unwrap())); + server.unwrap() } pub(crate) async fn fixture_runtime( fixture: FixtureServer, ) -> (McpRuntime, RunningService) { - let (client_io, server_io) = tokio::io::duplex(4096); - let (server, client) = tokio::join!(fixture.serve(server_io), ().serve(client_io)); let mut runtime = McpRuntime::new(); - runtime.insert("fixture".to_string(), Arc::new(client.unwrap())); - (runtime, server.unwrap()) + let server = add_fixture_server(&mut runtime, "fixture", fixture).await; + (runtime, server) } } #[cfg(test)] mod tests { - use super::test_fixtures::{FixtureServer, fixture_runtime}; + use super::test_fixtures::{FixtureServer, add_fixture_server, fixture_runtime}; use super::*; use crate::function::ToolCall; use log::{Level, LevelFilter, Log, Metadata, Record}; @@ -908,4 +1112,269 @@ mod tests { "file:///missing not found in fixture MCP server resource catalog" ); } + + #[tokio::test] + async fn prompt_returns_result_and_counts_calls() { + let fixture = FixtureServer { + prompts_capability: true, + ..Default::default() + }; + let get_prompt_calls = Arc::clone(&fixture.get_prompt_calls); + let (runtime, _server) = fixture_runtime(fixture).await; + + let result = runtime + .prompt( + "fixture", + "summarize", + HashMap::from([("path".to_string(), "notes.txt".to_string())]), + ) + .await + .unwrap(); + + assert_eq!(result.messages.len(), 3); + assert_eq!( + result.messages[0] + .content + .as_text() + .map(|t| t.text.as_str()), + Some("Summarize notes.txt") + ); + assert_eq!(get_prompt_calls.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn prompt_missing_server_errors() { + let runtime = McpRuntime::new(); + + let err = runtime + .prompt("ghost", "summarize", HashMap::new()) + .await + .unwrap_err() + .to_string(); + + assert_eq!(err, "Prompt MCP server does not exist: ghost"); + } + + #[tokio::test] + async fn prompt_surfaces_server_failure() { + let fixture = FixtureServer { + prompts_capability: true, + fail_get_prompt: true, + ..Default::default() + }; + let (runtime, _server) = fixture_runtime(fixture).await; + + let err = runtime + .prompt("fixture", "summarize", HashMap::new()) + .await + .unwrap_err() + .to_string(); + + assert!( + err.contains("get_prompt exploded"), + "unexpected error: {err}" + ); + } + + #[test] + fn flatten_prompt_messages_labels_every_message() { + let messages = vec![ + PromptMessage::new_text(Role::User, "Summarize notes.txt"), + PromptMessage::new_text(Role::Assistant, "In which style?"), + PromptMessage::new_text(Role::User, "Concise."), + ]; + + assert_eq!( + flatten_prompt_messages(&messages), + "[user]\nSummarize notes.txt\n\n[assistant]\nIn which style?\n\n[user]\nConcise." + ); + } + + #[test] + fn flatten_prompt_messages_labels_single_message() { + let messages = vec![PromptMessage::new_text(Role::User, "hello")]; + + assert_eq!(flatten_prompt_messages(&messages), "[user]\nhello"); + } + + #[test] + fn flatten_prompt_messages_replaces_non_text_content() { + let messages = vec![PromptMessage::new( + Role::User, + ContentBlock::image("aGk=", "image/png"), + )]; + + assert_eq!( + flatten_prompt_messages(&messages), + "[user]\n[image content omitted]" + ); + } + + #[tokio::test] + async fn flattened_hostile_prompt_cannot_start_a_command_line() { + let fixture = FixtureServer { + prompts_capability: true, + ..Default::default() + }; + let (runtime, _server) = fixture_runtime(fixture).await; + + let result = runtime + .prompt("fixture", "hostile", HashMap::new()) + .await + .unwrap(); + let flattened = flatten_prompt_messages(&result.messages); + + assert!(flattened.starts_with("[user]")); + assert!(!flattened.starts_with('!')); + assert!(!flattened.starts_with('.')); + assert!(flattened.contains("!rm -rf /")); + assert!(flattened.contains(".session hijack")); + } + + #[test] + fn format_prompt_arguments_marks_required() { + let arguments = vec![ + PromptArgument::new("path").with_required(true), + PromptArgument::new("style"), + ]; + + assert_eq!( + format_prompt_arguments(&arguments), + "path (required), style" + ); + assert_eq!(format_prompt_arguments(&[]), ""); + } + + #[tokio::test] + async fn prompt_catalog_carries_arguments() { + let fixture = FixtureServer { + prompts_capability: true, + ..Default::default() + }; + let (runtime, _server) = fixture_runtime(fixture).await; + + let items = runtime.prompt_catalog().await; + + assert_eq!(items.len(), 1); + let item = &items[0]; + assert_eq!(item.kind, CatalogItemKind::Prompt); + assert_eq!(item.server, "fixture"); + assert_eq!(item.name, "summarize"); + assert_eq!(item.description, "Summarize a document"); + let arguments = item.arguments.as_deref().unwrap(); + assert_eq!(format_prompt_arguments(arguments), "path (required), style"); + } + + #[tokio::test] + async fn prompt_catalog_degrades_when_one_server_fails() { + install_warn_collector(); + let healthy = FixtureServer { + prompts_capability: true, + ..Default::default() + }; + let failing = FixtureServer { + prompts_capability: true, + fail_prompt_listings: true, + ..Default::default() + }; + let (mut runtime, _server) = fixture_runtime(healthy).await; + let _failing_server = add_fixture_server(&mut runtime, "broken", failing).await; + + let items = runtime.prompt_catalog().await; + + assert_eq!(items.len(), 1); + assert_eq!(items[0].server, "fixture"); + let messages = warn_messages().lock().unwrap(); + assert!( + messages + .iter() + .any(|msg| msg.contains("Failed to list prompts on MCP server broken")), + "missing prompt-listing warning in: {messages:?}" + ); + } + + #[tokio::test] + async fn prompt_catalog_skips_servers_without_prompts_capability() { + let fixture = FixtureServer::default(); + let prompts_calls = Arc::clone(&fixture.list_prompts_calls); + let (runtime, _server) = fixture_runtime(fixture).await; + + let items = runtime.prompt_catalog().await; + + assert!(items.is_empty()); + assert_eq!(prompts_calls.load(Ordering::SeqCst), 0); + } + + #[tokio::test] + async fn prompt_completion_stage_one_uses_local_state_only() { + let fixture = FixtureServer { + prompts_capability: true, + ..Default::default() + }; + let list_prompts_calls = Arc::clone(&fixture.list_prompts_calls); + let get_prompt_calls = Arc::clone(&fixture.get_prompt_calls); + let (mut runtime, _server) = fixture_runtime(fixture).await; + let _tools_only = + add_fixture_server(&mut runtime, "tools-only", FixtureServer::default()).await; + + let stage = + runtime.prompt_completion(&["fixture".to_string(), "tools-only".to_string()], &[""]); + + let McpPromptCompletion::Ready(values) = stage else { + panic!("stage one must not require an RPC"); + }; + assert_eq!(values, vec![("fixture".to_string(), None)]); + assert_eq!(list_prompts_calls.load(Ordering::SeqCst), 0); + assert_eq!(get_prompt_calls.load(Ordering::SeqCst), 0); + } + + #[tokio::test] + async fn prompt_completion_stage_one_excludes_disabled_servers() { + let fixture = FixtureServer { + prompts_capability: true, + ..Default::default() + }; + let (runtime, _server) = fixture_runtime(fixture).await; + + let McpPromptCompletion::Ready(values) = runtime.prompt_completion(&[], &[""]) else { + panic!("stage one must not require an RPC"); + }; + + assert!(values.is_empty()); + } + + #[tokio::test] + async fn prompt_completion_unknown_server_is_empty() { + let (runtime, _server) = fixture_runtime(FixtureServer::default()).await; + + for args in [ + ["ghost", ""].as_slice(), + ["ghost", "summarize", ""].as_slice(), + ] { + let McpPromptCompletion::Ready(values) = runtime.prompt_completion(&[], args) else { + panic!("unknown server must degrade to empty suggestions"); + }; + assert!(values.is_empty()); + } + } + + #[tokio::test] + async fn prompt_completion_later_stages_carry_typed_keys() { + let fixture = FixtureServer { + prompts_capability: true, + ..Default::default() + }; + let (runtime, _server) = fixture_runtime(fixture).await; + + let stage = runtime.prompt_completion(&[], &["fixture", "summarize", "path=x", "sty"]); + + let McpPromptCompletion::ArgumentKeys { + prompt, typed_keys, .. + } = stage + else { + panic!("expected the argument-key stage"); + }; + assert_eq!(prompt, "summarize"); + assert_eq!(typed_keys, vec!["path".to_string()]); + } } diff --git a/src/mcp/mod.rs b/src/mcp/mod.rs index d426ff4..bf1d21a 100644 --- a/src/mcp/mod.rs +++ b/src/mcp/mod.rs @@ -16,7 +16,7 @@ use futures_util::{StreamExt, TryStreamExt, stream}; use http::{HeaderName, HeaderValue}; use indexmap::IndexMap; use indoc::formatdoc; -use rmcp::model::ServerCapabilities; +use rmcp::model::{PromptArgument, ServerCapabilities}; use rmcp::service::RunningService; use rmcp::transport::StreamableHttpClientTransport; use rmcp::transport::TokioChildProcess; @@ -123,6 +123,8 @@ pub struct CatalogItem { pub mime_type: Option, #[serde(skip_serializing_if = "Option::is_none")] pub size: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub arguments: Option>, } #[derive(Debug, Clone, Deserialize, Serialize)] diff --git a/src/repl/completer.rs b/src/repl/completer.rs index 440edbb..e71fbc3 100644 --- a/src/repl/completer.rs +++ b/src/repl/completer.rs @@ -1,11 +1,17 @@ use super::{REPL_COMMANDS, ReplCommand}; -use crate::{config::RequestContext, utils::fuzzy_filter}; +use crate::config::{McpPromptCompletion, RequestContext}; +use crate::mcp::ConnectedServer; +use crate::utils::fuzzy_filter; use parking_lot::RwLock; use reedline::{Completer, Span, Suggestion}; +use rmcp::model::Prompt; use std::collections::HashMap; use std::sync::Arc; +use std::time::Duration; + +const PROMPT_COMPLETION_RPC_TIMEOUT: Duration = Duration::from_secs(2); impl Completer for ReplCompleter { fn complete(&mut self, line: &str, pos: usize) -> Vec { @@ -29,6 +35,22 @@ impl Completer for ReplCompleter { return suggestions; } + if cmd == ".prompt" && parts_len > 1 { + let span = Span::new(parts[parts_len - 1].1, pos); + let args: Vec<&str> = parts.iter().skip(1).map(|(v, _)| *v).collect(); + let filter = args.last().copied().unwrap_or_default().to_string(); + let stage = { + let ctx = self.ctx.read(); + ctx.mcp_prompt_completion(&args) + }; + return complete_prompt_stage(stage, &filter, PROMPT_COMPLETION_RPC_TIMEOUT) + .iter() + .map(|(value, description)| { + create_suggestion(value, description.as_deref().unwrap_or_default(), span) + }) + .collect(); + } + let ctx = self.ctx.read(); let state = ctx.state(); let model_has_reasoning = !ctx.current_model().reasoning_levels().is_empty(); @@ -141,6 +163,57 @@ fn create_suggestion(value: &str, description: &str, span: Span) -> Suggestion { } } +fn complete_prompt_stage( + stage: McpPromptCompletion, + filter: &str, + rpc_timeout: Duration, +) -> Vec<(String, Option)> { + let values = match stage { + McpPromptCompletion::Ready(values) => values, + McpPromptCompletion::PromptNames { server } => list_prompts_blocking(server, rpc_timeout) + .unwrap_or_default() + .into_iter() + .map(|prompt| (prompt.name, prompt.description)) + .collect(), + McpPromptCompletion::ArgumentKeys { + server, + prompt, + typed_keys, + } => list_prompts_blocking(server, rpc_timeout) + .unwrap_or_default() + .into_iter() + .find(|candidate| candidate.name == prompt) + .and_then(|candidate| candidate.arguments) + .unwrap_or_default() + .into_iter() + .filter(|arg| !typed_keys.contains(&arg.name)) + .map(|arg| { + let description = match (arg.required == Some(true), arg.description) { + (true, Some(description)) => Some(format!("{description} (required)")), + (true, None) => Some("(required)".to_string()), + (false, description) => description, + }; + (format!("{}=", arg.name), description) + }) + .collect(), + }; + fuzzy_filter(values, |(value, _)| value.as_str(), filter) +} + +fn list_prompts_blocking( + server: Arc, + rpc_timeout: Duration, +) -> Option> { + let fut = async move { tokio::time::timeout(rpc_timeout, server.list_all_prompts()).await }; + // block_in_place is only sound because the REPL's read_line runs inside the + // main-thread block_on of the multi-thread runtime. + let result = match tokio::runtime::Handle::try_current().ok() { + Some(handle) => tokio::task::block_in_place(|| handle.block_on(fut)), + None => tokio::runtime::Runtime::new().ok()?.block_on(fut), + }; + result.ok()?.ok() +} + fn split_line(line: &str) -> Vec<(&str, usize)> { let mut parts = vec![]; let mut part_start = None; @@ -178,3 +251,169 @@ fn test_split_line() { vec![(".set", 0), ("highlight", 5), ("t", 15)], ); } + +#[cfg(test)] +mod prompt_completion_tests { + use super::*; + use crate::config::test_fixtures::{FixtureServer, fixture_runtime}; + use std::sync::atomic::Ordering; + + fn prompts_fixture() -> FixtureServer { + FixtureServer { + prompts_capability: true, + ..Default::default() + } + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn stage_two_lists_prompt_names_with_descriptions() { + let (runtime, _server) = fixture_runtime(prompts_fixture()).await; + let server = runtime.get("fixture").cloned().unwrap(); + + let values = complete_prompt_stage( + McpPromptCompletion::PromptNames { server }, + "", + Duration::from_secs(2), + ); + + assert_eq!( + values, + vec![( + "summarize".to_string(), + Some("Summarize a document".to_string()) + )] + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn stage_three_suggests_argument_keys_with_required_marker() { + let (runtime, _server) = fixture_runtime(prompts_fixture()).await; + let server = runtime.get("fixture").cloned().unwrap(); + + let values = complete_prompt_stage( + McpPromptCompletion::ArgumentKeys { + server, + prompt: "summarize".to_string(), + typed_keys: vec![], + }, + "", + Duration::from_secs(2), + ); + + assert_eq!( + values, + vec![ + ( + "path=".to_string(), + Some("Document path (required)".to_string()) + ), + ("style=".to_string(), None), + ] + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn stage_three_excludes_typed_keys_and_fuzzy_filters() { + let (runtime, _server) = fixture_runtime(prompts_fixture()).await; + let server = runtime.get("fixture").cloned().unwrap(); + + let values = complete_prompt_stage( + McpPromptCompletion::ArgumentKeys { + server: Arc::clone(&server), + prompt: "summarize".to_string(), + typed_keys: vec!["path".to_string()], + }, + "", + Duration::from_secs(2), + ); + assert_eq!(values, vec![("style=".to_string(), None)]); + + let values = complete_prompt_stage( + McpPromptCompletion::ArgumentKeys { + server, + prompt: "summarize".to_string(), + typed_keys: vec![], + }, + "sty", + Duration::from_secs(2), + ); + assert_eq!(values, vec![("style=".to_string(), None)]); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn stage_three_unknown_prompt_is_empty() { + let (runtime, _server) = fixture_runtime(prompts_fixture()).await; + let server = runtime.get("fixture").cloned().unwrap(); + + let values = complete_prompt_stage( + McpPromptCompletion::ArgumentKeys { + server, + prompt: "ghost".to_string(), + typed_keys: vec![], + }, + "", + Duration::from_secs(2), + ); + + assert!(values.is_empty()); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn slow_listing_times_out_to_empty() { + let fixture = FixtureServer { + prompt_delay: Some(Duration::from_millis(200)), + ..prompts_fixture() + }; + let (runtime, _server) = fixture_runtime(fixture).await; + let server = runtime.get("fixture").cloned().unwrap(); + + let values = complete_prompt_stage( + McpPromptCompletion::PromptNames { server }, + "", + Duration::from_millis(20), + ); + + assert!(values.is_empty()); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn failed_listing_is_swallowed_without_retry() { + let fixture = FixtureServer { + fail_prompt_listings: true, + ..prompts_fixture() + }; + let list_prompts_calls = Arc::clone(&fixture.list_prompts_calls); + let (runtime, _server) = fixture_runtime(fixture).await; + let server = runtime.get("fixture").cloned().unwrap(); + + let values = complete_prompt_stage( + McpPromptCompletion::PromptNames { server }, + "", + Duration::from_secs(2), + ); + + assert!(values.is_empty()); + assert_eq!(list_prompts_calls.load(Ordering::SeqCst), 1); + } + + #[test] + fn bridge_without_ambient_runtime_uses_fallback_runtime() { + let rt = tokio::runtime::Runtime::new().unwrap(); + let (runtime, _server) = rt.block_on(fixture_runtime(prompts_fixture())); + let server = runtime.get("fixture").cloned().unwrap(); + + let values = complete_prompt_stage( + McpPromptCompletion::PromptNames { server }, + "", + Duration::from_secs(2), + ); + + assert_eq!( + values, + vec![( + "summarize".to_string(), + Some("Summarize a document".to_string()) + )] + ); + } +} diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 3c159f2..6a4a7db 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -13,7 +13,7 @@ use crate::client::{ }; use crate::config::{ AgentVariables, AppConfig, AssertState, Input, LastMessage, MacroState, RequestContext, - StateFlags, macro_execute, + StateFlags, flatten_prompt_messages, macro_execute, }; use crate::config::{AssetCategory, paths}; use crate::function::supervisor::{GuardrailAction, check_pending_agents_guardrail}; @@ -29,6 +29,7 @@ use anyhow::{Context, Result, bail}; use crossterm::cursor::SetCursorStyle; use fancy_regex::Regex; use indoc::indoc; +use inquire::Text; use log::warn; use parking_lot::RwLock; use reedline::CursorConfig; @@ -38,6 +39,8 @@ use reedline::{ default_emacs_keybindings, default_vi_insert_keybindings, default_vi_normal_keybindings, }; use reedline::{MenuBuilder, Signal}; +use rmcp::model::PromptArgument; +use std::collections::HashMap; use std::sync::LazyLock; use std::{env, process, sync::Arc}; use tokio::task; @@ -53,7 +56,7 @@ pub const DEFAULT_CONTINUATION_PROMPT: &str = indoc! {" 4. Continue with the next pending item now. Call tools immediately." }; -static REPL_COMMANDS: LazyLock<[ReplCommand; 61]> = LazyLock::new(|| { +static REPL_COMMANDS: LazyLock<[ReplCommand; 62]> = LazyLock::new(|| { [ ReplCommand::new(".help", "Show this help guide", AssertState::pass()), ReplCommand::new(".info", "Show system info", AssertState::pass()), @@ -105,6 +108,11 @@ static REPL_COMMANDS: LazyLock<[ReplCommand; 61]> = LazyLock::new(|| { ReplCommand::new(".model", "Switch LLM model", AssertState::pass()), ReplCommand::new( ".prompt", + "Invoke an MCP prompt and submit the result as chat input", + AssertState::pass(), + ), + ReplCommand::new( + ".temp-role", "Set a temporary role using a prompt", AssertState::False(StateFlags::SESSION | StateFlags::AGENT), ), @@ -307,7 +315,7 @@ static REPL_COMMANDS: LazyLock<[ReplCommand; 61]> = LazyLock::new(|| { ), ReplCommand::new( ".list", - "List roles, sessions, agents, RAGs, macros, skills, tools, MCP servers, or bundles", + "List roles, sessions, agents, RAGs, macros, skills, prompts, tools, MCP servers, or bundles", AssertState::pass(), ), ReplCommand::new( @@ -774,12 +782,46 @@ pub async fn run_repl_command( .tool disable # Disable a single tool in the current context"# ), }, - ".prompt" => match args { + ".prompt" => { + let (words, _) = split_args_text(args.unwrap_or_default(), cfg!(windows)); + match words.as_slice() { + [server, name, rest @ ..] => { + let provided = parse_prompt_call_args(rest)?; + let prompts = ctx.tool_scope.mcp_runtime.list_prompts(server).await?; + let declared = prompts + .into_iter() + .find(|prompt| prompt.name == *name) + .with_context(|| { + format!("Prompt '{name}' not found on MCP server '{server}'") + })? + .arguments + .unwrap_or_default(); + let (mut arguments, missing) = resolve_prompt_args(&declared, provided); + for key in missing { + let value = + Text::new(&format!("{key}:")).prompt().with_context(|| { + format!("Failed to read prompt argument '{key}'") + })?; + arguments.insert(key, value); + } + let result = ctx + .tool_scope + .mcp_runtime + .prompt(server, name, arguments) + .await?; + let flattened = flatten_prompt_messages(&result.messages); + let input = Input::from_str(ctx, &flattened, None)?; + ask(ctx, abort_signal.clone(), input, true).await?; + } + _ => println!("Usage: .prompt [key=value ...]"), + } + } + ".temp-role" => match args { Some(text) => { let app = Arc::clone(&ctx.app.config); ctx.use_prompt(app.as_ref(), text)?; } - None => println!("Usage: .prompt ..."), + None => println!("Usage: .temp-role ..."), }, ".role" => match args { Some(args) => match args.split_once(['\n', ' ']) { @@ -1207,13 +1249,16 @@ pub async fn run_repl_command( println!("Usage: .uninstall [--yes] (see `.uninstall --help`)") } }, - ".list" => match args { + ".list" => match args.map(str::trim) { + Some("prompts") => { + ctx.list_prompt_assets().await?; + } Some(args) => { - ctx.list_assets(args.trim())?; + ctx.list_assets(args)?; } _ => { println!( - "Usage: .list " + "Usage: .list " ) } }, @@ -1737,6 +1782,40 @@ fn split_first_arg(args: Option<&str>) -> Option<(&str, Option<&str>)> { }) } +fn parse_prompt_call_args(words: &[String]) -> Result> { + let mut args = HashMap::new(); + for word in words { + let Some((key, value)) = word.split_once('=') else { + bail!("Invalid prompt argument '{word}': arguments must be key=value pairs"); + }; + args.insert(key.to_string(), unquote_prompt_value(value).to_string()); + } + Ok(args) +} + +fn unquote_prompt_value(value: &str) -> &str { + let quoted = value.len() >= 2 + && ((value.starts_with('"') && value.ends_with('"')) + || (value.starts_with('\'') && value.ends_with('\''))); + if quoted { + &value[1..value.len() - 1] + } else { + value + } +} + +fn resolve_prompt_args( + declared: &[PromptArgument], + provided: HashMap, +) -> (HashMap, Vec) { + let missing = declared + .iter() + .filter(|arg| arg.required == Some(true) && !provided.contains_key(&arg.name)) + .map(|arg| arg.name.clone()) + .collect(); + (provided, missing) +} + pub fn split_args_text(line: &str, is_win: bool) -> (Vec, &str) { let mut words = Vec::new(); let mut word = String::new(); @@ -1888,8 +1967,52 @@ mod tests { } #[test] - fn repl_commands_has_61_entries() { - assert_eq!(REPL_COMMANDS.len(), 61); + fn repl_commands_has_62_entries() { + assert_eq!(REPL_COMMANDS.len(), 62); + } + + #[test] + fn parse_prompt_call_args_splits_on_first_equals_and_unquotes() { + let words = vec![ + "path=notes.txt".to_string(), + r#"style="a b""#.to_string(), + "expr=a=b".to_string(), + ]; + + let args = parse_prompt_call_args(&words).unwrap(); + + assert_eq!(args["path"], "notes.txt"); + assert_eq!(args["style"], "a b"); + assert_eq!(args["expr"], "a=b"); + } + + #[test] + fn parse_prompt_call_args_rejects_words_without_equals() { + let err = parse_prompt_call_args(&["positional".to_string()]) + .unwrap_err() + .to_string(); + + assert_eq!( + err, + "Invalid prompt argument 'positional': arguments must be key=value pairs" + ); + } + + #[test] + fn resolve_prompt_args_reports_missing_required_only() { + let declared = vec![ + PromptArgument::new("path").with_required(true), + PromptArgument::new("style"), + ]; + + let (resolved, missing) = resolve_prompt_args(&declared, HashMap::new()); + assert!(resolved.is_empty()); + assert_eq!(missing, vec!["path".to_string()]); + + let provided = HashMap::from([("path".to_string(), "notes.txt".to_string())]); + let (resolved, missing) = resolve_prompt_args(&declared, provided); + assert_eq!(resolved["path"], "notes.txt"); + assert!(missing.is_empty()); } #[test] @@ -2105,10 +2228,22 @@ mod tests { } #[test] - fn repl_commands_prompt_blocked_in_session_or_agent() { + fn repl_commands_prompt_always_available() { let cmd = REPL_COMMANDS.iter().find(|c| c.name == ".prompt").unwrap(); assert!(cmd.is_valid(StateFlags::empty())); assert!(cmd.is_valid(StateFlags::ROLE)); + assert!(cmd.is_valid(StateFlags::SESSION)); + assert!(cmd.is_valid(StateFlags::AGENT)); + } + + #[test] + fn repl_commands_temp_role_blocked_in_session_or_agent() { + let cmd = REPL_COMMANDS + .iter() + .find(|c| c.name == ".temp-role") + .unwrap(); + assert!(cmd.is_valid(StateFlags::empty())); + assert!(cmd.is_valid(StateFlags::ROLE)); assert!(!cmd.is_valid(StateFlags::SESSION)); assert!(!cmd.is_valid(StateFlags::AGENT)); }