use crate::error::{Error, Result};
use crate::message::Message;
use dashmap::DashMap;
use futures_util::{SinkExt, StreamExt};
use serde::{Deserialize, Serialize};
use std::net::SocketAddr;
use std::sync::Arc;
use tokio::net::TcpStream;
use tokio::sync::mpsc;
use tokio_tungstenite::WebSocketStream;
use tracing::{debug, error, info, warn};
pub type ConnectionId = String;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConnectionInfo {
pub id: ConnectionId,
pub addr: SocketAddr,
pub connected_at: u64,
pub protocol: Option<String>,
}
pub struct Connection {
pub id: ConnectionId,
pub info: ConnectionInfo,
sender: mpsc::UnboundedSender<Message>,
}
impl Connection {
pub fn new(id: ConnectionId, addr: SocketAddr, sender: mpsc::UnboundedSender<Message>) -> Self {
let info = ConnectionInfo {
id: id.clone(),
addr,
connected_at: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs(),
protocol: None,
};
Self { id, info, sender }
}
pub fn send(&self, message: Message) -> Result<()> {
self.sender
.send(message)
.map_err(|e| Error::custom(format!("Failed to send message: {}", e)))
}
pub fn send_text(&self, text: impl Into<String>) -> Result<()> {
self.send(Message::text(text.into()))
}
pub fn send_binary(&self, data: Vec<u8>) -> Result<()> {
self.send(Message::binary(data))
}
pub fn send_json<T: Serialize>(&self, data: &T) -> Result<()> {
let json = serde_json::to_string(data)?;
self.send_text(json)
}
pub fn id(&self) -> &ConnectionId {
&self.id
}
pub fn info(&self) -> &ConnectionInfo {
&self.info
}
}
pub struct ConnectionManager {
connections: Arc<DashMap<ConnectionId, Connection>>,
}
impl ConnectionManager {
pub fn new() -> Self {
Self {
connections: Arc::new(DashMap::new()),
}
}
pub fn add(&self, conn: Connection) -> usize {
let id = conn.id.clone();
self.connections.insert(id.clone(), conn);
let count = self.connections.len();
info!("Added connection: {} (Total: {})", id, count);
count
}
pub fn remove(&self, id: &ConnectionId) -> Option<Connection> {
let result = self.connections.remove(id).map(|(_, conn)| conn);
let count = self.connections.len();
info!("Removed connection: {} (Total: {})", id, count);
result
}
pub fn get(&self, id: &ConnectionId) -> Option<Connection> {
self.connections.get(id).map(|entry| entry.value().clone())
}
pub fn broadcast(&self, message: Message) {
let count = self.connections.len();
debug!("Broadcasting message to {} connections", count);
let mut success = 0;
let mut failed = 0;
for entry in self.connections.iter() {
match entry.value().send(message.clone()) {
Ok(_) => {
success += 1;
debug!("✅ Broadcast sent to {}", entry.key());
}
Err(e) => {
failed += 1;
error!("❌ Failed to broadcast to {}: {}", entry.key(), e);
}
}
}
info!(
"Broadcast complete: {} success, {} failed out of {} total",
success, failed, count
);
}
pub fn broadcast_except(&self, except_id: &ConnectionId, message: Message) {
debug!(
"Broadcasting message to {} connections (except {})",
self.connections.len() - 1,
except_id
);
for entry in self.connections.iter() {
if entry.key() != except_id {
if let Err(e) = entry.value().send(message.clone()) {
error!("Failed to broadcast to {}: {}", entry.key(), e);
}
}
}
}
pub fn broadcast_to(&self, ids: &[ConnectionId], message: Message) {
for id in ids {
if let Some(conn) = self.get(id) {
if let Err(e) = conn.send(message.clone()) {
error!("Failed to send to {}: {}", id, e);
}
}
}
}
pub fn count(&self) -> usize {
self.connections.len()
}
pub fn all_ids(&self) -> Vec<ConnectionId> {
self.connections.iter().map(|e| e.key().clone()).collect()
}
pub fn all_connections(&self) -> Vec<Connection> {
self.connections.iter().map(|e| e.value().clone()).collect()
}
}
impl Clone for Connection {
fn clone(&self) -> Self {
Self {
id: self.id.clone(),
info: self.info.clone(),
sender: self.sender.clone(),
}
}
}
impl Default for ConnectionManager {
fn default() -> Self {
Self::new()
}
}
pub async fn handle_websocket(
stream: WebSocketStream<TcpStream>,
conn_id: ConnectionId,
peer_addr: SocketAddr,
manager: Arc<ConnectionManager>,
on_message: Arc<dyn Fn(ConnectionId, Message) + Send + Sync>,
on_connect: Arc<dyn Fn(ConnectionId) + Send + Sync>,
on_disconnect: Arc<dyn Fn(ConnectionId) + Send + Sync>,
) {
info!(
"WebSocket connection established: {} from {}",
conn_id, peer_addr
);
let (mut ws_sender, mut ws_receiver) = stream.split();
let (tx, mut rx) = mpsc::unbounded_channel::<Message>();
let conn = Connection::new(conn_id.clone(), peer_addr, tx);
let _count = manager.add(conn);
let verify_count = manager.count();
debug!(
"Connection {} added. Verified count: {}",
conn_id, verify_count
);
on_connect(conn_id.clone());
let conn_id_write = conn_id.clone();
let write_task = tokio::spawn(async move {
debug!("Write task started for {}", conn_id_write);
while let Some(message) = rx.recv().await {
debug!("📤 Sending message to {}", conn_id_write);
let msg = message.into_tungstenite();
if let Err(e) = ws_sender.send(msg).await {
error!("Failed to send message to {}: {}", conn_id_write, e);
break;
}
debug!("✅ Message sent to {}", conn_id_write);
}
info!("Write task ended for {}", conn_id_write);
});
let conn_id_read = conn_id.clone();
let read_task = tokio::spawn(async move {
debug!("Read task started for {}", conn_id_read);
while let Some(result) = ws_receiver.next().await {
match result {
Ok(msg) => {
if msg.is_close() {
info!("Close message received from {}", conn_id_read);
break;
}
debug!("📨 Received message from {}", conn_id_read);
let message = Message::from_tungstenite(msg);
on_message(conn_id_read.clone(), message);
}
Err(e) => {
warn!("WebSocket error for {}: {}", conn_id_read, e);
break;
}
}
}
debug!("Read task ended for {}", conn_id_read);
});
tokio::select! {
_ = write_task => {
debug!("Write task finished first for {}", conn_id);
},
_ = read_task => {
debug!("Read task finished first for {}", conn_id);
},
}
manager.remove(&conn_id);
on_disconnect(conn_id);
}