More slash commands (#5858)
Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
@@ -0,0 +1,315 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use anyhow::{anyhow, Result};
|
||||
|
||||
use crate::context_mgmt::compact_messages;
|
||||
use crate::conversation::message::{Message, SystemNotificationType};
|
||||
use crate::recipe::build_recipe::build_recipe_from_template_with_positional_params;
|
||||
use crate::session::SessionManager;
|
||||
|
||||
use super::Agent;
|
||||
|
||||
pub const COMPACT_TRIGGERS: &[&str] =
|
||||
&["/compact", "Please compact this conversation", "/summarize"];
|
||||
|
||||
pub struct CommandDef {
|
||||
pub name: &'static str,
|
||||
pub description: &'static str,
|
||||
}
|
||||
|
||||
static COMMANDS: &[CommandDef] = &[
|
||||
CommandDef {
|
||||
name: "prompts",
|
||||
description: "List available prompts, optionally filtered by extension",
|
||||
},
|
||||
CommandDef {
|
||||
name: "prompt",
|
||||
description: "Execute a prompt or show its info with --info",
|
||||
},
|
||||
CommandDef {
|
||||
name: "compact",
|
||||
description: "Compact the conversation history",
|
||||
},
|
||||
CommandDef {
|
||||
name: "clear",
|
||||
description: "Clear the conversation history",
|
||||
},
|
||||
];
|
||||
|
||||
pub fn list_commands() -> &'static [CommandDef] {
|
||||
COMMANDS
|
||||
}
|
||||
|
||||
impl Agent {
|
||||
pub async fn execute_command(&self, message_text: &str, session_id: &str) -> Option<Message> {
|
||||
let mut trimmed = message_text.trim().to_string();
|
||||
|
||||
if COMPACT_TRIGGERS.contains(&trimmed.as_str()) {
|
||||
trimmed = COMPACT_TRIGGERS[0].to_string();
|
||||
}
|
||||
|
||||
if !trimmed.starts_with('/') {
|
||||
return None;
|
||||
}
|
||||
|
||||
let command_str = trimmed.strip_prefix('/').unwrap_or(&trimmed);
|
||||
let (command, params) = command_str
|
||||
.split_once(' ')
|
||||
.map(|(cmd, p)| (cmd, p.trim()))
|
||||
.unwrap_or((command_str, ""));
|
||||
|
||||
let params: Vec<&str> = if params.is_empty() {
|
||||
vec![]
|
||||
} else {
|
||||
params.split_whitespace().collect()
|
||||
};
|
||||
|
||||
let result = match command {
|
||||
"prompts" => self.handle_prompts_command(¶ms, session_id).await,
|
||||
"prompt" => self.handle_prompt_command(¶ms, session_id).await,
|
||||
"compact" => self.handle_compact_command(session_id).await,
|
||||
"clear" => self.handle_clear_command(session_id).await,
|
||||
_ => {
|
||||
self.handle_recipe_command(command, ¶ms, session_id)
|
||||
.await
|
||||
}
|
||||
};
|
||||
|
||||
match result {
|
||||
Ok(msg) => msg,
|
||||
Err(e) => {
|
||||
Some(Message::assistant().with_text(format!("Error executing /{}: {}", command, e)))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_compact_command(&self, session_id: &str) -> Result<Option<Message>> {
|
||||
let session = SessionManager::get_session(session_id, true).await?;
|
||||
let conversation = session
|
||||
.conversation
|
||||
.ok_or_else(|| anyhow!("Session has no conversation"))?;
|
||||
|
||||
let (compacted_conversation, _usage) = compact_messages(
|
||||
self.provider().await?.as_ref(),
|
||||
&conversation,
|
||||
true, // is_manual_compact
|
||||
)
|
||||
.await?;
|
||||
|
||||
SessionManager::replace_conversation(session_id, &compacted_conversation).await?;
|
||||
|
||||
Ok(Some(Message::assistant().with_system_notification(
|
||||
SystemNotificationType::InlineMessage,
|
||||
"Compaction complete",
|
||||
)))
|
||||
}
|
||||
|
||||
async fn handle_clear_command(&self, session_id: &str) -> Result<Option<Message>> {
|
||||
use crate::conversation::Conversation;
|
||||
|
||||
SessionManager::replace_conversation(session_id, &Conversation::default()).await?;
|
||||
|
||||
SessionManager::update_session(session_id)
|
||||
.total_tokens(Some(0))
|
||||
.input_tokens(Some(0))
|
||||
.output_tokens(Some(0))
|
||||
.apply()
|
||||
.await?;
|
||||
|
||||
Ok(Some(Message::assistant().with_system_notification(
|
||||
SystemNotificationType::InlineMessage,
|
||||
"Conversation cleared",
|
||||
)))
|
||||
}
|
||||
|
||||
async fn handle_prompts_command(
|
||||
&self,
|
||||
params: &[&str],
|
||||
_session_id: &str,
|
||||
) -> Result<Option<Message>> {
|
||||
let extension_filter = params.first().map(|s| s.to_string());
|
||||
|
||||
let prompts = self.list_extension_prompts().await;
|
||||
|
||||
if let Some(filter) = &extension_filter {
|
||||
if !prompts.contains_key(filter) {
|
||||
let error_msg = format!("Extension '{}' not found", filter);
|
||||
return Ok(Some(Message::assistant().with_text(error_msg)));
|
||||
}
|
||||
}
|
||||
|
||||
let filtered_prompts: HashMap<String, Vec<String>> = prompts
|
||||
.into_iter()
|
||||
.filter(|(ext, _)| extension_filter.as_ref().is_none_or(|f| f == ext))
|
||||
.map(|(extension, prompt_list)| {
|
||||
let names = prompt_list.into_iter().map(|p| p.name).collect();
|
||||
(extension, names)
|
||||
})
|
||||
.collect();
|
||||
|
||||
let mut output = String::new();
|
||||
if filtered_prompts.is_empty() {
|
||||
output.push_str("No prompts available.\n");
|
||||
} else {
|
||||
output.push_str("Available prompts:\n\n");
|
||||
for (extension, prompt_names) in filtered_prompts {
|
||||
output.push_str(&format!("**{}**:\n", extension));
|
||||
for name in prompt_names {
|
||||
output.push_str(&format!(" - {}\n", name));
|
||||
}
|
||||
output.push('\n');
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Some(Message::assistant().with_text(output)))
|
||||
}
|
||||
|
||||
async fn handle_prompt_command(
|
||||
&self,
|
||||
params: &[&str],
|
||||
session_id: &str,
|
||||
) -> Result<Option<Message>> {
|
||||
if params.is_empty() {
|
||||
return Ok(Some(
|
||||
Message::assistant().with_text("Prompt name argument is required"),
|
||||
));
|
||||
}
|
||||
|
||||
let prompt_name = params[0].to_string();
|
||||
let is_info = params.get(1).map(|s| *s == "--info").unwrap_or(false);
|
||||
|
||||
if is_info {
|
||||
let prompts = self.list_extension_prompts().await;
|
||||
let mut prompt_info = None;
|
||||
|
||||
for (extension, prompt_list) in prompts {
|
||||
if let Some(prompt) = prompt_list.iter().find(|p| p.name == prompt_name) {
|
||||
let mut output = format!("**Prompt: {}**\n\n", prompt.name);
|
||||
if let Some(desc) = &prompt.description {
|
||||
output.push_str(&format!("Description: {}\n\n", desc));
|
||||
}
|
||||
output.push_str(&format!("Extension: {}\n\n", extension));
|
||||
|
||||
if let Some(args) = &prompt.arguments {
|
||||
output.push_str("Arguments:\n");
|
||||
for arg in args {
|
||||
output.push_str(&format!(" - {}", arg.name));
|
||||
if let Some(desc) = &arg.description {
|
||||
output.push_str(&format!(": {}", desc));
|
||||
}
|
||||
output.push('\n');
|
||||
}
|
||||
}
|
||||
|
||||
prompt_info = Some(output);
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
return Ok(Some(Message::assistant().with_text(
|
||||
prompt_info.unwrap_or_else(|| format!("Prompt '{}' not found", prompt_name)),
|
||||
)));
|
||||
}
|
||||
|
||||
let mut arguments = HashMap::new();
|
||||
for param in params.iter().skip(1) {
|
||||
if let Some((key, value)) = param.split_once('=') {
|
||||
let value = value.trim_matches('"');
|
||||
arguments.insert(key.to_string(), value.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
let arguments_value = serde_json::to_value(arguments)
|
||||
.map_err(|e| anyhow!("Failed to serialize arguments: {}", e))?;
|
||||
|
||||
match self.get_prompt(&prompt_name, arguments_value).await {
|
||||
Ok(prompt_result) => {
|
||||
for (i, prompt_message) in prompt_result.messages.into_iter().enumerate() {
|
||||
let msg = Message::from(prompt_message);
|
||||
|
||||
let expected_role = if i % 2 == 0 {
|
||||
rmcp::model::Role::User
|
||||
} else {
|
||||
rmcp::model::Role::Assistant
|
||||
};
|
||||
|
||||
if msg.role != expected_role {
|
||||
let error_msg = format!(
|
||||
"Expected {:?} message at position {}, but found {:?}",
|
||||
expected_role, i, msg.role
|
||||
);
|
||||
return Ok(Some(Message::assistant().with_text(error_msg)));
|
||||
}
|
||||
|
||||
SessionManager::add_message(session_id, &msg).await?;
|
||||
}
|
||||
|
||||
let last_message = SessionManager::get_session(session_id, true)
|
||||
.await?
|
||||
.conversation
|
||||
.ok_or_else(|| anyhow!("No conversation found"))?
|
||||
.messages()
|
||||
.last()
|
||||
.cloned()
|
||||
.ok_or_else(|| anyhow!("No messages in conversation"))?;
|
||||
|
||||
Ok(Some(last_message))
|
||||
}
|
||||
Err(e) => Ok(Some(
|
||||
Message::assistant().with_text(format!("Error getting prompt: {}", e)),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_recipe_command(
|
||||
&self,
|
||||
command: &str,
|
||||
params: &[&str],
|
||||
_session_id: &str,
|
||||
) -> Result<Option<Message>> {
|
||||
let full_command = format!("/{}", command);
|
||||
let recipe_path = match crate::slash_commands::get_recipe_for_command(&full_command) {
|
||||
Some(path) => path,
|
||||
None => return Ok(None),
|
||||
};
|
||||
|
||||
if !recipe_path.exists() {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let recipe_content = std::fs::read_to_string(&recipe_path)
|
||||
.map_err(|e| anyhow!("Failed to read recipe file: {}", e))?;
|
||||
|
||||
let recipe_dir = recipe_path
|
||||
.parent()
|
||||
.ok_or_else(|| anyhow!("Recipe path has no parent directory"))?;
|
||||
|
||||
let param_values: Vec<String> = params.iter().map(|s| s.to_string()).collect();
|
||||
|
||||
let recipe = match build_recipe_from_template_with_positional_params(
|
||||
recipe_content,
|
||||
recipe_dir,
|
||||
param_values,
|
||||
None::<fn(&str, &str) -> Result<String>>,
|
||||
) {
|
||||
Ok(recipe) => recipe,
|
||||
Err(crate::recipe::build_recipe::RecipeError::MissingParams { parameters }) => {
|
||||
return Ok(Some(Message::assistant().with_text(format!(
|
||||
"Recipe requires {} parameter(s): {}. Provided: {}",
|
||||
parameters.len(),
|
||||
parameters.join(", "),
|
||||
params.len()
|
||||
))));
|
||||
}
|
||||
Err(e) => return Err(anyhow!("Failed to build recipe: {}", e)),
|
||||
};
|
||||
|
||||
let prompt = [recipe.instructions.as_deref(), recipe.prompt.as_deref()]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n\n");
|
||||
|
||||
Ok(Some(Message::user().with_text(prompt)))
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user