summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/config/agent.rs33
-rw-r--r--src/config/mod.rs14
-rw-r--r--src/rag/mod.rs12
-rw-r--r--src/utils/mod.rs53
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());
+ }
}