use std::collections::{HashMap, HashSet, VecDeque};
use std::sync::Arc;
use tokio::sync::{mpsc, oneshot, Mutex};
use authsocket::client::AuthSocketClient;
use bsv::wallet::interfaces::WalletInterface;
use crate::encryption;
use crate::error::MessageBoxError;
use crate::types::{AuthenticatedPeerMessage, ServerPeerMessage};
type SubscriptionMap =
Arc<Mutex<HashMap<String, Arc<dyn Fn(AuthenticatedPeerMessage) + Send + Sync>>>>;
type PendingAcks = Arc<Mutex<HashMap<String, VecDeque<oneshot::Sender<bool>>>>>;
pub struct MessageBoxWebSocket {
ws: AuthSocketClient,
subscriptions: SubscriptionMap,
pending_acks: PendingAcks,
joined_rooms: Arc<Mutex<HashSet<String>>>,
}
impl MessageBoxWebSocket {
pub async fn connect<W>(
url: &str,
identity_key: &str,
wallet: W,
originator: Option<String>,
) -> Result<Self, MessageBoxError>
where
W: WalletInterface + Clone + Send + Sync + 'static,
{
let ws = AuthSocketClient::connect(url, identity_key, wallet.clone())
.await
.map_err(|e| MessageBoxError::WebSocket(e.to_string()))?;
let subscriptions: SubscriptionMap = Arc::new(Mutex::new(HashMap::new()));
let pending_acks: PendingAcks = Arc::new(Mutex::new(HashMap::new()));
let joined_rooms: Arc<Mutex<HashSet<String>>> = Arc::new(Mutex::new(HashSet::new()));
let (event_tx, mut event_rx) = mpsc::unbounded_channel::<(String, serde_json::Value)>();
ws.set_fallback(Arc::new(move |event_name, data| {
let _ = event_tx.send((event_name, data));
}))
.await;
{
let subscriptions = subscriptions.clone();
let pending_acks = pending_acks.clone();
let wallet = wallet.clone();
let originator = originator.clone();
tokio::spawn(async move {
while let Some((event_name, data)) = event_rx.recv().await {
if let Some(room_id) = event_name.strip_prefix("sendMessage-") {
let room_id = room_id.to_string();
let Ok(server_msg) =
serde_json::from_value::<ServerPeerMessage>(data.clone())
else {
continue;
};
let callback = {
let guard = subscriptions.lock().await;
guard.get(&event_name).cloned()
};
if let Some(cb) = callback {
let outcome = encryption::try_decrypt_message_typed(
&wallet,
&server_msg.body,
&server_msg.sender,
originator.as_deref(),
)
.await;
let authenticated_decrypt = outcome.is_authenticated();
let (recipient, message_box) = split_room_id(&room_id);
cb(AuthenticatedPeerMessage {
message_id: server_msg.message_id,
sender: server_msg.sender,
recipient,
message_box,
body: outcome.into_body(),
authenticated_decrypt,
});
}
} else if event_name.starts_with("sendMessageAck-") {
let success =
data.get("status").and_then(|s| s.as_str()) == Some("success");
let mut guard = pending_acks.lock().await;
if let Some(queue) = guard.get_mut(&event_name) {
if let Some(tx) = queue.pop_front() {
let _ = tx.send(success);
}
if queue.is_empty() {
guard.remove(&event_name);
}
}
}
}
});
}
Ok(Self {
ws,
subscriptions,
pending_acks,
joined_rooms,
})
}
pub fn is_connected(&self) -> bool {
self.ws.is_connected()
}
pub fn ms_since_last_inbound(&self) -> u64 {
self.ws.ms_since_last_inbound()
}
pub fn server_identity_key(&self) -> &str {
self.ws.server_identity_key()
}
pub async fn join_room(&self, room_id: &str) -> Result<(), MessageBoxError> {
{
let guard = self.joined_rooms.lock().await;
if guard.contains(room_id) {
return Ok(());
}
}
self.ws
.join_room(room_id)
.await
.map_err(|e| MessageBoxError::WebSocket(e.to_string()))?;
self.joined_rooms.lock().await.insert(room_id.to_string());
Ok(())
}
pub async fn leave_room(&self, room_id: &str) -> Result<(), MessageBoxError> {
self.joined_rooms.lock().await.remove(room_id);
let event_key = format!("sendMessage-{room_id}");
self.subscriptions.lock().await.remove(&event_key);
self.ws
.leave_room(room_id)
.await
.map_err(|e| MessageBoxError::WebSocket(e.to_string()))
}
pub async fn subscribe(
&self,
event_key: String,
callback: Arc<dyn Fn(AuthenticatedPeerMessage) + Send + Sync>,
) {
self.subscriptions.lock().await.insert(event_key, callback);
}
pub async fn emit_send_message(
&self,
payload: serde_json::Value,
ack_key: String,
ack_tx: oneshot::Sender<bool>,
) -> Result<(), MessageBoxError> {
self.pending_acks
.lock()
.await
.entry(ack_key.clone())
.or_default()
.push_back(ack_tx);
if let Err(e) = self.ws.emit("sendMessage", &payload).await {
let mut guard = self.pending_acks.lock().await;
if let Some(queue) = guard.get_mut(&ack_key) {
queue.pop_back();
if queue.is_empty() {
guard.remove(&ack_key);
}
}
return Err(MessageBoxError::WebSocket(e.to_string()));
}
Ok(())
}
pub async fn remove_pending_ack(&self, key: &str) {
let mut guard = self.pending_acks.lock().await;
if let Some(queue) = guard.get_mut(key) {
queue.pop_front();
if queue.is_empty() {
guard.remove(key);
}
}
}
pub async fn disconnect(&self) -> Result<(), MessageBoxError> {
{
let mut guard = self.pending_acks.lock().await;
for (_, queue) in guard.drain() {
for tx in queue {
let _ = tx.send(false);
}
}
}
self.subscriptions.lock().await.clear();
self.joined_rooms.lock().await.clear();
self.ws
.disconnect()
.await
.map_err(|e| MessageBoxError::WebSocket(e.to_string()))
}
}
fn split_room_id(room_id: &str) -> (String, String) {
const HEX_KEY_LEN: usize = 66;
if room_id.len() > HEX_KEY_LEN && room_id.as_bytes()[HEX_KEY_LEN] == b'-' {
let key = room_id[..HEX_KEY_LEN].to_string();
let mb = room_id[HEX_KEY_LEN + 1..].to_string();
return (key, mb);
}
if let Some(pos) = room_id.find('-') {
(room_id[..pos].to_string(), room_id[pos + 1..].to_string())
} else {
(room_id.to_string(), String::new())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::{WsSendMessageData, WsSendMessagePayload};
use serde_json::json;
use std::collections::HashSet;
#[test]
fn room_id_format() {
let my_key = "03abc";
let recipient = "03def";
let message_box = "payment_inbox";
let listen_room = format!("{my_key}-{message_box}");
let send_room = format!("{recipient}-{message_box}");
assert_eq!(listen_room, "03abc-payment_inbox");
assert_eq!(send_room, "03def-payment_inbox");
assert_ne!(listen_room, send_room);
}
#[test]
fn send_message_data_serializes_camel_case() {
let data = WsSendMessageData {
room_id: "03abc-payment_inbox".to_string(),
message: WsSendMessagePayload {
message_id: "deadbeef".to_string(),
recipient: "03abc".to_string(),
body: "encrypted".to_string(),
},
};
let json = serde_json::to_string(&data).unwrap();
assert!(json.contains("\"roomId\""), "roomId field name");
assert!(json.contains("\"messageId\""), "messageId field name");
assert!(!json.contains("room_id"), "no snake_case leakage");
assert!(!json.contains("message_id"), "no snake_case leakage");
}
#[test]
fn send_message_payload_round_trip() {
let payload = WsSendMessagePayload {
message_id: "abc123".to_string(),
recipient: "03def456".to_string(),
body: r#"{"encryptedMessage":"abc=="}"#.to_string(),
};
let json = serde_json::to_string(&payload).unwrap();
let back: WsSendMessagePayload = serde_json::from_str(&json).unwrap();
assert_eq!(back.message_id, "abc123");
assert_eq!(back.recipient, "03def456");
assert_eq!(back.body, r#"{"encryptedMessage":"abc=="}"#);
}
#[test]
fn authenticated_event_format() {
let identity_key = "03abcdef1234567890";
let v = json!({"identityKey": identity_key});
let json = serde_json::to_string(&v).unwrap();
assert_eq!(json, r#"{"identityKey":"03abcdef1234567890"}"#);
}
#[test]
fn join_room_idempotency_uses_hashset() {
let mut rooms: HashSet<String> = HashSet::new();
let room_id = "03abc-payment_inbox";
let first = rooms.insert(room_id.to_string());
let second = rooms.insert(room_id.to_string());
assert!(first, "first insert returns true");
assert!(!second, "second insert returns false (already present)");
assert_eq!(rooms.len(), 1, "only one entry in set");
}
#[test]
fn split_room_id_hex_key() {
let key = "a".repeat(66);
let mb = "my_inbox";
let room_id = format!("{key}-{mb}");
let (got_key, got_mb) = split_room_id(&room_id);
assert_eq!(got_key, key);
assert_eq!(got_mb, mb);
}
#[test]
fn split_room_id_mb_with_hyphen() {
let key = "b".repeat(66);
let mb = "payment-inbox-v2";
let room_id = format!("{key}-{mb}");
let (got_key, got_mb) = split_room_id(&room_id);
assert_eq!(got_key, key);
assert_eq!(got_mb, mb);
}
#[tokio::test]
async fn fifo_acks_resolve_concurrent_same_room_sends_in_order() {
let acks: PendingAcks = Arc::new(Mutex::new(HashMap::new()));
let key = "sendMessageAck-03abc-inbox".to_string();
let (tx1, rx1) = oneshot::channel::<bool>();
let (tx2, rx2) = oneshot::channel::<bool>();
{
let mut g = acks.lock().await;
g.entry(key.clone()).or_default().push_back(tx1);
g.entry(key.clone()).or_default().push_back(tx2);
assert_eq!(g.get(&key).unwrap().len(), 2, "both waiters queued");
}
{
let mut g = acks.lock().await;
let q = g.get_mut(&key).unwrap();
let _ = q.pop_front().unwrap().send(true);
assert!(!q.is_empty(), "second waiter still queued");
}
assert!(rx1.await.unwrap(), "first send resolved by first ack");
{
let mut g = acks.lock().await;
let q = g.get_mut(&key).unwrap();
let _ = q.pop_front().unwrap().send(false);
if q.is_empty() {
g.remove(&key);
}
assert!(!g.contains_key(&key), "key removed once queue drains");
}
assert!(!rx2.await.unwrap(), "second send resolved by second ack");
}
#[tokio::test]
async fn remove_pending_ack_pops_oldest_and_clears_empty_key() {
let acks: PendingAcks = Arc::new(Mutex::new(HashMap::new()));
let key = "sendMessageAck-03abc-inbox".to_string();
let (tx, _rx) = oneshot::channel::<bool>();
acks.lock().await.entry(key.clone()).or_default().push_back(tx);
{
let mut g = acks.lock().await;
if let Some(q) = g.get_mut(&key) {
q.pop_front();
if q.is_empty() {
g.remove(&key);
}
}
}
assert!(acks.lock().await.is_empty(), "no leaked ack entries");
}
#[tokio::test]
async fn acks_are_per_room_independent() {
let acks: PendingAcks = Arc::new(Mutex::new(HashMap::new()));
let key_a = "sendMessageAck-03aaa-inbox".to_string();
let key_b = "sendMessageAck-03bbb-inbox".to_string();
let (tx_a, rx_a) = oneshot::channel::<bool>();
let (tx_b, rx_b) = oneshot::channel::<bool>();
{
let mut g = acks.lock().await;
g.entry(key_a.clone()).or_default().push_back(tx_a);
g.entry(key_b.clone()).or_default().push_back(tx_b);
}
{
let mut g = acks.lock().await;
let q = g.get_mut(&key_a).unwrap();
let _ = q.pop_front().unwrap().send(true);
if q.is_empty() {
g.remove(&key_a);
}
}
assert!(rx_a.await.unwrap(), "room A resolved");
assert!(acks.lock().await.contains_key(&key_b), "room B still pending");
drop(rx_b);
}
}