From d40913d7396e1dc5f3b6ddf47b61467b1e04430f Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 6 Mar 2023 05:30:20 +0800 Subject: feat: add `.prompt` command (#21) * feat: add `.prompt` command * set temp role name --- src/config.rs | 51 +++++++++++++++++++++++++++++++++++++-------------- src/repl.rs | 45 +++++++++++++++++++++++++++++++++++---------- 2 files changed, 72 insertions(+), 24 deletions(-) diff --git a/src/config.rs b/src/config.rs index 2f0a361..48b386b 100644 --- a/src/config.rs +++ b/src/config.rs @@ -18,6 +18,7 @@ const CONFIG_FILE_NAME: &str = "config.yaml"; const ROLES_FILE_NAME: &str = "roles.yaml"; const HISTORY_FILE_NAME: &str = "history.txt"; const MESSAGE_FILE_NAME: &str = "messages.md"; +const TEMP_ROLE: &str = "%TEMP%"; #[derive(Debug, Clone, Deserialize)] pub struct Config { @@ -119,20 +120,34 @@ impl Config { return; } if let Some(file) = file { - let role_name = self - .role - .as_ref() - .map(|v| format!("({})", v.name)) - .unwrap_or_default(); - let timestamp = format!("[{}]", now()); - let _ = file.write_all( - format!( - "# CHAT:{timestamp} {role_name}\n{}\n\n--------\n{}\n--------\n\n", - input.trim(), - output.trim(), - ) - .as_bytes(), - ); + let timestamp = now(); + let output = match self.role.as_ref() { + None => { + format!( + "# CHAT:[{timestamp}]\n{}\n\n--------\n{}\n--------\n\n", + input.trim(), + output.trim(), + ) + } + Some(v) => { + if v.name == TEMP_ROLE { + format!( + "# CHAT:[{timestamp}]\n{}\n{}\n\n--------\n{}\n--------\n\n", + v.prompt, + input.trim(), + output.trim(), + ) + } else { + format!( + "# CHAT:[{timestamp}] ({})\n{}\n\n--------\n{}\n--------\n\n", + v.name, + input.trim(), + output.trim(), + ) + } + } + }; + let _ = file.write_all(output.as_bytes()); } } @@ -163,6 +178,14 @@ impl Config { } } + pub fn create_temp_role(&mut self, prompt: &str) -> String { + self.role = Some(Role { + name: TEMP_ROLE.into(), + prompt: prompt.into(), + }); + "Done".into() + } + pub fn get_prompt(&self) -> Option { self.role.as_ref().and_then(|v| { if v.prompt.is_empty() { diff --git a/src/repl.rs b/src/repl.rs index 3c101b7..79b7bdb 100644 --- a/src/repl.rs +++ b/src/repl.rs @@ -17,9 +17,10 @@ use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; use std::thread::spawn; -const REPL_COMMANDS: [(&str, &str); 11] = [ +const REPL_COMMANDS: [(&str, &str); 12] = [ (".role", "Specifies the role the AI will play"), (".clear role", "Clear the currently selected role"), + (".prompt", "Add prompt, aka create a temporary role"), (".history", "Print the history"), (".clear history", "Clear the history"), (".multiline", "Enter multiline editor mode"), @@ -52,7 +53,9 @@ impl Repl { .with_edit_mode(edit_mode) .with_quick_completions(true) .with_partial_completions(true) - .with_validator(Box::new(ReplValidator)) + .with_validator(Box::new(ReplValidator { + multiline_cmds: [".multiline", ".prompt"].to_vec(), + })) .with_ansi_colors(true); let prompt = Self::create_prompt(); Ok(Self { editor, prompt }) @@ -138,11 +141,14 @@ impl Repl { handler.handle(ReplCmd::Info)?; } ".multiline" => { - let text = args.unwrap_or_default().to_string(); - if text.starts_with('{') && text.ends_with('}') { - handler.handle(ReplCmd::Submit(text))?; + let mut text = args.unwrap_or_default().to_string(); + if text.is_empty() { + dump("Usage: .multiline { }", 2); } else { - dump("Usage: .multiline { put your content here }", 2); + if text.starts_with('{') && text.ends_with('}') { + text = text[1..text.len() - 1].to_string() + } + handler.handle(ReplCmd::Submit(text))?; } } ".copy" => { @@ -157,6 +163,17 @@ impl Repl { ".set" => { handler.handle(ReplCmd::UpdateConfig(args.unwrap_or_default().to_string()))? } + ".prompt" => { + let mut text = args.unwrap_or_default().to_string(); + if text.is_empty() { + dump("Usage: .prompt { }.", 2); + } else { + if text.starts_with('{') && text.ends_with('}') { + text = text[1..text.len() - 1].to_string() + } + handler.handle(ReplCmd::Prompt(text))?; + } + } _ => dump_unknown_command(), } } else { @@ -219,11 +236,13 @@ impl Repl { )) } } -pub struct ReplValidator; +pub struct ReplValidator { + multiline_cmds: Vec<&'static str>, +} impl Validator for ReplValidator { fn validate(&self, line: &str) -> ValidationResult { - if line.split('"').count() % 2 == 0 || incomplete_brackets(line) { + if line.split('"').count() % 2 == 0 || incomplete_brackets(line, &self.multiline_cmds) { ValidationResult::Incomplete } else { ValidationResult::Complete @@ -231,9 +250,10 @@ impl Validator for ReplValidator { } } -fn incomplete_brackets(line: &str) -> bool { +fn incomplete_brackets(line: &str, multiline_cmds: &[&str]) -> bool { let mut balance: Vec = Vec::new(); - if !line.trim_start().starts_with(".multiline") { + let line = line.trim_start(); + if !multiline_cmds.iter().any(|v| line.starts_with(v)) { return false; } @@ -323,6 +343,10 @@ impl ReplCmdHandler { self.config.borrow_mut().role = None; dump("Done", 2); } + ReplCmd::Prompt(prompt) => { + let output = self.config.borrow_mut().create_temp_role(&prompt); + dump(output.trim(), 2); + } ReplCmd::Info => { let output = self.config.borrow().info()?; dump(output.trim(), 2); @@ -386,6 +410,7 @@ enum ReplCmd { Submit(String), SetRole(String), UpdateConfig(String), + Prompt(String), ClearRole, Info, } -- cgit v1.2.3