diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-11 11:00:12 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-11 11:00:12 +0800 |
| commit | bb867c4fcbc6f42770471f0a0117cd8909f660d7 (patch) | |
| tree | 6293f7f1108309160d1951f53f6429e9b004870d /src/rag/mod.rs | |
| parent | 5635ca6a58fb4a590419335b098b7317285bfb82 (diff) | |
| download | aichat-bb867c4fcbc6f42770471f0a0117cd8909f660d7.tar.gz | |
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
Diffstat (limited to 'src/rag/mod.rs')
| -rw-r--r-- | src/rag/mod.rs | 19 |
1 files changed, 11 insertions, 8 deletions
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<Self> { 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<usize> { value.parse().map_err(|_| anyhow!("Invalid chunk_size")) } -fn add_document_paths() -> Result<Vec<String>> { +fn add_doc_paths() -> Result<Vec<String>> { let text = Text::new("Add document paths:") .with_validator(required!("This field is required")) .with_help_message("e.g. file1;dir2/;dir3/**/*.md") |
