summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-27 07:51:30 +0800
committerGitHub <noreply@github.com>2024-06-27 07:51:30 +0800
commit1ced451c2723c36b7e0cb70e6c0755cdbad457c3 (patch)
treeea994b6bf4a6ba0f209129419a54e04679c88562 /src/config
parentf60df039979b8642aa82fec9e6ce56acb3a80f50 (diff)
downloadaichat-1ced451c2723c36b7e0cb70e6c0755cdbad457c3.tar.gz
refactor: agent rag use `documents` field other than `embeddings` dir (#658)
Diffstat (limited to 'src/config')
-rw-r--r--src/config/agent.rs33
-rw-r--r--src/config/mod.rs14
2 files changed, 22 insertions, 25 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 {