From bb867c4fcbc6f42770471f0a0117cd8909f660d7 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 11 Jun 2024 11:00:12 +0800 Subject: feat: support bot (#579) * feat: support bots * refactor with RoleLike * improve exiting session * make bot works with rag * refactor repl assert state * add bot banner * repl complete bots according bots.txt * fix on windows * remove threadpool executing function callings * adjust repl left_prompt * move bot config to global config.yaml * `.bot` throw err if funciton callings is not configured --- src/rag/mod.rs | 19 +++++++++++-------- 1 file changed, 11 insertions(+), 8 deletions(-) (limited to 'src/rag/mod.rs') diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 5ab9de5..029c59e 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -20,7 +20,6 @@ use std::fmt::Debug; use std::{io::BufReader, path::Path}; use tokio::sync::mpsc; -pub const TEMP_RAG_NAME: &str = "temp"; pub const CHUNK_OVERLAP: usize = 20; pub const SIMILARITY_THRESHOLD: f32 = 0.25; @@ -48,7 +47,8 @@ impl Rag { pub async fn init( config: &GlobalConfig, name: &str, - path: &Path, + save_path: &Path, + doc_paths: &[String], abort_signal: AbortSignal, ) -> Result { debug!("init rag: {name}"); @@ -56,9 +56,12 @@ impl Rag { let chunk_size = model.default_chunk_size(); let chunk_size = set_chunk_size(chunk_size)?; let data = RagData::new(&model.id(), chunk_size); - let mut rag = Self::create(config, name, path, data)?; - let paths = add_document_paths()?; - debug!("document paths: {paths:?}"); + let mut rag = Self::create(config, name, save_path, data)?; + let mut paths = doc_paths.to_vec(); + if paths.is_empty() { + paths = add_doc_paths()?; + }; + debug!("doc paths: {paths:?}"); let (stop_spinner_tx, set_spinner_message_tx) = run_spinner("Starting").await; tokio::select! { ret = rag.add_paths(&paths, Some(set_spinner_message_tx)) => { @@ -71,8 +74,8 @@ impl Rag { }, }; if !rag.is_temp() { - rag.save(path)?; - println!("✨ Saved rag to '{}'", path.display()); + rag.save(save_path)?; + println!("✨ Saved rag to '{}'", save_path.display()); } Ok(rag) } @@ -408,7 +411,7 @@ fn set_chunk_size(chunk_size: usize) -> Result { value.parse().map_err(|_| anyhow!("Invalid chunk_size")) } -fn add_document_paths() -> Result> { +fn add_doc_paths() -> Result> { let text = Text::new("Add document paths:") .with_validator(required!("This field is required")) .with_help_message("e.g. file1;dir2/;dir3/**/*.md") -- cgit v1.2.3