use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use futures_util::{SinkExt, StreamExt};
use serde_json::Value;
use tokio::sync::mpsc::{self, Sender};
use tokio_tungstenite::tungstenite::Message;
use super::{SignalingProtocol, WsStream};
const OUTBOUND_CAP: usize = 256;
#[derive(Default)]
struct RelayHub {
next_id: AtomicU64,
rooms: Mutex<HashMap<String, HashMap<u64, Sender<Message>>>>,
}
impl RelayHub {
fn next_conn_id(&self) -> u64 {
self.next_id.fetch_add(1, Ordering::Relaxed)
}
fn join(&self, room: &str, conn_id: u64, tx: Sender<Message>) {
let mut rooms = self.rooms.lock().unwrap();
rooms
.entry(room.to_string())
.or_default()
.insert(conn_id, tx);
}
fn broadcast(&self, room: &str, from: u64, msg: &Message) {
let rooms = self.rooms.lock().unwrap();
if let Some(members) = rooms.get(room) {
for (&cid, tx) in members.iter() {
if cid != from {
let _ = tx.try_send(msg.clone());
}
}
}
}
fn leave(&self, room: &str, conn_id: u64) {
let mut rooms = self.rooms.lock().unwrap();
if let Some(members) = rooms.get_mut(room) {
members.remove(&conn_id);
if members.is_empty() {
rooms.remove(room);
}
}
}
}
fn room_of(text: &str) -> Option<String> {
let v: Value = serde_json::from_str(text).ok()?;
let room = v.get("room")?.as_str()?;
if room.is_empty() {
None
} else {
Some(room.to_string())
}
}
pub struct JsonRelayProtocol {
hub: Arc<RelayHub>,
}
impl JsonRelayProtocol {
pub fn new() -> Self {
Self {
hub: Arc::new(RelayHub::default()),
}
}
}
impl Default for JsonRelayProtocol {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl SignalingProtocol for JsonRelayProtocol {
fn id(&self) -> &'static str {
"json-relay"
}
async fn handle(&self, ws: WsStream, _peer: SocketAddr) {
let hub = self.hub.clone();
let conn_id = hub.next_conn_id();
let (mut sink, mut stream) = ws.split();
let (tx, mut rx) = mpsc::channel::<Message>(OUTBOUND_CAP);
let writer = tokio::spawn(async move {
while let Some(m) = rx.recv().await {
if sink.send(m).await.is_err() {
break;
}
}
let _ = sink.close().await;
});
let mut room: Option<String> = None;
while let Some(next) = stream.next().await {
match next {
Ok(Message::Text(t)) => {
let text = t.as_str();
match &room {
Some(r) => hub.broadcast(r, conn_id, &Message::text(text.to_string())),
None => {
if let Some(r) = room_of(text) {
hub.join(&r, conn_id, tx.clone());
hub.broadcast(&r, conn_id, &Message::text(text.to_string()));
room = Some(r);
}
}
}
}
Ok(Message::Ping(p)) => {
let _ = tx.try_send(Message::Pong(p));
}
Ok(Message::Close(_)) => break,
Ok(_) => {}
Err(_) => break,
}
}
if let Some(r) = room {
hub.leave(&r, conn_id);
}
writer.abort();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn broadcast_reaches_others_not_sender() {
let hub = RelayHub::default();
let a = hub.next_conn_id();
let b = hub.next_conn_id();
let (atx, mut arx) = mpsc::channel::<Message>(OUTBOUND_CAP);
let (btx, mut brx) = mpsc::channel::<Message>(OUTBOUND_CAP);
hub.join("r", a, atx);
hub.join("r", b, btx);
let msg = Message::text(r#"{"room":"r","hello":"world"}"#.to_string());
hub.broadcast("r", a, &msg);
match brx.try_recv().expect("B should receive the broadcast") {
Message::Text(t) => {
let v: Value = serde_json::from_str(t.as_str()).unwrap();
assert_eq!(v["room"], "r");
assert_eq!(v["hello"], "world");
}
other => panic!("unexpected message: {other:?}"),
}
assert!(
arx.try_recv().is_err(),
"sender must not receive its own message"
);
}
#[test]
fn leaving_empties_and_drops_the_room() {
let hub = RelayHub::default();
let a = hub.next_conn_id();
let (atx, _arx) = mpsc::channel::<Message>(OUTBOUND_CAP);
hub.join("r", a, atx);
assert!(hub.rooms.lock().unwrap().contains_key("r"));
hub.leave("r", a);
assert!(!hub.rooms.lock().unwrap().contains_key("r"));
}
#[test]
fn room_of_parses_field() {
assert_eq!(room_of(r#"{"room":"abc"}"#).as_deref(), Some("abc"));
assert_eq!(room_of(r#"{"room":""}"#), None);
assert_eq!(room_of(r#"{"no":"room"}"#), None);
assert_eq!(room_of("not json"), None);
}
}