diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/mod.rs | 16 | ||||
| -rw-r--r-- | src/rag/loader.rs | 23 | ||||
| -rw-r--r-- | src/rag/mod.rs | 4 | ||||
| -rw-r--r-- | src/utils/mod.rs | 11 | ||||
| -rw-r--r-- | src/utils/request.rs | 309 |
5 files changed, 340 insertions, 23 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs index 1c9ae9d..8db7085 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1740,16 +1740,12 @@ impl Config { } fn setup_document_loaders(&mut self) { - [ - ("pdf", "pdftotext $1 -"), - ("docx", "pandoc --to plain $1"), - (RECURSIVE_URL_LOADER, "rag-crawler $1 $2"), - ] - .into_iter() - .for_each(|(k, v)| { - let (k, v) = (k.to_string(), v.to_string()); - self.document_loaders.entry(k).or_insert(v); - }); + [("pdf", "pdftotext $1 -"), ("docx", "pandoc --to plain $1")] + .into_iter() + .for_each(|(k, v)| { + let (k, v) = (k.to_string(), v.to_string()); + self.document_loaders.entry(k).or_insert(v); + }); } } 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(); diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 6f89fbe..d9d6617 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -121,6 +121,17 @@ pub fn fuzzy_match(text: &str, pattern: &str) -> bool { pattern_index == pattern_chars.len() } +pub fn pretty_error(err: &anyhow::Error) -> String { + let mut output = vec![]; + output.push(err.to_string()); + output.push("Caused by:".to_string()); + for (i, cause) in err.chain().skip(1).enumerate() { + output.push(format!(" {i}: {cause}")); + } + output.push(String::new()); + output.join("\n") +} + pub fn error_text(input: &str) -> String { nu_ansi_term::Style::new() .fg(nu_ansi_term::Color::Red) diff --git a/src/utils/request.rs b/src/utils/request.rs index 4a4e7c9..873cbcc 100644 --- a/src/utils/request.rs +++ b/src/utils/request.rs @@ -1,15 +1,28 @@ use super::*; -use anyhow::{bail, Result}; +use anyhow::{anyhow, bail, Context, Result}; +use fancy_regex::Regex; +use futures_util::{stream, StreamExt}; use http::header::CONTENT_TYPE; +use reqwest::Url; +use scraper::{Html, Selector}; +use serde::Deserialize; +use serde_json::Value; use std::{collections::HashMap, time::Duration}; +use std::{collections::HashSet, sync::Arc}; use tokio::io::AsyncWriteExt; +use tokio::sync::Semaphore; pub const URL_LOADER: &str = "url"; pub const RECURSIVE_URL_LOADER: &str = "recursive_url"; + pub const MEDIA_URL_EXTENSION: &str = "media_url"; pub const DEFAULT_EXTENSION: &str = "txt"; +const MAX_CRAWLS: usize = 5; +const BREAK_ON_ERROR: bool = false; +const USER_AGENT: &str = "curl/8.6.0"; + lazy_static::lazy_static! { static ref CLIENT: Result<reqwest::Client> = { let builder = reqwest::ClientBuilder::new().timeout(Duration::from_secs(30)); @@ -17,6 +30,21 @@ lazy_static::lazy_static! { let client = builder.build()?; Ok(client) }; + + static ref PRESET: Vec<(Regex, CrawlOptions)> = vec![ + ( + Regex::new(r"github.com/([^/]+)/([^/]+)/tree/([^/]+)").unwrap(), + CrawlOptions { exclude: vec!["changelog".into(), "changes".into(), "license".into()], ..Default::default() } + ), + ( + Regex::new(r"github.com/([^/]+)/([^/]+)/wiki").unwrap(), + CrawlOptions { exclude: vec!["_history".into()], extract: Some("#wiki-body".into()), ..Default::default() } + + ) + ]; + + static ref EXTENSION_RE: Regex = Regex::new(r"\.[^.]+$").unwrap(); + static ref GITHUB_REPO_RE: Regex = Regex::new(r"^https://github\.com/([^/]+)/([^/]+)/tree/([^/]+)").unwrap(); } pub async fn fetch( @@ -120,3 +148,282 @@ pub async fn fetch( }; Ok(result) } + +#[derive(Debug, Clone, Default)] +pub struct CrawlOptions { + extract: Option<String>, + exclude: Vec<String>, + no_log: bool, +} + +impl CrawlOptions { + pub fn preset(start_url: &str) -> CrawlOptions { + for (re, options) in PRESET.iter() { + if let Ok(true) = re.is_match(start_url) { + return options.clone(); + } + } + CrawlOptions::default() + } +} + +pub async fn crawl_website(start_url: &str, options: CrawlOptions) -> Result<Vec<Page>> { + let start_url = Url::parse(start_url)?; + let mut paths = vec![start_url.path().to_string()]; + let normalized_start_url = normalize_start_url(&start_url); + if !options.no_log { + println!( + "Start crawling url={start_url} exclude={} extract={}", + options.exclude.join(","), + options.extract.as_deref().unwrap_or_default() + ); + } + + if let Ok(true) = GITHUB_REPO_RE.is_match(start_url.as_str()) { + paths = crawl_gh_tree(&start_url, &options.exclude) + .await + .with_context(|| "Failed to craw github repo".to_string())?; + } + + let semaphore = Arc::new(Semaphore::new(MAX_CRAWLS)); + let mut result_pages = Vec::new(); + + let mut index = 0; + while index < paths.len() { + let batch = paths[index..std::cmp::min(index + MAX_CRAWLS, paths.len())].to_vec(); + + let tasks: Vec<_> = batch + .iter() + .map(|path| { + let options = options.clone(); + let permit = semaphore.clone().acquire_owned(); // acquire a permit for concurrency control + let normalized_start_url = normalized_start_url.clone(); + let path = path.clone(); + + async move { + let _permit = permit.await?; + let url = normalized_start_url + .join(&path) + .map_err(|_| anyhow!("Invalid crawl page at {}", path))?; + let mut page = crawl_page(&normalized_start_url, &path, options) + .await + .with_context(|| format!("Failed to crawl page {}", url.as_str()))?; + page.0 = url.as_str().to_string(); + Ok(page) + } + }) + .collect(); + + let results = stream::iter(tasks) + .buffer_unordered(MAX_CRAWLS) + .collect::<Vec<_>>() + .await; + + let mut new_paths = Vec::new(); + + for res in results { + match res { + Ok((path, text, links)) => { + if !options.no_log { + println!("Crawled {path}"); + } + if !text.is_empty() { + result_pages.push(Page { path, text }); + } + for link in links { + if !paths.iter().any(|p| match_link(p, &link)) { + new_paths.push(link); + } + } + } + Err(err) => { + if BREAK_ON_ERROR { + return Err(err); + } else if !options.no_log { + println!("{}", error_text(&pretty_error(&err))); + } + } + } + } + paths.extend(new_paths); + + index += batch.len(); + } + + Ok(result_pages) +} + +#[derive(Debug, Deserialize)] +pub struct Page { + pub path: String, + pub text: String, +} + +async fn crawl_gh_tree(start_url: &Url, exclude: &[String]) -> Result<Vec<String>> { + let path_segs: Vec<&str> = start_url.path().split('/').collect(); + if path_segs.len() < 4 { + bail!("Invalid gh tree {}", start_url.as_str()); + } + let client = match *CLIENT { + Ok(ref client) => client, + Err(ref err) => bail!("{err}"), + }; + let owner = path_segs[1]; + let repo = path_segs[2]; + let branch = path_segs[4]; + let root_path = path_segs[5..].join("/"); + + let url = format!( + "https://api.github.com/repos/{}/{}/git/ref/heads/{}", + owner, repo, branch + ); + + let res_body: Value = client + .get(&url) + .header("User-Agent", USER_AGENT) + .header("Accept", "application/vnd.github+json") + .header("X-GitHub-Api-Version", "2022-11-28") + .send() + .await? + .json() + .await?; + + let sha = res_body["object"]["sha"] + .as_str() + .ok_or_else(|| anyhow!("Not found branch or tag"))?; + + let url = format!( + "https://api.github.com/repos/{}/{}/git/trees/{}?recursive=true", + owner, repo, sha + ); + + let res_body: Value = client + .get(&url) + .header("User-Agent", USER_AGENT) + .header("Accept", "application/vnd.github+json") + .header("X-GitHub-Api-Version", "2022-11-28") + .send() + .await? + .json() + .await?; + let tree = res_body["tree"] + .as_array() + .ok_or_else(|| anyhow!("Invalid github repo tree"))?; + let paths = tree + .iter() + .flat_map(|v| { + let path = v["path"].as_str()?; + if (path.ends_with(".md") || path.ends_with(".MD")) + && path.starts_with(&root_path) + && !should_exclude_link(path, exclude) + { + Some(format!( + "https://raw.githubusercontent.com/{owner}/{repo}/{branch}/{path}" + )) + } else { + None + } + }) + .collect(); + + Ok(paths) +} + +async fn crawl_page( + start_url: &Url, + path: &str, + options: CrawlOptions, +) -> Result<(String, String, Vec<String>)> { + let client = match *CLIENT { + Ok(ref client) => client, + Err(ref err) => bail!("{err}"), + }; + let location = start_url.join(path)?; + let response = client + .get(location.as_str()) + .header("User-Agent", USER_AGENT) + .send() + .await?; + let body = response.text().await?; + + if let Ok(true) = GITHUB_REPO_RE.is_match(start_url.as_str()) { + return Ok((path.to_string(), body, vec![])); + } + + let mut links = HashSet::new(); + let document = Html::parse_document(&body); + let selector = Selector::parse("a").map_err(|err| anyhow!("Invalid link selector, {}", err))?; + + for element in document.select(&selector) { + if let Some(href) = element.value().attr("href") { + let href = Url::parse(href).ok().or_else(|| location.join(href).ok()); + match href { + None => continue, + Some(href) => { + if href.as_str().starts_with(location.as_str()) + && !should_exclude_link(href.path(), &options.exclude) + { + links.insert(href.path().to_string()); + } + } + } + } + } + + let text = if let Some(selector) = &options.extract { + let selector = Selector::parse(selector) + .map_err(|err| anyhow!("Invalid extract selector, {}", err))?; + document + .select(&selector) + .map(|v| html_to_md(&v.html())) + .collect::<Vec<String>>() + .join("\n\n") + } else { + html_to_md(&body) + }; + + Ok((path.to_string(), text, links.into_iter().collect())) +} + +fn html_to_md(html: &str) -> String { + html2text::from_read(html.as_bytes(), usize::MAX) +} + +fn should_exclude_link(link: &str, exclude: &[String]) -> bool { + if link.contains("#") { + return true; + } + let parts: Vec<&str> = link.trim_end_matches('/').split('/').collect(); + let name = parts.last().unwrap_or(&"").to_lowercase(); + + for exclude_name in exclude { + let yes = match EXTENSION_RE.is_match(exclude_name) { + Ok(true) => exclude_name.to_lowercase() == name.to_lowercase(), + _ => exclude_name.to_lowercase() == EXTENSION_RE.replace(&name, "").to_lowercase(), + }; + if yes { + return true; + } + } + false +} + +fn normalize_start_url(start_url: &Url) -> Url { + let mut start_url = start_url.clone(); + start_url.set_query(None); + start_url.set_fragment(None); + let new_path = match start_url.path().rfind('/') { + Some(last_slash_index) => start_url.path()[..last_slash_index + 1].to_string(), + None => start_url.path().to_string(), + }; + start_url.set_path(&new_path); + start_url +} + +fn match_link(path: &str, link: &str) -> bool { + path == link + || path + == link + .trim_end_matches("/index.html") + .trim_end_matches("/index.htm") +} |
