diff options
| author | sigoden <sigoden@gmail.com> | 2024-08-12 22:31:59 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-08-12 22:31:59 +0800 |
| commit | d79ad491067669fcf67991238e73a70aae413618 (patch) | |
| tree | 738d8d944e78955fd7b7aa0fd4fc9f6c33f24ab2 /src/rag | |
| parent | 92ce440b0bc8fd72b4c35f8500a8e10b488001b1 (diff) | |
| download | aichat-d79ad491067669fcf67991238e73a70aae413618.tar.gz | |
feat: support builtin website crawling (recursive_url) (#786)
Diffstat (limited to 'src/rag')
| -rw-r--r-- | src/rag/loader.rs | 23 | ||||
| -rw-r--r-- | src/rag/mod.rs | 4 |
2 files changed, 15 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, diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 266a859..6dbb6be 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -371,6 +371,10 @@ impl Rag { self.data.add(next_file_id, files, document_ids, embeddings); self.data.document_paths = document_paths; + if self.data.files.is_empty() { + bail!("No RAG files"); + } + progress(&spinner, "Building store".into()); self.hnsw = self.data.build_hnsw(); self.bm25 = self.data.build_bm25(); |
