From f9847475b8eae99a70f16dc63bac67da2632474a Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 13 Jun 2024 19:41:54 +0800 Subject: feat: add rag and bot related cli options (#595) --- src/cli.rs | 16 ++++++++-- src/config/bot.rs | 20 ++++-------- src/config/mod.rs | 24 +++++++++++--- src/config/session.rs | 3 -- src/main.rs | 80 ++++++++++++++++++++++++++++++----------------- src/rag/mod.rs | 3 ++ src/repl/mod.rs | 12 +++---- src/utils/abort_signal.rs | 4 +-- 8 files changed, 103 insertions(+), 59 deletions(-) (limited to 'src') diff --git a/src/cli.rs b/src/cli.rs index c84db2a..651d23a 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -18,6 +18,12 @@ pub struct Cli { /// Forces the session to be saved #[clap(long)] pub save_session: bool, + /// Start a bot + #[clap(short = 'b', long)] + pub bot: Option, + /// Start a RAG + #[clap(long)] + pub rag: Option, /// Serve the LLM API and WebAPP #[clap(long, value_name = "ADDRESS")] pub serve: Option>, @@ -51,12 +57,18 @@ pub struct Cli { /// List all available models #[clap(long)] pub list_models: bool, - /// List all available roles + /// List all roles #[clap(long)] pub list_roles: bool, - /// List all available sessions + /// List all sessions #[clap(long)] pub list_sessions: bool, + /// List all bots + #[clap(long)] + pub list_bots: bool, + /// List all RAGs + #[clap(long)] + pub list_rags: bool, /// Input text #[clap(trailing_var_arg = true)] text: Vec, diff --git a/src/config/bot.rs b/src/config/bot.rs index e1e1df6..e02cce7 100644 --- a/src/config/bot.rs +++ b/src/config/bot.rs @@ -53,22 +53,14 @@ impl Bot { None => config.current_model().clone(), } }; - let rag = if rag_path.exists() { Some(Arc::new(Rag::load(config, "rag", &rag_path)?)) } else if embeddings_dir.is_dir() { - println!("The bot has an embeddings directory, RAG is initializing..."); - let ans = Confirm::new("The bot attached embeddings, init RAG?") - .with_default(true) - .prompt()?; - if ans { - let doc_path = embeddings_dir.display().to_string(); - Some(Arc::new( - Rag::init(config, "rag", &rag_path, &[doc_path], abort_signal).await?, - )) - } else { - None - } + println!("The bot uses an embeddings directory, initializing RAG..."); + let doc_path = embeddings_dir.display().to_string(); + Some(Arc::new( + Rag::init(config, "rag", &rag_path, &[doc_path], abort_signal).await?, + )) } else { None }; @@ -121,7 +113,7 @@ impl Bot { self.rag.clone() } - pub fn converstaion_staters(&self) -> &[String] { + pub fn conversation_staters(&self) -> &[String] { &self.definition.conversation_starters } } diff --git a/src/config/mod.rs b/src/config/mod.rs index b1eb41f..571bf2c 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -442,7 +442,20 @@ impl Config { } pub fn info(&self) -> Result { - if let Some(session) = &self.session { + if let Some(bot) = &self.bot { + let output = bot.export()?; + if let Some(session) = &self.session { + let session = session + .export()? + .split('\n') + .map(|v| format!(" {v}")) + .collect::>() + .join("\n"); + Ok(format!("{output}session:\n{session}")) + } else { + Ok(output) + } + } else if let Some(session) = &self.session { session.export() } else if let Some(role) = &self.role { role.export() @@ -896,6 +909,7 @@ impl Config { pub async fn use_bot( config: &GlobalConfig, name: &str, + session: Option<&str>, abort_signal: AbortSignal, ) -> Result<()> { if !config.read().function_calling { @@ -904,11 +918,13 @@ impl Config { if config.read().bot.is_some() { bail!("Already in a bot, please run '.exit bot' first to exit the current bot."); } - let prelude = config.read().bot_prelude.clone(); let bot = Bot::init(config, name, abort_signal).await?; config.write().rag = bot.rag(); config.write().bot = Some(bot); - if let Some(session) = prelude { + let session = session + .map(|v| v.to_string()) + .or_else(|| config.read().bot_prelude.clone()); + if let Some(session) = session { config.write().use_session(Some(&session))?; } Ok(()) @@ -1033,7 +1049,7 @@ impl Config { ".bot" => list_bots().into_iter().map(|v| (v, None)).collect(), ".starter" => match &self.bot { Some(bot) => bot - .converstaion_staters() + .conversation_staters() .iter() .map(|v| (v.clone(), None)) .collect(), diff --git a/src/config/session.rs b/src/config/session.rs index 2aff467..38855cc 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -126,9 +126,6 @@ impl Session { } pub fn export(&self) -> Result { - if self.path.is_none() { - bail!("Not found session '{}'", self.name) - } let mut data = json!({ "path": self.path, "model": self.model().id(), diff --git a/src/main.rs b/src/main.rs index 72be91c..3f197eb 100644 --- a/src/main.rs +++ b/src/main.rs @@ -16,15 +16,13 @@ extern crate log; use crate::cli::Cli; use crate::client::{list_chat_models, send_stream, ChatCompletionsOutput}; use crate::config::{ - Config, GlobalConfig, Input, WorkingMode, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE, + list_bots, Config, GlobalConfig, Input, WorkingMode, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE, + TEMP_SESSION_NAME, }; use crate::function::{eval_tool_calls, need_send_call_results}; use crate::render::{render_error, MarkdownRender}; use crate::repl::Repl; -use crate::utils::{ - create_abort_signal, detect_shell, extract_block, run_command, run_spinner, Shell, - CODE_BLOCK_RE, IS_STDOUT_TERMINAL, -}; +use crate::utils::*; use anyhow::{bail, Result}; use async_recursion::async_recursion; @@ -53,9 +51,17 @@ async fn main() -> Result<()> { crate::logger::setup_logger(working_mode)?; let config = Arc::new(RwLock::new(Config::init(working_mode)?)); + let abort_signal = create_abort_signal(); + if let Some(addr) = cli.serve { return serve::run(config, addr).await; } + if cli.list_models { + for model in list_chat_models(&config.read()) { + println!("{}", model.id()); + } + return Ok(()); + } if cli.list_roles { config .read() @@ -64,15 +70,14 @@ async fn main() -> Result<()> { .for_each(|v| println!("{}", v.name())); return Ok(()); } - if cli.list_models { - for model in list_chat_models(&config.read()) { - println!("{}", model.id()); - } + if cli.list_bots { + let bots = list_bots().join("\n"); + println!("{bots}"); return Ok(()); } - if cli.list_sessions { - let sessions = config.read().list_sessions().join("\n"); - println!("{sessions}"); + if cli.list_rags { + let rags = config.read().list_rags().join("\n"); + println!("{rags}"); return Ok(()); } if let Some(wrap) = &cli.wrap { @@ -84,19 +89,36 @@ async fn main() -> Result<()> { if cli.dry_run { config.write().dry_run = true; } - if let Some(prompt) = &cli.prompt { - config.write().use_prompt(prompt)?; - } else if let Some(name) = &cli.role { - config.write().use_role(name)?; - } else if cli.execute { - config.write().use_role(SHELL_ROLE)?; - } else if cli.code { - config.write().use_role(CODE_ROLE)?; + + if let Some(bot) = &cli.bot { + let session = cli.session.as_ref().map(|v| match v { + Some(v) => v.as_str(), + None => TEMP_SESSION_NAME, + }); + Config::use_bot(&config, bot, session, abort_signal.clone()).await? + } else { + if let Some(prompt) = &cli.prompt { + config.write().use_prompt(prompt)?; + } else if let Some(name) = &cli.role { + config.write().use_role(name)?; + } else if cli.execute { + config.write().use_role(SHELL_ROLE)?; + } else if cli.code { + config.write().use_role(CODE_ROLE)?; + } + if let Some(session) = &cli.session { + config + .write() + .use_session(session.as_ref().map(|v| v.as_str()))?; + } + if let Some(rag) = &cli.rag { + Config::use_rag(&config, Some(rag), abort_signal.clone()).await?; + } } - if let Some(session) = &cli.session { - config - .write() - .use_session(session.as_ref().map(|v| v.as_str()))?; + if cli.list_sessions { + let sessions = config.read().list_sessions().join("\n"); + println!("{sessions}"); + return Ok(()); } if let Some(model_id) = &cli.model { config.write().set_model(model_id)?; @@ -124,8 +146,9 @@ async fn main() -> Result<()> { config.write().apply_prelude()?; if let Err(err) = match no_input { false => { - let input = create_input(&config, text, file)?; - start_directive(&config, input, cli.no_stream, cli.code).await + let mut input = create_input(&config, text, file)?; + input.use_embeddings(abort_signal.clone()).await?; + start_directive(&config, input, cli.no_stream, cli.code, abort_signal).await } true => start_interactive(&config).await, } { @@ -142,6 +165,7 @@ async fn start_directive( mut input: Input, no_stream: bool, code_mode: bool, + abort_signal: AbortSignal, ) -> Result<()> { let client = input.create_client()?; let extract_code = !*IS_STDOUT_TERMINAL && code_mode; @@ -167,8 +191,7 @@ async fn start_directive( (text, vec![]) } } else { - let abort = create_abort_signal(); - send_stream(&input, client.as_ref(), config, abort).await? + send_stream(&input, client.as_ref(), config, abort_signal.clone()).await? }; config .write() @@ -180,6 +203,7 @@ async fn start_directive( input.merge_tool_call(output, tool_call_results), no_stream, code_mode, + abort_signal, ) .await } else { diff --git a/src/rag/mod.rs b/src/rag/mod.rs index f88e80e..c98bab7 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -50,6 +50,9 @@ impl Rag { doc_paths: &[String], abort_signal: AbortSignal, ) -> Result { + if !*IS_STDOUT_TERMINAL { + bail!("An interactive shell is required to initialize rag.") + } debug!("init rag: {name}"); let model = select_embedding_model(config)?; let chunk_size = set_chunk_size(&model)?; diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 0b74e73..d56b790 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -79,7 +79,7 @@ lazy_static! { ), ReplCommand::new( ".exit session", - "End the current session", + "End the session", AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION) ), ReplCommand::new( @@ -105,7 +105,7 @@ lazy_static! { ), ReplCommand::new( ".starter", - "Use converstaion starters", + "Use the conversation starter", AssertState::True(StateFlags::BOT) ), ReplCommand::new( @@ -138,14 +138,13 @@ impl Repl { let editor = Self::create_editor(config)?; let prompt = ReplPrompt::new(config); - - let abort = create_abort_signal(); + let abort_signal = create_abort_signal(); Ok(Self { config: config.clone(), editor, prompt, - abort_signal: abort, + abort_signal, }) } @@ -254,7 +253,8 @@ impl Repl { } ".bot" => match args { Some(name) => { - Config::use_bot(&self.config, name, self.abort_signal.clone()).await?; + Config::use_bot(&self.config, name, None, self.abort_signal.clone()) + .await?; } None => println!(r#"Usage: .bot "#), }, diff --git a/src/utils/abort_signal.rs b/src/utils/abort_signal.rs index ac93653..3a8007d 100644 --- a/src/utils/abort_signal.rs +++ b/src/utils/abort_signal.rs @@ -54,9 +54,9 @@ impl AbortSignalInner { } } -pub async fn watch_abort_signal(abort: AbortSignal) { +pub async fn watch_abort_signal(abort_signal: AbortSignal) { loop { - if abort.aborted() { + if abort_signal.aborted() { break; } tokio::time::sleep(std::time::Duration::from_millis(100)).await; -- cgit v1.2.3