use std::net::{SocketAddr, TcpListener as StdListener};
use std::sync::mpsc;
use std::thread;
use futures_util::{SinkExt as _, StreamExt as _};
use tokio::net::TcpListener;
use tokio::sync::broadcast;
use tokio_tungstenite::tungstenite::Message;
use crate::{Command, Error, MessageEvent, Response};
const EVENT_BACKLOG: usize = 256;
#[derive(Debug, Clone)]
enum Outbound {
Event(Box<MessageEvent>),
Reply {
client: u64,
response: Box<Response>,
},
}
#[derive(Debug, Clone)]
pub struct Request {
pub client: u64,
pub command: Command,
}
#[derive(Debug)]
pub struct WebSocketBridge {
outbound: broadcast::Sender<Outbound>,
commands: mpsc::Receiver<Request>,
address: SocketAddr,
}
impl WebSocketBridge {
pub fn bind(address: &str) -> Result<Self, Error> {
let listener = StdListener::bind(address)?;
listener.set_nonblocking(true)?;
let address = listener.local_addr()?;
let (outbound, _) = broadcast::channel(EVENT_BACKLOG);
let (command_sender, commands) = mpsc::channel();
let publisher = outbound.clone();
thread::Builder::new()
.name("rs-teststand-websocket".to_owned())
.spawn(move || {
let Ok(runtime) = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
else {
return;
};
runtime.block_on(serve(listener, publisher, command_sender));
})
.map_err(|error| Error::ThreadNotStarted {
reason: error.to_string(),
})?;
Ok(Self {
outbound,
commands,
address,
})
}
#[must_use]
pub const fn address(&self) -> SocketAddr {
self.address
}
#[must_use]
pub fn client_count(&self) -> usize {
self.outbound.receiver_count()
}
pub fn publish(&self, event: &MessageEvent) {
let _ = self.outbound.send(Outbound::Event(Box::new(event.clone())));
}
pub fn reply(&self, request: &Request, response: &Response) {
let _ = self.outbound.send(Outbound::Reply {
client: request.client,
response: Box::new(response.clone()),
});
}
#[must_use]
pub fn next_command(&self) -> Option<Request> {
self.commands.try_recv().ok()
}
}
async fn serve(
listener: StdListener,
outbound: broadcast::Sender<Outbound>,
commands: mpsc::Sender<Request>,
) {
let Ok(listener) = TcpListener::from_std(listener) else {
return;
};
let mut next_client = 0_u64;
while let Ok((stream, _)) = listener.accept().await {
next_client += 1;
let client = next_client;
let subscription = outbound.subscribe();
let commands = commands.clone();
tokio::spawn(async move {
if let Ok(websocket) = tokio_tungstenite::accept_async(stream).await {
serve_client(websocket, client, subscription, commands).await;
}
});
}
}
async fn serve_client<S>(
websocket: tokio_tungstenite::WebSocketStream<S>,
client: u64,
mut outbound: broadcast::Receiver<Outbound>,
commands: mpsc::Sender<Request>,
) where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
let (mut sink, mut stream) = websocket.split();
loop {
tokio::select! {
received = outbound.recv() => {
let Ok(message) = received else {
break;
};
let text = match message {
Outbound::Event(event) => serde_json::to_string(&event),
Outbound::Reply { client: target, response } if target == client => {
serde_json::to_string(&response)
}
Outbound::Reply { .. } => continue,
};
let Ok(text) = text else { continue };
if sink.send(Message::text(text)).await.is_err() {
break;
}
}
received = stream.next() => {
let Some(Ok(frame)) = received else { break };
match frame {
Message::Text(text) => match serde_json::from_str::<Command>(&text) {
Ok(command) => {
if commands.send(Request { client, command }).is_err() {
break;
}
}
Err(error) => {
let reply = Response::Failed {
command: "unparsed".to_owned(),
reason: error.to_string(),
};
let Ok(text) = serde_json::to_string(&reply) else { continue };
if sink.send(Message::text(text)).await.is_err() {
break;
}
}
},
Message::Close(_) => break,
_ => {}
}
}
}
}
let _ = sink.close().await;
}