use futures_util::{SinkExt, StreamExt};
use std::sync::{Arc, Mutex};
use tokio::join;
use tokio::net::TcpListener;
use tokio::sync::Mutex as TokioMutex; use tokio_rustls::TlsAcceptor;
use tokio_tungstenite::tungstenite::protocol::Message;
use tokio_tungstenite::{WebSocketStream, accept_async};
use std::fs::File;
use std::io;
use std::io::BufReader;
use rustls::ServerConfig;
use rustls::pki_types::{CertificateDer, PrivateKeyDer};
use rustls_pemfile::{certs, pkcs8_private_keys};
use serde::{Deserialize, Serialize};
use tokio::sync::mpsc::{self, Receiver, Sender};
pub trait WebSocketStreamTraits: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin {}
impl<T> WebSocketStreamTraits for T where T: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin {}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct WebsocketConfig {
ip_address: String,
port: String,
cert_path: Option<String>,
key_path: Option<String>,
}
pub enum WsType {
Plain(TcpListener),
Secure(TcpListener, TlsAcceptor),
}
#[derive(Debug, Clone)]
pub enum WebSocketMessage {
Text(String),
Binary(Vec<u8>),
Ping(Vec<u8>),
Pong(Vec<u8>),
Close(Option<(u16, String)>),
}
impl From<Message> for WebSocketMessage {
fn from(msg: Message) -> Self {
match msg {
Message::Text(text) => WebSocketMessage::Text(text.to_string()), Message::Binary(data) => WebSocketMessage::Binary(data.to_vec()),
Message::Ping(data) => WebSocketMessage::Ping(data.to_vec()),
Message::Pong(data) => WebSocketMessage::Pong(data.to_vec()),
Message::Close(close_frame) => WebSocketMessage::Close(close_frame.map(|frame| (frame.code.into(), frame.reason.to_string()))),
_ => WebSocketMessage::Close(None),
}
}
}
impl From<WebSocketMessage> for Message {
fn from(msg: WebSocketMessage) -> Self {
match msg {
WebSocketMessage::Text(text) => Message::Text(text.into()),
WebSocketMessage::Binary(data) => Message::Binary(data.into()),
WebSocketMessage::Ping(data) => Message::Ping(data.into()),
WebSocketMessage::Pong(data) => Message::Pong(data.into()),
WebSocketMessage::Close(close_info) => {
if let Some((code, reason)) = close_info {
Message::Close(Some(tokio_tungstenite::tungstenite::protocol::CloseFrame {
code: code.into(),
reason: reason.into(),
}))
} else {
Message::Close(None)
}
}
}
}
}
#[derive(Clone)]
pub struct ClientConnection {
pub id: String,
pub tx: Sender<WebSocketMessage>,
}
pub struct WebsocketServer {
config: Arc<WebsocketConfig>,
clients: Arc<TokioMutex<Vec<ClientConnection>>>, message_tx: Sender<(String, WebSocketMessage)>,
message_rx: Arc<Mutex<Option<Receiver<(String, WebSocketMessage)>>>>,
}
impl WebsocketServer {
pub fn new(ip_address: String, port: String, cert_path: Option<String>, key_path: Option<String>) -> Self {
let (tx, rx) = mpsc::channel::<(String, WebSocketMessage)>(100);
WebsocketServer {
config: Arc::new(WebsocketConfig {
ip_address,
port,
cert_path,
key_path,
}),
clients: Arc::new(TokioMutex::new(Vec::new())), message_tx: tx,
message_rx: Arc::new(Mutex::new(Some(rx))),
}
}
pub fn get_message_sender(&self) -> Sender<(String, WebSocketMessage)> {
self.message_tx.clone()
}
pub fn take_message_receiver(&self) -> Option<Receiver<(String, WebSocketMessage)>> {
self.message_rx.lock().unwrap().take()
}
pub async fn start(&self) -> Receiver<(String, WebSocketMessage)> {
let (cert_path, key_path) = match (&self.config.cert_path, &self.config.key_path) {
(Some(cert), Some(key)) => (cert.clone(), key.clone()),
_ => ("".into(), "".into()),
};
let (certs_result, key_result) = join!(async { self.load_certs(cert_path).await }, async {
self.load_private_key(key_path).await
});
let ws_type = match (certs_result, key_result) {
(Ok(certs), Ok(key)) => {
let tls_config = ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(certs, key)
.expect("Invalid TLS config");
let tls_acceptor = TlsAcceptor::from(Arc::new(tls_config));
let secure_listener = TcpListener::bind(format!("{}:{}", &self.config.ip_address, &self.config.port))
.await
.unwrap();
tracing::info!("Starting secure WebSocket server...");
tracing::info!("Websocket listening on wss://{}:{}", &self.config.ip_address, &self.config.port);
WsType::Secure(secure_listener, tls_acceptor)
}
_ => {
tracing::info!("TLS not configured or cert/key files missing - falling back to ws://");
let plain_listener = TcpListener::bind(format!("{}:{}", &self.config.ip_address, &self.config.port))
.await
.unwrap();
tracing::info!("Starting plain WebSocket server...");
tracing::info!("Websocket listening on ws://{}:{}", &self.config.ip_address, &self.config.port);
WsType::Plain(plain_listener)
}
};
let (broadcast_tx, mut broadcast_rx) = mpsc::channel::<(Option<String>, WebSocketMessage)>(100);
let server_clone = self.clone();
let broadcast_tx_clone = broadcast_tx.clone();
let receiver = self.take_message_receiver().expect("Message receiver already taken");
tokio::spawn(async move {
match ws_type {
WsType::Secure(listener, tls_acceptor) => {
let acceptor = tls_acceptor.clone();
loop {
match listener.accept().await {
Ok((stream, addr)) => {
tracing::info!("New connection from: {}", addr);
let server = server_clone.clone();
let broadcast_tx = broadcast_tx_clone.clone();
let acceptor = acceptor.clone();
tokio::spawn(async move {
match acceptor.accept(stream).await {
Ok(tls_stream) => match accept_async(tls_stream).await {
Ok(ws_stream) => {
let client_id = uuid::Uuid::new_v4().to_string();
server.handle_connection(ws_stream, client_id, broadcast_tx).await;
}
Err(e) => tracing::error!("WebSocket upgrade failed: {}", e),
},
Err(e) => tracing::error!("TLS handshake failed: {}", e),
}
});
}
Err(e) => tracing::error!("Failed to accept connection: {}", e),
}
}
}
WsType::Plain(listener) => loop {
match listener.accept().await {
Ok((stream, addr)) => {
tracing::info!("New connection from: {}", addr);
let server = server_clone.clone();
let broadcast_tx = broadcast_tx_clone.clone();
tokio::spawn(async move {
match accept_async(stream).await {
Ok(ws_stream) => {
let client_id = uuid::Uuid::new_v4().to_string();
server.handle_connection(ws_stream, client_id, broadcast_tx).await;
}
Err(e) => tracing::error!("WebSocket upgrade failed: {}", e),
}
});
}
Err(e) => tracing::error!("Failed to accept connection: {}", e),
}
},
}
});
let server_clone = self.clone();
tokio::spawn(async move {
while let Some((target_client_id, message)) = broadcast_rx.recv().await {
server_clone.broadcast_message(target_client_id, message).await;
}
});
receiver
}
async fn broadcast_message(&self, target_client_id: Option<String>, message: WebSocketMessage) {
let clients = self.clients.lock().await;
for client in clients.iter() {
if let Some(target_id) = &target_client_id {
if &client.id != target_id {
continue;
}
}
if let Err(e) = client.tx.send(message.clone()).await {
tracing::error!("Failed to send message to client {}: {}", client.id, e);
}
}
}
async fn handle_connection<S>(&self, stream: WebSocketStream<S>, client_id: String, _broadcast_tx: Sender<(Option<String>, WebSocketMessage)>)
where
S: WebSocketStreamTraits + Send + 'static,
{
tracing::info!("Handling new WebSocket connection for client: {}", client_id);
let (mut ws_sender, mut ws_receiver) = stream.split();
let (client_tx, mut client_rx) = mpsc::channel::<WebSocketMessage>(100);
{
let mut clients = self.clients.lock().await; clients.push(ClientConnection {
id: client_id.clone(),
tx: client_tx.clone(),
});
}
let message_tx = self.message_tx.clone();
let client_id_clone = client_id.clone();
let server_arc = Arc::new(self.clone());
tokio::spawn(async move {
while let Some(result) = ws_receiver.next().await {
match result {
Ok(msg) => {
let ws_msg = WebSocketMessage::from(msg);
tracing::info!("Received message from client {}: {:?}", client_id_clone, ws_msg);
if let Err(e) = message_tx.send((client_id_clone.clone(), ws_msg)).await {
tracing::error!("Failed to forward message: {}", e);
break;
}
}
Err(e) => {
tracing::error!("Error receiving message from {}: {}", client_id_clone, e);
break;
}
}
}
tracing::info!("Client {} disconnected", client_id_clone);
let server_arc_clone = server_arc.clone();
{
let mut clients = server_arc_clone.clients.lock().await; if let Some(pos) = clients.iter().position(|c| c.id == client_id_clone) {
clients.remove(pos);
}
}
});
let client_id_clone = client_id.clone();
tokio::spawn(async move {
while let Some(msg) = client_rx.recv().await {
let tungstenite_msg: Message = msg.into();
if let Err(e) = ws_sender.send(tungstenite_msg).await {
tracing::error!("Error sending message to {}: {}", client_id_clone, e);
break;
}
}
let _ = ws_sender.close().await;
});
}
pub async fn send_to_client(&self, client_id: String, message: WebSocketMessage) -> Result<(), String> {
let clients = self.clients.lock().await;
for client in clients.iter() {
if client.id == client_id {
return client.tx.send(message).await.map_err(|e| format!("Failed to send message: {}", e));
}
}
Err(format!("Client {} not found", client_id))
}
pub async fn broadcast(&self, message: WebSocketMessage) -> Result<(), String> {
let clients = self.clients.lock().await;
for client in clients.iter() {
if let Err(e) = client.tx.send(message.clone()).await {
return Err(format!("Failed to broadcast to client {}: {}", client.id, e));
}
}
Ok(())
}
pub async fn get_clients(&self) -> Vec<String> {
let clients = self.clients.lock().await; clients.iter().map(|client| client.id.clone()).collect()
}
async fn load_certs(&self, path: String) -> Result<Vec<CertificateDer<'static>>, io::Error> {
let file = File::open(path)?;
let mut reader = BufReader::new(file);
let cert_list = certs(&mut reader).filter_map(Result::ok).collect::<Vec<_>>();
if cert_list.is_empty() {
Err(io::Error::new(io::ErrorKind::InvalidData, "No valid certificates found"))
} else {
Ok(cert_list)
}
}
async fn load_private_key(&self, path: String) -> Result<PrivateKeyDer<'static>, io::Error> {
let file = File::open(path)?;
let mut reader = BufReader::new(file);
let key = pkcs8_private_keys(&mut reader).filter_map(Result::ok).next();
match key {
Some(k) => Ok(PrivateKeyDer::Pkcs8(k)),
None => Err(io::Error::new(io::ErrorKind::InvalidData, "No valid private key found")),
}
}
}
impl Clone for WebsocketServer {
fn clone(&self) -> Self {
WebsocketServer {
config: self.config.clone(),
clients: self.clients.clone(),
message_tx: self.message_tx.clone(),
message_rx: self.message_rx.clone(),
}
}
}