use crate::Error;
use futures_util::{SinkExt, StreamExt};
use notify::{Config, Event, EventKind, RecommendedWatcher, RecursiveMode, Watcher};
use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::{RwLock, broadcast};
use tokio::task::JoinHandle;
use tokio_tungstenite::WebSocketStream;
use tokio_tungstenite::tungstenite::Message as WsMessage;
#[derive(Debug, Clone)]
pub struct HmrEvent {
pub kind: HmrEventKind,
pub path: PathBuf,
pub extension: Option<String>,
pub timestamp: std::time::SystemTime,
}
#[derive(Debug, Clone, PartialEq)]
pub enum HmrEventKind {
Modified,
Created,
Deleted,
}
#[derive(Clone, Debug)]
pub struct HmrConfig {
pub enabled: bool,
pub watch_paths: Vec<PathBuf>,
pub watch_extensions: Vec<String>,
pub ignore_patterns: Vec<String>,
pub debounce_ms: u64,
pub websocket_port: u16,
pub verbose: bool,
}
impl Default for HmrConfig {
fn default() -> Self {
Self {
enabled: true,
watch_paths: vec![PathBuf::from("src"), PathBuf::from("public")],
watch_extensions: vec![
"js".to_string(),
"ts".to_string(),
"jsx".to_string(),
"tsx".to_string(),
"css".to_string(),
"scss".to_string(),
"less".to_string(),
"html".to_string(),
"vue".to_string(),
"svelte".to_string(),
],
ignore_patterns: vec![
"node_modules".to_string(),
"dist".to_string(),
"build".to_string(),
".git".to_string(),
"target".to_string(),
],
debounce_ms: 100,
websocket_port: 3001,
verbose: false,
}
}
}
impl HmrConfig {
pub fn new() -> Self {
Self::default()
}
pub fn watch_path(mut self, path: PathBuf) -> Self {
self.watch_paths.push(path);
self
}
pub fn watch_extension(mut self, ext: String) -> Self {
self.watch_extensions.push(ext);
self
}
pub fn ignore_pattern(mut self, pattern: String) -> Self {
self.ignore_patterns.push(pattern);
self
}
pub fn debounce(mut self, ms: u64) -> Self {
self.debounce_ms = ms;
self
}
pub fn websocket_port(mut self, port: u16) -> Self {
self.websocket_port = port;
self
}
pub fn verbose(mut self, enabled: bool) -> Self {
self.verbose = enabled;
self
}
}
pub struct HmrManager {
config: HmrConfig,
event_tx: broadcast::Sender<HmrEvent>,
client_count: Arc<AtomicUsize>,
last_events: Arc<RwLock<HashMap<PathBuf, std::time::SystemTime>>>,
server_handle: std::sync::Mutex<Option<JoinHandle<()>>>,
}
impl HmrManager {
pub fn new(config: HmrConfig) -> Self {
let (event_tx, _) = broadcast::channel(100);
Self {
config,
event_tx,
client_count: Arc::new(AtomicUsize::new(0)),
last_events: Arc::new(RwLock::new(HashMap::new())),
server_handle: std::sync::Mutex::new(None),
}
}
pub async fn start_watching(&self) -> Result<(), Error> {
if !self.config.enabled {
println!("🔥 HMR disabled");
return Ok(());
}
println!("🔥 HMR enabled - watching for changes...");
if self.config.verbose {
println!(" Watching paths:");
for path in &self.config.watch_paths {
println!(" - {}", path.display());
}
println!(" Extensions: {:?}", self.config.watch_extensions);
}
let event_tx = self.event_tx.clone();
let config = self.config.clone();
let last_events = self.last_events.clone();
tokio::spawn(async move {
if let Err(e) = Self::watch_files(event_tx, config, last_events).await {
eprintln!("❌ HMR file watcher error: {}", e);
}
});
Ok(())
}
async fn watch_files(
event_tx: broadcast::Sender<HmrEvent>,
config: HmrConfig,
last_events: Arc<RwLock<HashMap<PathBuf, std::time::SystemTime>>>,
) -> Result<(), Error> {
let (tx, mut rx) = tokio::sync::mpsc::channel(100);
let mut watcher = RecommendedWatcher::new(
move |res: Result<Event, notify::Error>| {
if let Ok(event) = res {
let _ = tx.blocking_send(event);
}
},
Config::default().with_poll_interval(Duration::from_millis(config.debounce_ms)),
)
.map_err(|e| Error::Internal(format!("Failed to create watcher: {}", e)))?;
for path in &config.watch_paths {
if path.exists() {
watcher
.watch(path, RecursiveMode::Recursive)
.map_err(|e| Error::Internal(format!("Failed to watch path: {}", e)))?;
} else if config.verbose {
println!("⚠️ Path not found: {}", path.display());
}
}
while let Some(event) = rx.recv().await {
if let Some(hmr_event) = Self::process_event(event, &config, &last_events).await {
let _ = event_tx.send(hmr_event);
}
}
Ok(())
}
async fn process_event(
event: Event,
config: &HmrConfig,
last_events: &Arc<RwLock<HashMap<PathBuf, std::time::SystemTime>>>,
) -> Option<HmrEvent> {
let kind = match event.kind {
EventKind::Modify(_) => HmrEventKind::Modified,
EventKind::Create(_) => HmrEventKind::Created,
EventKind::Remove(_) => HmrEventKind::Deleted,
_ => return None,
};
let path = event.paths.first()?.clone();
if Self::should_ignore(&path, &config.ignore_patterns) {
return None;
}
let extension = path.extension()?.to_str()?.to_string();
if !config.watch_extensions.contains(&extension) {
return None;
}
let now = std::time::SystemTime::now();
let mut last_events_map = last_events.write().await;
if let Some(last_time) = last_events_map.get(&path)
&& let Ok(duration) = now.duration_since(*last_time)
&& duration.as_millis() < config.debounce_ms as u128
{
return None;
}
last_events_map.insert(path.clone(), now);
if config.verbose {
println!("🔄 HMR: {:?} - {}", kind, path.display());
}
Some(HmrEvent {
kind,
path: path.clone(),
extension: Some(extension),
timestamp: now,
})
}
fn should_ignore(path: &Path, ignore_patterns: &[String]) -> bool {
let path_str = path.to_string_lossy();
ignore_patterns
.iter()
.any(|pattern| path_str.contains(pattern))
}
pub fn subscribe(&self) -> broadcast::Receiver<HmrEvent> {
self.event_tx.subscribe()
}
pub fn get_client_script(&self) -> String {
format!(
r#"<script>
(function() {{
console.log('🔥 HMR Client initialized');
let ws;
let reconnectAttempts = 0;
const maxReconnectAttempts = 10;
function connect() {{
ws = new WebSocket('ws://localhost:{}');
ws.onopen = function() {{
console.log('🔥 HMR Connected');
reconnectAttempts = 0;
}};
ws.onmessage = function(event) {{
const data = JSON.parse(event.data);
console.log('🔥 HMR Update:', data);
if (data.type === 'full-reload') {{
console.log('🔥 HMR: Full page reload');
window.location.reload();
}} else if (data.type === 'css-update') {{
console.log('🔥 HMR: CSS hot reload');
reloadCSS(data.path);
}} else if (data.type === 'js-update') {{
console.log('🔥 HMR: JavaScript update, reloading...');
window.location.reload();
}}
}};
ws.onclose = function() {{
console.log('🔥 HMR Disconnected');
if (reconnectAttempts < maxReconnectAttempts) {{
reconnectAttempts++;
setTimeout(connect, 1000 * reconnectAttempts);
}}
}};
ws.onerror = function(error) {{
console.error('🔥 HMR Error:', error);
}};
}}
function reloadCSS(path) {{
const links = document.querySelectorAll('link[rel="stylesheet"]');
links.forEach(link => {{
if (!path || link.href.includes(path)) {{
const href = link.href.split('?')[0];
link.href = href + '?t=' + Date.now();
}}
}});
}}
connect();
}})();
</script>"#,
self.config.websocket_port
)
}
pub async fn start_websocket_server(&self) -> Result<u16, Error> {
let listener = TcpListener::bind(("127.0.0.1", self.config.websocket_port))
.await
.map_err(|e| Error::Internal(format!("HMR WebSocket bind failed: {}", e)))?;
let port = listener
.local_addr()
.map_err(|e| Error::Internal(format!("HMR WebSocket local_addr failed: {}", e)))?
.port();
if self.config.verbose {
println!("🔥 HMR WebSocket server listening on ws://127.0.0.1:{port}");
}
let event_tx = self.event_tx.clone();
let client_count = self.client_count.clone();
let verbose = self.config.verbose;
let handle = tokio::spawn(async move {
loop {
let (stream, _addr) = match listener.accept().await {
Ok(conn) => conn,
Err(e) => {
if verbose {
eprintln!("❌ HMR WebSocket accept error: {}", e);
}
continue;
}
};
let rx = event_tx.subscribe();
let client_count = client_count.clone();
tokio::spawn(async move {
let handshake = tokio::time::timeout(
Duration::from_secs(10),
tokio_tungstenite::accept_async(stream),
)
.await;
match handshake {
Ok(Ok(ws)) => {
client_count.fetch_add(1, Ordering::SeqCst);
Self::forward_events(ws, rx, verbose).await;
client_count.fetch_sub(1, Ordering::SeqCst);
}
Ok(Err(e)) => {
if verbose {
eprintln!("❌ HMR WebSocket handshake failed: {}", e);
}
}
Err(_elapsed) => {
if verbose {
eprintln!(
"❌ HMR WebSocket handshake timed out; dropping connection"
);
}
}
}
});
}
});
let mut slot = self
.server_handle
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if let Some(old) = slot.replace(handle) {
old.abort();
}
Ok(port)
}
async fn forward_events(
mut ws: WebSocketStream<TcpStream>,
mut rx: broadcast::Receiver<HmrEvent>,
verbose: bool,
) {
loop {
tokio::select! {
event = rx.recv() => match event {
Ok(event) => {
let message = Self::reload_message(&event);
if ws.send(WsMessage::text(message)).await.is_err() {
break;
}
}
Err(broadcast::error::RecvError::Lagged(skipped)) => {
if verbose {
eprintln!("⚠️ HMR client lagged, skipped {} events", skipped);
}
}
Err(broadcast::error::RecvError::Closed) => break,
},
incoming = ws.next() => match incoming {
Some(Ok(msg)) if msg.is_close() => break,
Some(Ok(_)) => {} Some(Err(_)) | None => break,
},
}
}
let _ = ws.close(None).await;
}
fn reload_message(event: &HmrEvent) -> String {
let msg_type = match event.extension.as_deref() {
Some("css") | Some("scss") | Some("less") => "css-update",
Some("js") | Some("ts") | Some("jsx") | Some("tsx") => "js-update",
_ => "full-reload",
};
serde_json::json!({
"type": msg_type,
"path": event.path.to_string_lossy(),
})
.to_string()
}
pub fn publish(&self, event: HmrEvent) -> usize {
self.event_tx.send(event).unwrap_or(0)
}
pub fn client_count(&self) -> usize {
self.client_count.load(Ordering::SeqCst)
}
pub fn stop(&self) {
let mut slot = self
.server_handle
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if let Some(handle) = slot.take() {
handle.abort();
}
}
}
impl Drop for HmrManager {
fn drop(&mut self) {
self.stop();
}
}
pub async fn inject_hmr_script(html: String, hmr_manager: &HmrManager) -> String {
if !hmr_manager.config.enabled {
return html;
}
let script = hmr_manager.get_client_script();
if let Some(pos) = html.rfind("</body>") {
let mut result = html;
result.insert_str(pos, &script);
result
} else {
format!("{}{}", html, script)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_hmr_config_builder() {
let config = HmrConfig::new()
.watch_path(PathBuf::from("src"))
.watch_extension("rs".to_string())
.debounce(200)
.websocket_port(3002);
assert_eq!(config.debounce_ms, 200);
assert_eq!(config.websocket_port, 3002);
assert!(config.watch_paths.contains(&PathBuf::from("src")));
}
#[test]
fn test_should_ignore() {
let ignore_patterns = vec!["node_modules".to_string(), ".git".to_string()];
assert!(HmrManager::should_ignore(
Path::new("node_modules/package/index.js"),
&ignore_patterns
));
assert!(HmrManager::should_ignore(
Path::new(".git/config"),
&ignore_patterns
));
assert!(!HmrManager::should_ignore(
Path::new("src/main.ts"),
&ignore_patterns
));
}
#[tokio::test]
async fn test_hmr_manager_creation() {
let config = HmrConfig::new();
let manager = HmrManager::new(config);
assert_eq!(manager.client_count(), 0);
}
#[test]
fn test_reload_message_types() {
let event = |ext: &str| HmrEvent {
kind: HmrEventKind::Modified,
path: PathBuf::from(format!("src/app.{ext}")),
extension: Some(ext.to_string()),
timestamp: std::time::SystemTime::now(),
};
let css: serde_json::Value =
serde_json::from_str(&HmrManager::reload_message(&event("css"))).unwrap();
assert_eq!(css["type"], "css-update");
assert_eq!(css["path"], "src/app.css");
let js: serde_json::Value =
serde_json::from_str(&HmrManager::reload_message(&event("ts"))).unwrap();
assert_eq!(js["type"], "js-update");
let html: serde_json::Value =
serde_json::from_str(&HmrManager::reload_message(&event("html"))).unwrap();
assert_eq!(html["type"], "full-reload");
}
#[tokio::test]
async fn test_websocket_server_delivers_reload_message() {
let config = HmrConfig::new().websocket_port(0);
let manager = HmrManager::new(config);
let port = manager.start_websocket_server().await.unwrap();
let (mut ws, _) = tokio_tungstenite::connect_async(format!("ws://127.0.0.1:{port}"))
.await
.expect("client should connect to HMR WebSocket server");
for _ in 0..200 {
if manager.client_count() == 1 {
break;
}
tokio::time::sleep(Duration::from_millis(5)).await;
}
assert_eq!(manager.client_count(), 1);
let delivered = manager.publish(HmrEvent {
kind: HmrEventKind::Modified,
path: PathBuf::from("public/styles/app.css"),
extension: Some("css".to_string()),
timestamp: std::time::SystemTime::now(),
});
assert_eq!(delivered, 1);
let msg = tokio::time::timeout(Duration::from_secs(5), ws.next())
.await
.expect("timed out waiting for reload message")
.expect("stream ended before reload message")
.expect("websocket error");
let text = msg.into_text().expect("reload message should be text");
let json: serde_json::Value = serde_json::from_str(&text).unwrap();
assert_eq!(json["type"], "css-update");
assert_eq!(json["path"], "public/styles/app.css");
drop(ws);
for _ in 0..200 {
if manager.client_count() == 0 {
break;
}
tokio::time::sleep(Duration::from_millis(5)).await;
}
assert_eq!(manager.client_count(), 0);
manager.stop();
}
#[tokio::test]
async fn test_incomplete_handshake_never_counts_as_client() {
let config = HmrConfig::new().websocket_port(0);
let manager = HmrManager::new(config);
let port = manager.start_websocket_server().await.unwrap();
let _sock = tokio::net::TcpStream::connect(("127.0.0.1", port))
.await
.expect("raw TCP connect should succeed");
for _ in 0..40 {
assert_eq!(manager.client_count(), 0);
tokio::time::sleep(Duration::from_millis(5)).await;
}
assert_eq!(manager.client_count(), 0);
manager.stop();
}
#[test]
fn test_client_script_generation() {
let config = HmrConfig::new().websocket_port(3333);
let manager = HmrManager::new(config);
let script = manager.get_client_script();
assert!(script.contains("ws://localhost:3333"));
assert!(script.contains("HMR Client initialized"));
}
#[tokio::test]
async fn test_inject_hmr_script() {
let config = HmrConfig::new();
let manager = HmrManager::new(config);
let html = "<html><body><h1>Hello</h1></body></html>".to_string();
let result = inject_hmr_script(html, &manager).await;
assert!(result.contains("<script>"));
assert!(result.contains("HMR Client"));
assert!(result.contains("</body>"));
}
#[tokio::test]
async fn test_inject_hmr_script_no_body() {
let config = HmrConfig::new();
let manager = HmrManager::new(config);
let html = "<html><div>Content</div></html>".to_string();
let result = inject_hmr_script(html, &manager).await;
assert!(result.contains("<script>"));
assert!(result.ends_with("</script>"));
}
}