use crate::{Error, HttpRequest, HttpResponse};
use notify::{Config, Event, EventKind, RecommendedWatcher, RecursiveMode, Watcher};
use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::{RwLock, broadcast};
#[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>,
clients: Arc<RwLock<Vec<ClientConnection>>>,
last_events: Arc<RwLock<HashMap<PathBuf, std::time::SystemTime>>>,
}
#[allow(dead_code)]
struct ClientConnection {
id: String,
connected_at: std::time::SystemTime,
}
impl HmrManager {
pub fn new(config: HmrConfig) -> Self {
let (event_tx, _) = broadcast::channel(100);
Self {
config,
event_tx,
clients: Arc::new(RwLock::new(Vec::new())),
last_events: Arc::new(RwLock::new(HashMap::new())),
}
}
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 handle_websocket(&self, _req: &HttpRequest) -> Result<HttpResponse, Error> {
Ok(HttpResponse::ok().with_body(b"WebSocket upgrade".to_vec()))
}
pub async fn register_client(&self, client_id: String) {
let mut clients = self.clients.write().await;
clients.push(ClientConnection {
id: client_id,
connected_at: std::time::SystemTime::now(),
});
}
pub async fn client_count(&self) -> usize {
self.clients.read().await.len()
}
}
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().await, 0);
}
#[tokio::test]
async fn test_client_registration() {
let config = HmrConfig::new();
let manager = HmrManager::new(config);
manager.register_client("test-client-1".to_string()).await;
manager.register_client("test-client-2".to_string()).await;
assert_eq!(manager.client_count().await, 2);
}
#[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>"));
}
}