use crate::engine::{collect_manifest, run_once_parallel, run_tasks_parallel};
use crate::{Environment, Mode, Website};
use std::collections::HashSet;
use std::env;
use std::net::{TcpListener, TcpStream};
use std::sync::mpsc::Sender;
use std::sync::{Arc, Mutex};
use std::thread::JoinHandle;
use std::time::Duration;
use camino::{Utf8Path, Utf8PathBuf};
use glob::Pattern;
use notify::RecursiveMode;
use notify_debouncer_full::{DebounceEventResult, new_debouncer};
use petgraph::visit::IntoNodeReferences;
use tungstenite::WebSocket;
pub fn watch<G: Send + Sync>(
site: &mut Website<G>,
data: G,
copied: Vec<(String, String)>,
out_dir: &Utf8Path,
cache_dir: &Utf8Path,
) -> anyhow::Result<()> {
let (tcp, port) = reserve_port()?;
let pwd = env::current_dir()?;
let globals = Environment {
generator: "hauchiwa",
mode: Mode::Watch,
port: Some(port),
data,
};
let prev_meta = crate::snapshot::SnapshotMeta::load(cache_dir)
.ok()
.flatten();
let mut static_files = crate::utils::collect_static(&copied, out_dir)?;
let (mut cache, mut snapshot, _) = run_once_parallel(site, &globals)?;
for entry in &static_files {
snapshot.insert_static_file(entry.dist_rel.clone(), entry.source_utf8.clone())?;
}
crate::utils::copy_static_entries(&static_files, &site.progress.copy)?;
tracing::info!("collected {} pages", snapshot.page_count());
match prev_meta {
Some(ref prev) => snapshot.commit_diff_meta(prev, out_dir)?,
None => snapshot.commit(out_dir)?,
}
snapshot.to_meta().save(cache_dir)?;
let mut prev_snapshot = snapshot;
tracing::info!("initial build completed, now watching for changes...");
let clients = Arc::new(Mutex::new(vec![]));
let _thread_i = new_thread_ws_incoming(tcp, clients.clone());
let (tx_reload, _thread_o) = new_thread_ws_reload(clients.clone());
let (tx, rx) = std::sync::mpsc::channel::<DebounceEventResult>();
let mut debouncer = new_debouncer(Duration::from_millis(250), None, move |result| {
tx.send(result).ok();
})?;
let mut watched = HashSet::new();
let mut task_filters = HashSet::new();
let mut static_filters = HashSet::new();
for (_, task) in site.graph.node_references() {
for path in &task.watched() {
if let Ok((path, pattern)) = resolve_watch_path(path) {
watched.insert(path);
task_filters.insert(pattern);
} else {
tracing::error!("failed to resolve path: {}", &path);
};
}
}
for (_, source) in &copied {
match resolve_static_watch_path(source) {
Ok((path, patterns)) => {
watched.insert(path);
static_filters.extend(patterns);
}
Err(e) => tracing::error!("failed to resolve static path `{}`: {}", source, e),
}
}
let watched = collapse_watch_paths(watched);
for path in watched {
tracing::info!("watching {}", path);
debouncer.watch(path, RecursiveMode::Recursive)?;
}
#[cfg(feature = "server")]
let _thread_http = super::http::start(out_dir.to_string());
loop {
match rx.recv() {
Ok(Ok(events)) => {
tracing::debug!("{:?} events received", events);
let mut dirty_nodes = HashSet::new();
let mut static_dirty = false;
for de in events {
for path in &de.event.paths {
let task_match =
task_filters.iter().any(|filter| filter.matches_path(path));
let static_match = static_filters
.iter()
.any(|filter| filter.matches_path(path));
if !task_match && !static_match {
continue;
}
static_dirty |= static_match;
if task_match && let Some(path) = Utf8Path::from_path(path) {
let Ok(path) = path.strip_prefix(&pwd) else {
continue;
};
for index in site.graph.node_indices() {
let task = &site.graph[index];
if task.is_dirty(path) {
dirty_nodes.insert(index);
}
}
}
}
}
if !dirty_nodes.is_empty() || static_dirty {
tracing::info!("change detected, re-running tasks...");
let mut to_rerun = HashSet::new();
for start_node in &dirty_nodes {
let mut dfs = petgraph::visit::Dfs::new(&site.graph, *start_node);
while let Some(nx) = dfs.next(&site.graph) {
to_rerun.insert(nx);
}
}
if !to_rerun.is_empty() {
let _diagnostics = match run_tasks_parallel(
site,
&globals,
&mut cache,
&to_rerun,
&dirty_nodes,
) {
Ok(res) => res,
Err(e) => {
tracing::error!("Error running tasks: {}", e);
continue;
}
};
}
if static_dirty {
static_files = match crate::utils::collect_static(&copied, out_dir) {
Ok(files) => files,
Err(e) => {
tracing::error!("failed to collect static files: {}", e);
continue;
}
};
}
let mut snapshot = match collect_manifest(&cache, &site.graph) {
Ok(snapshot) => snapshot,
Err(e) => {
tracing::error!("failed to collect output manifest: {}", e);
continue;
}
};
let mut static_manifest_ok = true;
for entry in &static_files {
if let Err(e) = snapshot
.insert_static_file(entry.dist_rel.clone(), entry.source_utf8.clone())
{
tracing::error!("failed to add static file to manifest: {}", e);
static_manifest_ok = false;
break;
}
}
if !static_manifest_ok {
continue;
}
if let Err(e) = crate::utils::copy_static_entries(&static_files, &site.progress.copy) {
tracing::error!("failed to copy static files: {}", e);
continue;
}
tracing::info!("collected {} pages", snapshot.page_count());
if let Err(e) = snapshot.commit_diff(&prev_snapshot, out_dir) {
tracing::error!("failed to write pages to dist: {}", e);
continue;
}
if let Err(e) = snapshot.to_meta().save(cache_dir) {
tracing::warn!("failed to save snapshot meta: {}", e);
}
prev_snapshot = snapshot;
tx_reload.send(()).ok();
tracing::info!("rebuild complete, watching for changes...");
}
}
Ok(Err(e)) => tracing::error!("watch error: {:?}", e),
Err(e) => tracing::error!("watch error: {:?}", e),
}
}
}
fn reserve_port() -> std::io::Result<(TcpListener, u16)> {
let listener = match TcpListener::bind("127.0.0.1:1337") {
Ok(sock) => sock,
Err(_) => TcpListener::bind("127.0.0.1:0")?,
};
let addr = listener.local_addr()?;
let port = addr.port();
Ok((listener, port))
}
fn new_thread_ws_incoming(
server: TcpListener,
client: Arc<Mutex<Vec<WebSocket<TcpStream>>>>,
) -> JoinHandle<()> {
std::thread::spawn(move || {
for stream in server.incoming() {
let stream = match stream {
Ok(s) => s,
Err(e) => {
tracing::warn!("WebSocket: incoming stream error: {e}");
continue;
}
};
let socket = match tungstenite::accept(stream) {
Ok(s) => s,
Err(e) => {
tracing::warn!("WebSocket: handshake failed: {e}");
continue;
}
};
#[allow(clippy::unwrap_used)] client.lock().unwrap().push(socket);
}
})
}
fn new_thread_ws_reload(
client: Arc<Mutex<Vec<WebSocket<TcpStream>>>>,
) -> (Sender<()>, JoinHandle<()>) {
let (tx, rx) = std::sync::mpsc::channel();
let thread = std::thread::spawn(move || {
while rx.recv().is_ok() {
#[allow(clippy::unwrap_used)] let mut clients = client.lock().unwrap();
let mut broken = vec![];
for (i, socket) in clients.iter_mut().enumerate() {
match socket.send("reload".into()) {
Ok(_) => {}
Err(tungstenite::error::Error::Io(e)) => {
if e.kind() == std::io::ErrorKind::BrokenPipe {
broken.push(i);
}
}
Err(e) => {
tracing::error!("Error: {e:?}");
}
}
}
for i in broken.into_iter().rev() {
clients.remove(i);
}
let len = clients.len();
if len > 10 {
for mut socket in clients.drain(0..len - 10) {
socket.close(None).ok();
}
}
}
});
(tx, thread)
}
pub fn resolve_watch_path(glob_str: impl AsRef<str>) -> anyhow::Result<(Utf8PathBuf, Pattern)> {
let path = Utf8Path::new(glob_str.as_ref());
let components: Vec<_> = path.components().collect();
let split_idx = components
.iter()
.position(|c| c.as_str().contains(['*', '?', '[']))
.unwrap_or(components.len());
let root_part: Utf8PathBuf = components.iter().take(split_idx).collect();
let suffix_part: Utf8PathBuf = components.iter().skip(split_idx).collect();
let absolute_root = root_part.canonicalize_utf8()?;
let (watch_root, match_pattern_str) =
if suffix_part.as_str().is_empty() && absolute_root.is_file() {
let parent = absolute_root
.parent()
.unwrap_or(&absolute_root)
.to_path_buf();
(parent, absolute_root)
} else {
let pattern_str = absolute_root.join(&suffix_part);
(absolute_root, pattern_str)
};
let pattern = Pattern::new(watch_root.join(match_pattern_str).as_str())?;
Ok((watch_root, pattern))
}
fn resolve_static_watch_path(
source: impl AsRef<str>,
) -> anyhow::Result<(Utf8PathBuf, Vec<Pattern>)> {
let absolute = Utf8Path::new(source.as_ref()).canonicalize_utf8()?;
if absolute.is_file() {
let watch_root = absolute.parent().unwrap_or(&absolute).to_path_buf();
let pattern = Pattern::new(absolute.as_str())?;
return Ok((watch_root, vec![pattern]));
}
let patterns = vec![
Pattern::new(absolute.as_str())?,
Pattern::new(absolute.join("**").as_str())?,
];
Ok((absolute, patterns))
}
fn collapse_watch_paths(paths: HashSet<Utf8PathBuf>) -> Vec<Utf8PathBuf> {
let mut paths: Vec<_> = paths.into_iter().collect();
paths.sort();
let mut filtered = Vec::new();
for path in paths {
if let Some(last) = filtered.last()
&& path.starts_with(last)
{
continue;
}
filtered.push(path);
}
filtered
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used)]
mod tests {
use super::*;
#[test]
fn test_concrete_file() {
let (watch, pattern) = resolve_watch_path("README.md").expect("Should resolve");
let cwd = Utf8PathBuf::try_from(std::env::current_dir().unwrap()).unwrap();
assert_eq!(watch.as_str(), cwd);
assert_eq!(pattern.as_str(), cwd.join("README.md"));
}
#[test]
fn test_concrete_directory() {
let (watch, pattern) = resolve_watch_path("src").expect("Should resolve");
let cwd = Utf8PathBuf::try_from(std::env::current_dir().unwrap()).unwrap();
assert_eq!(watch.as_str(), cwd.join("src"));
assert_eq!(pattern.as_str(), cwd.join("src"));
}
#[test]
fn test_directory_wildcard() {
let (watch, pattern) = resolve_watch_path("src/**/*.rs").expect("Should resolve");
let cwd = Utf8PathBuf::try_from(std::env::current_dir().unwrap()).unwrap();
assert_eq!(watch.as_str(), cwd.join("src/"));
assert_eq!(pattern.as_str(), cwd.join("src/**/*.rs"));
}
#[test]
fn test_static_directory_watch_path() {
let (watch, patterns) = resolve_static_watch_path("src").expect("Should resolve");
let cwd = Utf8PathBuf::try_from(std::env::current_dir().unwrap()).unwrap();
let source = cwd.join("src");
assert_eq!(watch, source);
assert!(
patterns
.iter()
.any(|p| p.matches_path(source.as_std_path()))
);
assert!(
patterns
.iter()
.any(|p| p.matches_path(source.join("lib.rs").as_std_path()))
);
}
#[test]
fn test_static_file_watch_path() {
let (watch, patterns) = resolve_static_watch_path("README.md").expect("Should resolve");
let cwd = Utf8PathBuf::try_from(std::env::current_dir().unwrap()).unwrap();
let source = cwd.join("README.md");
assert_eq!(watch, cwd);
assert_eq!(patterns.len(), 1);
assert!(patterns[0].matches_path(source.as_std_path()));
}
#[test]
fn test_collapse_watch_paths() {
let mut paths = HashSet::new();
paths.insert(Utf8PathBuf::from("/a"));
paths.insert(Utf8PathBuf::from("/a/b"));
paths.insert(Utf8PathBuf::from("/a/b/c"));
paths.insert(Utf8PathBuf::from("/b"));
paths.insert(Utf8PathBuf::from("/c/d"));
let collapsed = collapse_watch_paths(paths);
assert_eq!(
collapsed,
vec![
Utf8PathBuf::from("/a"),
Utf8PathBuf::from("/b"),
Utf8PathBuf::from("/c/d")
]
);
}
#[test]
fn test_collapse_watch_paths_siblings() {
let mut paths = HashSet::new();
paths.insert(Utf8PathBuf::from("/a/x"));
paths.insert(Utf8PathBuf::from("/a/y"));
let collapsed = collapse_watch_paths(paths);
assert_eq!(
collapsed,
vec![Utf8PathBuf::from("/a/x"), Utf8PathBuf::from("/a/y")]
);
}
#[test]
fn test_collapse_watch_paths_similar_names() {
let mut paths = HashSet::new();
paths.insert(Utf8PathBuf::from("/foo"));
paths.insert(Utf8PathBuf::from("/foo-bar"));
let collapsed = collapse_watch_paths(paths);
assert_eq!(
collapsed,
vec![Utf8PathBuf::from("/foo"), Utf8PathBuf::from("/foo-bar")]
);
}
}