summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-13 07:41:29 +0800
committerGitHub <noreply@github.com>2024-06-13 07:41:29 +0800
commit255b194bcc538b2557caa60c5b607b4c4bfc0abd (patch)
tree2a2995f7a33df2d3024f489de4435a06e94dcef8
parent64982b4510e38153885bfd0a78c250110b3e03c5 (diff)
downloadaichat-255b194bcc538b2557caa60c5b607b4c4bfc0abd.tar.gz
feat: add `.starter` repl command (#594)
-rw-r--r--src/config/bot.rs17
-rw-r--r--src/config/input.rs2
-rw-r--r--src/config/mod.rs18
-rw-r--r--src/repl/mod.rs50
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 {