diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/agent.rs | 33 | ||||
| -rw-r--r-- | src/config/mod.rs | 14 | ||||
| -rw-r--r-- | src/rag/mod.rs | 12 | ||||
| -rw-r--r-- | src/utils/mod.rs | 53 |
4 files changed, 82 insertions, 30 deletions
diff --git a/src/config/agent.rs b/src/config/agent.rs index aaf0b07..c7ff9ad 100644 --- a/src/config/agent.rs +++ b/src/config/agent.rs @@ -29,13 +29,13 @@ impl Agent { name: &str, abort_signal: AbortSignal, ) -> Result<Self> { - let definition_path = Config::agent_definition_file(name)?; - let functions_path = Config::agent_functions_file(name)?; + let functions_dir = Config::agent_functions_dir(name)?; + let definition_file_path = functions_dir.join("index.yaml"); + let functions_file_path = functions_dir.join("functions.json"); let rag_path = Config::agent_rag_file(name)?; - let embeddings_dir = Config::agent_embeddings_dir(name)?; - let definition = AgentDefinition::load(&definition_path)?; - let functions = if functions_path.exists() { - Functions::init(&functions_path)? + let definition = AgentDefinition::load(&definition_file_path)?; + let functions = if functions_file_path.exists() { + Functions::init(&functions_file_path)? } else { Functions::default() }; @@ -55,11 +55,20 @@ impl Agent { }; let rag = if rag_path.exists() { Some(Arc::new(Rag::load(config, "rag", &rag_path)?)) - } else if embeddings_dir.is_dir() { - println!("The agent uses an embeddings directory, initializing RAG..."); - let doc_path = embeddings_dir.display().to_string(); + } else if !definition.documents.is_empty() { + println!("The agent has the documents, initializing RAG..."); + let mut document_paths = vec![]; + for path in &definition.documents { + if Rag::is_url_path(path) { + document_paths.push(path.to_string()); + } else { + let new_path = safe_join_path(&functions_dir, path) + .ok_or_else(|| anyhow!("Invalid document path: '{path}'"))?; + document_paths.push(new_path.display().to_string()) + } + } Some(Arc::new( - Rag::init(config, "rag", &rag_path, &[doc_path], abort_signal).await?, + Rag::init(config, "rag", &rag_path, &document_paths, abort_signal).await?, )) } else { None @@ -197,6 +206,8 @@ pub struct AgentDefinition { pub instructions: String, #[serde(default)] pub conversation_starters: Vec<String>, + #[serde(default)] + pub documents: Vec<String>, } impl AgentDefinition { @@ -249,7 +260,7 @@ fn list_agents_impl() -> Result<Vec<String>> { .split('\n') .filter_map(|line| { let line = line.trim(); - if line.is_empty() { + if line.is_empty() || line.starts_with('#') { None } else { Some(line.to_string()) diff --git a/src/config/mod.rs b/src/config/mod.rs index 7ca1972..f82f36a 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -47,8 +47,6 @@ const FUNCTIONS_DIR_NAME: &str = "functions"; const FUNCTIONS_FILE_NAME: &str = "functions.json"; const FUNCTIONS_BIN_DIR_NAME: &str = "bin"; const AGENTS_DIR_NAME: &str = "agents"; -const AGENT_DEFINITION_FILE_NAME: &str = "index.yaml"; -const AGENT_EMBEDDINGS_DIR: &str = "embeddings"; const AGENT_RAG_FILE_NAME: &str = "rag.bin"; pub const TEMP_ROLE_NAME: &str = "%%"; @@ -352,18 +350,6 @@ impl Config { Ok(Self::agents_functions_dir()?.join(name)) } - pub fn agent_functions_file(name: &str) -> Result<PathBuf> { - Ok(Self::agent_functions_dir(name)?.join(FUNCTIONS_FILE_NAME)) - } - - pub fn agent_definition_file(name: &str) -> Result<PathBuf> { - Ok(Self::agent_functions_dir(name)?.join(AGENT_DEFINITION_FILE_NAME)) - } - - pub fn agent_embeddings_dir(name: &str) -> Result<PathBuf> { - Ok(Self::agent_functions_dir(name)?.join(AGENT_EMBEDDINGS_DIR)) - } - pub fn state(&self) -> StateFlags { let mut flags = StateFlags::empty(); if let Some(session) = &self.session { diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 3139b7e..f01fc0c 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -56,7 +56,7 @@ impl Rag { 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()?; + paths = add_document_paths()?; }; debug!("doc paths: {paths:?}"); let loaders = config.read().rag_document_loaders.clone(); @@ -233,7 +233,7 @@ impl Rag { progress(&spinner, "Gathering paths".into()); for path in paths { let path = path.as_ref(); - if path.starts_with("http://") || path.starts_with("https://") { + if Self::is_url_path(path) { if let Some(path) = path.strip_suffix("**") { new_paths.push((path.to_string(), RECURSIVE_URL_LOADER.into())); } else { @@ -337,6 +337,10 @@ impl Rag { Ok(()) } + pub fn is_url_path(path: &str) -> bool { + path.starts_with("http://") || path.starts_with("https://") + } + async fn hybird_search( &self, query: &str, @@ -628,8 +632,8 @@ fn set_chunk_overlay(default_value: usize) -> Result<usize> { value.parse().map_err(|_| anyhow!("Invalid chunk_overlay")) } -fn add_doc_paths() -> Result<Vec<String>> { - let text = Text::new("Add document paths:") +fn add_document_paths() -> Result<Vec<String>> { + let text = Text::new("Add documents:") .with_validator(required!("This field is required")) .with_help_message("e.g. file;dir/;dir/**/*.md;url;sites/**") .prompt()?; diff --git a/src/utils/mod.rs b/src/utils/mod.rs index f8bfcc6..7d1d07a 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -17,7 +17,10 @@ pub use self::spinner::{create_spinner, Spinner}; use fancy_regex::Regex; use is_terminal::IsTerminal; use lazy_static::lazy_static; -use std::env; +use std::{ + env, + path::{self, Path, PathBuf}, +}; lazy_static! { pub static ref CODE_BLOCK_RE: Regex = Regex::new(r"(?ms)```\w*(.*)```").unwrap(); @@ -151,6 +154,32 @@ pub fn dimmed_text(input: &str) -> String { nu_ansi_term::Style::new().dimmed().paint(input).to_string() } +pub fn safe_join_path<T1: AsRef<Path>, T2: AsRef<Path>>( + base_path: T1, + sub_path: T2, +) -> Option<PathBuf> { + let base_path = base_path.as_ref(); + let sub_path = sub_path.as_ref(); + if sub_path.is_absolute() { + return None; + } + + let mut joined_path = PathBuf::from(base_path); + + for component in sub_path.components() { + if path::Component::ParentDir == component { + return None; + } + joined_path.push(component); + } + + if joined_path.starts_with(base_path) { + Some(joined_path) + } else { + None + } +} + #[cfg(test)] mod tests { use super::*; @@ -161,4 +190,26 @@ mod tests { assert!(fuzzy_match("openai:gpt-4-turbo", "oai4")); assert!(!fuzzy_match("openai:gpt-4-turbo", "4gpt")); } + + #[test] + #[cfg(not(target_os = "windows"))] + fn test_safe_join_path() { + assert_eq!( + safe_join_path("/home/user/dir1", "files/file1"), + Some(PathBuf::from("/home/user/dir1/files/file1")) + ); + assert!(safe_join_path("/home/user/dir1", "/files/file1").is_none()); + assert!(safe_join_path("/home/user/dir1", "../file1").is_none()); + } + + #[test] + #[cfg(target_os = "windows")] + fn test_safe_join_path() { + assert_eq!( + safe_join_path("C:\\Users\\user\\dir1", "files/file1"), + Some(PathBuf::from("C:\\Users\\user\\dir1\\files\\file1")) + ); + assert!(safe_join_path("C:\\Users\\user\\dir1", "/files/file1").is_none()); + assert!(safe_join_path("C:\\Users\\user\\dir1", "../file1").is_none()); + } } |
