summaryrefslogtreecommitdiffstats
path: root/src/rag/loader.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-08-12 22:31:59 +0800
committerGitHub <noreply@github.com>2024-08-12 22:31:59 +0800
commitd79ad491067669fcf67991238e73a70aae413618 (patch)
tree738d8d944e78955fd7b7aa0fd4fc9f6c33f24ab2 /src/rag/loader.rs
parent92ce440b0bc8fd72b4c35f8500a8e10b488001b1 (diff)
downloadaichat-d79ad491067669fcf67991238e73a70aae413618.tar.gz
feat: support builtin website crawling (recursive_url) (#786)
Diffstat (limited to 'src/rag/loader.rs')
-rw-r--r--src/rag/loader.rs23
1 files changed, 11 insertions, 12 deletions
diff --git a/src/rag/loader.rs b/src/rag/loader.rs
index c764d78..f6c3b8e 100644
--- a/src/rag/loader.rs
+++ b/src/rag/loader.rs
@@ -11,15 +11,20 @@ pub async fn load_recursive_url(
path: &str,
) -> Result<Vec<(String, RagMetadata)>> {
let extension = RECURSIVE_URL_LOADER;
- let loader_command = loaders
- .get(extension)
- .with_context(|| format!("Document loader '{extension}' not configured"))?;
- let contents = run_loader_command(path, extension, loader_command)?;
- let pages: Vec<WebPage> = serde_json::from_str(&contents).context(r#"The crawler response is invalid. It should follow the JSON format: `[{"path":"...", "text":"..."}]`."#)?;
+ let pages: Vec<Page> = match loaders.get(extension) {
+ Some(loader_command) => {
+ let contents = run_loader_command(path, extension, loader_command)?;
+ serde_json::from_str(&contents).context(r#"The crawler response is invalid. It should follow the JSON format: `[{"path":"...", "text":"..."}]`."#)?
+ }
+ None => {
+ let options = CrawlOptions::preset(path);
+ crawl_website(path, options).await?
+ }
+ };
let output = pages
.into_iter()
.map(|v| {
- let WebPage { path, text } = v;
+ let Page { path, text } = v;
let mut metadata: RagMetadata = Default::default();
metadata.insert(PATH_METADATA.into(), path);
metadata.insert(EXTENSION_METADATA.into(), "md".into());
@@ -29,12 +34,6 @@ pub async fn load_recursive_url(
Ok(output)
}
-#[derive(Debug, Deserialize)]
-struct WebPage {
- path: String,
- text: String,
-}
-
pub async fn load_path(
loaders: &HashMap<String, String>,
path: &str,