use crate::Config;
use crate::actor::{Actor, ActorContext, Addr};
use crate::adapters::ws_conn::WsConn;
use crate::message::Message;
use crate::metrics::Metrics;
use crate::utils::FxHashSet;
use async_trait::async_trait;
use std::fs::File;
use std::io::Read;
use std::sync::Arc;
use tokio::sync::RwLock;
use log::{debug, info};
use tokio::net::TcpListener;
use tokio_native_tls::native_tls::Identity;
use tokio_websockets::ServerBuilder;
type Clients = Arc<RwLock<FxHashSet<Addr>>>;
#[derive(Clone)]
pub struct WsServerConfig {
pub port: u16,
pub cert_path: Option<String>,
pub key_path: Option<String>,
}
impl Default for WsServerConfig {
fn default() -> Self {
WsServerConfig {
port: 4944,
cert_path: None,
key_path: None,
}
}
}
#[derive(Clone)]
pub struct WsServer {
config: Config,
ws_config: WsServerConfig,
clients: Clients,
}
impl WsServer {
pub fn new(config: Config) -> Self {
Self::new_with_config(config, WsServerConfig::default())
}
pub fn new_with_config(config: Config, ws_config: WsServerConfig) -> Self {
Self {
config,
ws_config,
clients: Clients::default(),
}
}
async fn handle_stream<S>(
stream: S,
ctx: &ActorContext,
clients: Clients,
allow_public_space: bool,
) where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
{
let ws_config = tokio_websockets::Config::default().flush_threshold(usize::MAX);
let ws_stream = match ServerBuilder::new().config(ws_config).accept(stream).await {
Ok((_req, s)) => s,
Err(e) => {
log::warn!("WsServer WebSocket handshake failed: {}", e);
return;
}
};
let conn = WsConn::new(ws_stream, allow_public_space);
let addr = ctx.start_actor(Box::new(conn));
clients.write().await.insert(addr);
}
async fn start_web_server(config: WsServerConfig, peer_id: String, metrics: Arc<Metrics>) {
let port = config.port + 1;
if let Some(cert_path) = config.cert_path {
let key_path = config.key_path.unwrap();
let _addr = format!("https://localhost:{}", port);
let cert = std::fs::read(cert_path).expect("failed to read cert file");
let key = std::fs::read(key_path).expect("failed to read key file");
let identity = tokio_native_tls::native_tls::Identity::from_pkcs8(&cert, &key)
.expect("failed to create TLS identity");
let acceptor = tokio_native_tls::TlsAcceptor::from(
tokio_native_tls::native_tls::TlsAcceptor::new(identity).unwrap(),
);
let listener = tokio::net::TcpListener::bind(("0.0.0.0", port))
.await
.expect("failed to bind web UI port");
loop {
let (stream, _) = match listener.accept().await {
Ok(s) => s,
Err(e) => {
log::error!("web UI accept error: {}", e);
continue;
}
};
let acceptor = acceptor.clone();
let peer_id = peer_id.clone();
let metrics_clone = metrics.clone();
crate::tokio_spawn::spawn(async move {
let stream = match acceptor.accept(stream).await {
Ok(s) => s,
Err(_) => return,
};
Self::handle_http_request(stream, &peer_id, &metrics_clone).await;
});
}
}
let _addr = format!("http://localhost:{}", port);
let listener = tokio::net::TcpListener::bind(("0.0.0.0", port))
.await
.expect("failed to bind web UI port");
loop {
let (stream, _) = match listener.accept().await {
Ok(s) => s,
Err(e) => {
log::error!("web UI accept error: {}", e);
continue;
}
};
let peer_id = peer_id.clone();
let metrics_clone = metrics.clone();
crate::tokio_spawn::spawn(async move {
Self::handle_http_request_plain(stream, &peer_id, &metrics_clone).await;
});
}
}
fn build_http_response(request: &str, peer_id: &str, metrics: &Metrics) -> Option<String> {
let request_line = request.lines().next().unwrap_or("");
if request_line.contains("GET /peer_id") {
Some(format!(
"HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
peer_id.len(),
peer_id
))
} else if request_line.contains("GET /metrics") {
let body = serde_json::to_string_pretty(&metrics.snapshot())
.unwrap_or_else(|_| "{}".to_string());
Some(format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
body.len(),
body
))
} else {
None
}
}
async fn handle_http_request(
mut stream: tokio_native_tls::TlsStream<tokio::net::TcpStream>,
peer_id: &str,
metrics: &Metrics,
) {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut buf = [0u8; 1024];
let n = match stream.read(&mut buf).await {
Ok(n) => n,
Err(_) => return,
};
let request = String::from_utf8_lossy(&buf[..n]);
let response = match Self::build_http_response(&request, peer_id, metrics) {
Some(r) => r,
None => "HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
.to_string(),
};
let _ = stream.write_all(response.as_bytes()).await;
let _ = stream.shutdown().await;
}
async fn handle_http_request_plain(
mut stream: tokio::net::TcpStream,
peer_id: &str,
metrics: &Metrics,
) {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut buf = [0u8; 1024];
let n = match stream.read(&mut buf).await {
Ok(n) => n,
Err(_) => return,
};
let request = String::from_utf8_lossy(&buf[..n]);
let response = match Self::build_http_response(&request, peer_id, metrics) {
Some(r) => r,
None => "HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
.to_string(),
};
let _ = stream.write_all(response.as_bytes()).await;
let _ = stream.shutdown().await;
}
pub fn peer_count(&self) -> usize {
if let Ok(count) = self.clients.try_read() {
count.len()
} else {
0
}
}
}
#[async_trait]
impl Actor for WsServer {
async fn handle(&mut self, msg: Arc<Message>, _ctx: &ActorContext) {
match &*msg {
Message::Put(_)
| Message::Get(_)
| Message::BatchPut(_)
| Message::Hi { .. }
| Message::Flush(_)
| Message::RtcSignal(_) => {}
Message::CheckQuorumTimeouts | Message::RegisterQuorum { .. } => return,
}
for conn in self.clients.read().await.iter() {
if msg.is_from(conn) {
continue;
}
if conn.send((*msg).clone()).is_err() {
self.clients.write().await.remove(conn);
}
}
}
async fn pre_start(&mut self, ctx: &ActorContext) {
let addr = format!("0.0.0.0:{}", self.ws_config.port).to_string();
let ctx = ctx.clone();
let peer_id = ctx.peer_id.read().clone();
let config_clone = self.ws_config.clone();
let metrics = ctx.metrics.clone();
ctx.child_task(async move {
Self::start_web_server(config_clone, peer_id, metrics).await;
});
let try_socket = TcpListener::bind(&addr).await;
let listener = try_socket.expect("Failed to bind");
let allow_public_space = self.config.allow_public_space;
let clients = self.clients.clone();
if let Some(cert_path) = &self.ws_config.cert_path {
let mut cert_file = File::open(cert_path).unwrap();
let mut cert = vec![];
cert_file.read_to_end(&mut cert).unwrap();
let key_path = self.ws_config.key_path.as_ref().unwrap();
let mut key_file = File::open(key_path).unwrap();
let mut key = vec![];
key_file.read_to_end(&mut key).unwrap();
let identity = Identity::from_pkcs8(&cert, &key).unwrap();
let acceptor = tokio_native_tls::native_tls::TlsAcceptor::new(identity).unwrap();
let acceptor = tokio_native_tls::TlsAcceptor::from(acceptor);
let acceptor = Arc::new(acceptor);
let mut shutdown_rx = ctx.shutdown_rx.clone();
ctx.clone().child_task(async move {
loop {
tokio::select! {
biased;
_ = shutdown_rx.changed() => {
debug!("WsServer TLS accept loop shutting down");
break;
}
result = listener.accept() => {
if let Ok((stream, _)) = result {
let acceptor = acceptor.clone();
let clients = clients.clone();
let ctx = ctx.clone();
crate::tokio_spawn::spawn(async move {
let stream = acceptor.accept(stream).await;
if let Ok(stream) = stream {
Self::handle_stream(
stream,
&ctx,
clients.clone(),
allow_public_space,
)
.await;
}
});
}
}
}
}
});
} else {
let mut shutdown_rx = ctx.shutdown_rx.clone();
ctx.clone().child_task(async move {
loop {
tokio::select! {
biased;
_ = shutdown_rx.changed() => {
debug!("WsServer plain accept loop shutting down");
break;
}
result = listener.accept() => {
if let Ok((stream, _)) = result {
Self::handle_stream(
stream,
&ctx,
clients.clone(),
allow_public_space,
)
.await;
}
}
}
}
});
}
}
fn subscribe_to_everything(&self) -> bool {
true
}
fn is_relay_server(&self) -> bool {
true
}
async fn stopping(&mut self, _context: &ActorContext) {
info!(
"WsServer stopping — closing {} client connections",
self.clients.read().await.len()
);
self.clients.write().await.clear();
}
}