diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-13 07:41:29 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-13 07:41:29 +0800 |
| commit | 255b194bcc538b2557caa60c5b607b4c4bfc0abd (patch) | |
| tree | 2a2995f7a33df2d3024f489de4435a06e94dcef8 /src | |
| parent | 64982b4510e38153885bfd0a78c250110b3e03c5 (diff) | |
| download | aichat-255b194bcc538b2557caa60c5b607b4c4bfc0abd.tar.gz | |
feat: add `.starter` repl command (#594)
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/bot.rs | 17 | ||||
| -rw-r--r-- | src/config/input.rs | 2 | ||||
| -rw-r--r-- | src/config/mod.rs | 18 | ||||
| -rw-r--r-- | src/repl/mod.rs | 50 |
4 files changed, 69 insertions, 18 deletions
diff --git a/src/config/bot.rs b/src/config/bot.rs index fef61f2..e1e1df6 100644 --- a/src/config/bot.rs +++ b/src/config/bot.rs @@ -54,10 +54,6 @@ impl Bot { } }; - let render_options = config.read().render_options()?; - let mut markdown_render = MarkdownRender::init(render_options)?; - println!("{}", markdown_render.render(&definition.banner())); - let rag = if rag_path.exists() { Some(Arc::new(Rag::load(config, "rag", &rag_path)?)) } else if embeddings_dir.is_dir() { @@ -101,6 +97,10 @@ impl Bot { Ok(data) } + pub fn banner(&self) -> String { + self.definition.banner() + } + pub fn name(&self) -> &str { &self.name } @@ -120,6 +120,10 @@ impl Bot { pub fn rag(&self) -> Option<Arc<Rag>> { self.rag.clone() } + + pub fn converstaion_staters(&self) -> &[String] { + &self.definition.conversation_starters + } } impl RoleLike for Bot { @@ -227,14 +231,13 @@ impl BotDefinition { format!( r#" -**Conversation Starters** +## Conversation Starters {starters}"# ) }; format!( r#"# {name} {version} -{description}{starters} -"# +{description}{starters}"# ) } } diff --git a/src/config/input.rs b/src/config/input.rs index 1686245..0c93c1f 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -129,7 +129,7 @@ impl Input { self.text = text; } - pub async fn maybe_embeddings(&mut self, abort_signal: AbortSignal) -> Result<()> { + pub async fn use_embeddings(&mut self, abort_signal: AbortSignal) -> Result<()> { if self.text.is_empty() { return Ok(()); } diff --git a/src/config/mod.rs b/src/config/mod.rs index 7e5d6cf..b1eb41f 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -918,7 +918,15 @@ impl Config { if let Some(bot) = &self.bot { bot.export() } else { - bail!("No rag") + bail!("No bot") + } + } + + pub fn bot_banner(&self) -> Result<String> { + if let Some(bot) = &self.bot { + Ok(bot.banner()) + } else { + bail!("No bot") } } @@ -1023,6 +1031,14 @@ impl Config { .collect(), ".rag" => self.list_rags().into_iter().map(|v| (v, None)).collect(), ".bot" => list_bots().into_iter().map(|v| (v, None)).collect(), + ".starter" => match &self.bot { + Some(bot) => bot + .converstaion_staters() + .iter() + .map(|v| (v.clone(), None)) + .collect(), + None => vec![], + }, ".set" => vec![ "max_output_tokens", "temperature", diff --git a/src/repl/mod.rs b/src/repl/mod.rs index dca5fc4..0b74e73 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -33,7 +33,7 @@ lazy_static! { const MENU_NAME: &str = "completion_menu"; lazy_static! { - static ref REPL_COMMANDS: [ReplCommand; 22] = [ + static ref REPL_COMMANDS: [ReplCommand; 23] = [ ReplCommand::new(".help", "Show this help message", AssertState::pass()), ReplCommand::new(".info", "View system info", AssertState::pass()), ReplCommand::new(".model", "Change the current LLM", AssertState::pass()), @@ -104,6 +104,11 @@ lazy_static! { AssertState::True(StateFlags::BOT), ), ReplCommand::new( + ".starter", + "Use converstaion starters", + AssertState::True(StateFlags::BOT) + ), + ReplCommand::new( ".exit bot", "Leave the bot", AssertState::True(StateFlags::BOT) @@ -233,7 +238,7 @@ impl Repl { Some((name, text)) => { let role = self.config.read().retrieve_role(name.trim())?; let input = Input::from_str(&self.config, text.trim(), Some(role)); - ask(&self.config, self.abort_signal.clone(), input).await?; + ask(&self.config, self.abort_signal.clone(), input, false).await?; } None => { self.config.write().use_role(args)?; @@ -253,6 +258,25 @@ impl Repl { } None => println!(r#"Usage: .bot <name>"#), }, + ".starter" => match args { + Some(value) => { + let input = Input::from_str(&self.config, value, None); + ask(&self.config, self.abort_signal.clone(), input, true).await?; + } + None => { + let banner = self.config.read().bot_banner()?; + let output = format!( + r#"Usage: .starter <text>... + +Tips: use <tab> to autocomplete conversation starter text. +--------------------------------------------------------- + +{banner}"# + ); + + println!("{output}"); + } + }, ".save" => { match args.map(|v| match v.split_once(' ') { Some((subcmd, args)) => (subcmd, args.trim()), @@ -284,7 +308,7 @@ impl Repl { let (files, text) = split_files_text(args); let files = shell_words::split(files).with_context(|| "Invalid args")?; let input = Input::new(&self.config, text, files, None)?; - ask(&self.config, self.abort_signal.clone(), input).await?; + ask(&self.config, self.abort_signal.clone(), input, true).await?; } None => println!("Usage: .file <files>... [-- <text>...]"), }, @@ -315,9 +339,8 @@ impl Repl { _ => unknown_command()?, }, None => { - let mut input = Input::from_str(&self.config, line, None); - input.maybe_embeddings(self.abort_signal.clone()).await?; - ask(&self.config, self.abort_signal.clone(), input).await?; + let input = Input::from_str(&self.config, line, None); + ask(&self.config, self.abort_signal.clone(), input, true).await?; } } @@ -455,16 +478,24 @@ impl Validator for ReplValidator { } #[async_recursion] -async fn ask(config: &GlobalConfig, abort: AbortSignal, mut input: Input) -> Result<()> { +async fn ask( + config: &GlobalConfig, + abort_signal: AbortSignal, + mut input: Input, + with_embeddings: bool, +) -> Result<()> { if input.is_empty() { return Ok(()); } + if with_embeddings { + input.use_embeddings(abort_signal.clone()).await?; + } while config.read().is_compressing_session() { std::thread::sleep(std::time::Duration::from_millis(100)); } let client = input.create_client()?; let (output, tool_call_results) = - send_stream(&input, client.as_ref(), config, abort.clone()).await?; + send_stream(&input, client.as_ref(), config, abort_signal.clone()).await?; config .write() @@ -493,8 +524,9 @@ async fn ask(config: &GlobalConfig, abort: AbortSignal, mut input: Input) -> Res if need_send_call_results(&tool_call_results) { ask( config, - abort, + abort_signal, input.merge_tool_call(output, tool_call_results), + false, ) .await } else { |
