From 9fa3d8cd134571b6a16fb1b328bed3c025742576 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 4 Nov 2024 19:44:11 +0800 Subject: refactor: save RAG document path as an absolute path (#965) --- src/rag/loader.rs | 25 +++++++++++++++++-------- src/rag/mod.rs | 1 - 2 files changed, 17 insertions(+), 9 deletions(-) (limited to 'src') diff --git a/src/rag/loader.rs b/src/rag/loader.rs index 2ea75a9..425187a 100644 --- a/src/rag/loader.rs +++ b/src/rag/loader.rs @@ -1,6 +1,7 @@ use super::*; use anyhow::{Context, Result}; +use path_absolutize::Absolutize; use std::collections::HashMap; pub const EXTENSION_METADATA: &str = "__extension__"; @@ -11,31 +12,40 @@ pub async fn load_document( path: &str, has_error: &mut bool, ) -> (String, Vec<(String, RagMetadata)>) { + let mut path = path.to_string(); let mut maybe_error = None; let mut files = vec![]; - if is_url(path) { + if is_url(&path) { if let Some(path) = path.strip_suffix("**") { match load_recursive_url(loaders, path).await { Ok(v) => files.extend(v), Err(err) => maybe_error = Some(err), } } else { - match load_url(loaders, path).await { + match load_url(loaders, &path).await { Ok(v) => files.push(v), Err(err) => maybe_error = Some(err), } } } else { - match load_path(loaders, path, has_error).await { - Ok(v) => files.extend(v), - Err(err) => maybe_error = Some(err), - } + match Path::new(&path).absolutize() { + Ok(v) => { + path = v.display().to_string(); + match load_path(loaders, &path, has_error).await { + Ok(v) => files.extend(v), + Err(err) => maybe_error = Some(err), + } + } + Err(_) => { + maybe_error = Some(anyhow!("Invalid path")); + } + }; } if let Some(err) = maybe_error { *has_error = true; println!("{}", warning_text(&format!("⚠️ {err:?}"))); } - (path.to_string(), files) + (path, files) } pub async fn load_recursive_url( @@ -71,7 +81,6 @@ pub async fn load_path( path: &str, has_error: &mut bool, ) -> Result> { - let path = Path::new(path).absolutize()?.display().to_string(); let file_paths = expand_glob_paths(&[path]).await?; let mut output = vec![]; let file_paths_len = file_paths.len(); diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 5bb4ea4..e490884 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -15,7 +15,6 @@ use hnsw_rs::prelude::*; use indexmap::{IndexMap, IndexSet}; use inquire::{required, validator::Validation, Confirm, Select, Text}; use parking_lot::RwLock; -use path_absolutize::Absolutize; use serde::{Deserialize, Serialize}; use serde_json::json; use std::{collections::HashMap, env, fmt::Debug, fs, hash::Hash, path::Path, time::Duration}; -- cgit v1.2.3