summaryrefslogtreecommitdiffstats
path: root/src/rag
diff options
context:
space:
mode:
Diffstat (limited to 'src/rag')
-rw-r--r--src/rag/loader.rs23
-rw-r--r--src/rag/mod.rs4
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();