use crate::connection::{ConnectionId, ConnectionManager, handle_websocket};
use crate::error::{Error, Result};
use crate::extractor::Extensions;
use crate::handler::Handler;
use crate::message::Message;
use crate::middleware::{Middleware, MiddlewareChain};
use crate::state::AppState;
use dashmap::DashMap;
use std::net::SocketAddr;
use std::path::PathBuf;
use std::sync::Arc;
use tokio::net::{TcpListener, TcpStream};
use tokio_tungstenite::accept_async;
use tracing::{error, info};
pub struct Route {
pub path: String,
pub chain: Arc<MiddlewareChain>,
}
pub struct Router {
routes: Arc<DashMap<String, Arc<MiddlewareChain>>>,
global_middlewares: Vec<Arc<dyn Middleware>>,
state: AppState,
connection_manager: Arc<ConnectionManager>,
on_connect: Option<Arc<dyn Fn(&Arc<ConnectionManager>, ConnectionId) + Send + Sync>>,
on_disconnect: Option<Arc<dyn Fn(&Arc<ConnectionManager>, ConnectionId) + Send + Sync>>,
default_chain: Option<Arc<MiddlewareChain>>,
static_handler: Option<crate::static_files::StaticFileHandler>,
}
impl Router {
pub fn new() -> Self {
Self {
routes: Arc::new(DashMap::new()),
global_middlewares: Vec::new(),
state: AppState::new(),
connection_manager: Arc::new(ConnectionManager::new()),
on_connect: None,
on_disconnect: None,
default_chain: None,
static_handler: None,
}
}
pub fn layer(mut self, middleware: Arc<dyn Middleware>) -> Self {
self.global_middlewares.push(middleware);
self
}
pub fn route(self, path: impl Into<String>, handler: Arc<dyn Handler>) -> Self {
let mut chain = MiddlewareChain::new();
for middleware in &self.global_middlewares {
chain = chain.layer(middleware.clone());
}
chain = chain.handler(handler);
self.routes.insert(path.into(), Arc::new(chain));
self
}
pub fn route_with_layers(
self,
path: impl Into<String>,
layers: Vec<Arc<dyn Middleware>>,
handler: Arc<dyn Handler>,
) -> Self {
let mut chain = MiddlewareChain::new();
for middleware in &self.global_middlewares {
chain = chain.layer(middleware.clone());
}
for middleware in layers {
chain = chain.layer(middleware);
}
chain = chain.handler(handler);
self.routes.insert(path.into(), Arc::new(chain));
self
}
pub fn with_state<T: Send + Sync + 'static>(self, data: Arc<T>) -> Self {
self.state.insert(data);
self
}
pub fn on_connect<F>(mut self, f: F) -> Self
where
F: Fn(&Arc<ConnectionManager>, ConnectionId) + Send + Sync + 'static,
{
self.on_connect = Some(Arc::new(f));
self
}
pub fn on_disconnect<F>(mut self, f: F) -> Self
where
F: Fn(&Arc<ConnectionManager>, ConnectionId) + Send + Sync + 'static,
{
self.on_disconnect = Some(Arc::new(f));
self
}
pub fn default_handler(mut self, handler: Arc<dyn Handler>) -> Self {
let mut chain = MiddlewareChain::new();
for middleware in &self.global_middlewares {
chain = chain.layer(middleware.clone());
}
chain = chain.handler(handler);
self.default_chain = Some(Arc::new(chain));
self
}
pub fn serve_static(mut self, path: impl Into<PathBuf>) -> Self {
self.static_handler = Some(crate::static_files::StaticFileHandler::new(path.into()));
self
}
pub fn connection_manager(&self) -> Arc<ConnectionManager> {
self.connection_manager.clone()
}
pub async fn listen(self, addr: impl AsRef<str>) -> Result<()> {
let addr: SocketAddr = addr
.as_ref()
.parse()
.map_err(|e| Error::custom(format!("Invalid address: {}", e)))?;
self.state.insert(self.connection_manager.clone());
let listener = TcpListener::bind(addr).await?;
info!("WebSocket server listening on {}", addr);
let router = Arc::new(self);
loop {
let (stream, peer_addr) = listener.accept().await?;
let router = router.clone();
tokio::spawn(async move {
if let Err(e) = router.handle_connection(stream, peer_addr).await {
error!("Connection error: {}", e);
}
});
}
}
async fn handle_connection(&self, stream: TcpStream, peer_addr: SocketAddr) -> Result<()> {
let mut buffer = [0u8; 1024];
let n = tokio::time::timeout(std::time::Duration::from_secs(5), stream.peek(&mut buffer))
.await
.map_err(|_| Error::custom("Connection timeout"))?
.map_err(|e| Error::custom(format!("Failed to read: {}", e)))?;
let header = String::from_utf8_lossy(&buffer[..n]);
if header.contains("Upgrade: websocket") || header.contains("upgrade: websocket") {
self.handle_websocket_connection(stream, peer_addr).await
} else if let Some(ref static_handler) = self.static_handler {
self.handle_http_request(stream, static_handler, &header)
.await
} else {
Err(Error::custom("No handler for HTTP requests"))
}
}
async fn handle_http_request(
&self,
mut stream: TcpStream,
static_handler: &crate::static_files::StaticFileHandler,
header: &str,
) -> Result<()> {
use crate::static_files::http_response;
use tokio::io::AsyncWriteExt;
let path = header
.lines()
.next()
.and_then(|line| {
let parts: Vec<&str> = line.split_whitespace().collect();
if parts.len() >= 2 && (parts[0] == "GET" || parts[0] == "HEAD") {
Some(parts[1])
} else {
None
}
})
.unwrap_or("/");
let response = match static_handler.serve(path).await {
Ok((content, mime_type)) => {
info!("Served: {} ({} bytes)", path, content.len());
http_response(200, &mime_type, content)
}
Err(e) => {
tracing::warn!("File not found: {} - {}", path, e);
let html = b"<html><body><h1>404 Not Found</h1></body></html>".to_vec();
http_response(404, "text/html", html)
}
};
stream.write_all(&response).await?;
stream.flush().await?;
Ok(())
}
async fn handle_websocket_connection(
&self,
stream: TcpStream,
peer_addr: SocketAddr,
) -> Result<()> {
let ws_stream = accept_async(stream).await?;
let conn_id = Self::generate_connection_id();
let router = self.clone();
let manager = self.connection_manager.clone();
let on_message = Arc::new(move |conn_id: ConnectionId, message: Message| {
let router = router.clone();
tokio::spawn(async move {
if let Err(e) = router.handle_message(conn_id, message).await {
error!("Message handling error: {}", e);
}
});
});
let manager_ref = manager.clone();
let on_connect = self
.on_connect
.clone()
.map(move |cb| {
let manager = manager_ref.clone();
Arc::new(move |conn_id: ConnectionId| {
cb(&manager, conn_id);
}) as Arc<dyn Fn(ConnectionId) + Send + Sync>
})
.unwrap_or_else(|| {
Arc::new(|conn_id: ConnectionId| {
info!("Client connected: {}", conn_id);
})
});
let manager_ref = manager.clone();
let on_disconnect = self
.on_disconnect
.clone()
.map(move |cb| {
let manager = manager_ref.clone();
Arc::new(move |conn_id: ConnectionId| {
cb(&manager, conn_id);
}) as Arc<dyn Fn(ConnectionId) + Send + Sync>
})
.unwrap_or_else(|| {
Arc::new(|conn_id: ConnectionId| {
info!("Client disconnected: {}", conn_id);
})
});
handle_websocket(
ws_stream,
conn_id,
peer_addr,
manager,
on_message,
on_connect,
on_disconnect,
)
.await;
Ok(())
}
async fn handle_message(&self, conn_id: ConnectionId, message: Message) -> Result<()> {
let conn = self
.connection_manager
.get(&conn_id)
.ok_or_else(|| Error::ConnectionNotFound(conn_id.clone()))?;
let extensions = Extensions::new();
let chain = if let Some(text) = message.as_text() {
if text.starts_with('/') {
if let Some((route, _)) = text.split_once(' ') {
self.routes.get(route).map(|c| c.value().clone())
} else {
self.routes.get(text).map(|c| c.value().clone())
}
} else {
None
}
} else {
None
};
let chain = chain.or_else(|| self.default_chain.clone());
if let Some(chain) = chain {
match chain
.execute(message, conn.clone(), self.state.clone(), extensions)
.await
{
Ok(Some(response)) => {
if let Err(e) = conn.send(response) {
error!("Failed to send response to {}: {}", conn_id, e);
}
}
Ok(None) => {
tracing::debug!("Handler processed message without response");
}
Err(e) => {
error!("Handler error for {}: {}", conn_id, e);
}
}
} else {
tracing::warn!("No handler found for message from {}", conn_id);
}
Ok(())
}
fn generate_connection_id() -> ConnectionId {
use std::sync::atomic::{AtomicU64, Ordering};
static COUNTER: AtomicU64 = AtomicU64::new(0);
format!("conn_{}", COUNTER.fetch_add(1, Ordering::SeqCst))
}
}
impl Clone for Router {
fn clone(&self) -> Self {
Self {
routes: self.routes.clone(),
global_middlewares: self.global_middlewares.clone(),
state: self.state.clone(),
connection_manager: self.connection_manager.clone(),
on_connect: self.on_connect.clone(),
on_disconnect: self.on_disconnect.clone(),
default_chain: self.default_chain.clone(),
static_handler: self.static_handler.clone(),
}
}
}
impl Default for Router {
fn default() -> Self {
Self::new()
}
}