use anyhow::{bail, Result};
use futures_util::{stream::SplitSink, SinkExt, StreamExt};
use serde_json::{json, Value};
use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use tokio::net::TcpStream;
use tokio::sync::{mpsc, oneshot, Mutex};
use tokio_tungstenite::{
connect_async, MaybeTlsStream, WebSocketStream,
};
use tokio_tungstenite::tungstenite::Message;
type WsSink = SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>;
struct SioPacket {
sio_type: char,
ack_id: Option<u64>,
data: Value,
}
#[derive(Debug)]
pub enum Inbound {
Event {
event: String,
payload: Value,
ack_id: Option<u64>,
},
ConnectOk,
ConnectError(String),
Closed,
}
pub struct SioClient {
sink: Arc<Mutex<WsSink>>,
acks: Arc<Mutex<HashMap<u64, oneshot::Sender<Value>>>>,
next_ack: Arc<AtomicU64>,
}
impl SioClient {
async fn send_raw(&self, s: String) -> Result<()> {
let mut g = self.sink.lock().await;
g.send(Message::Text(s)).await?;
Ok(())
}
pub async fn emit(&self, event: &str, payload: Value) -> Result<()> {
let body = serde_json::to_string(&json!([event, payload]))?;
self.send_raw(format!("42/cli,{body}")).await
}
pub async fn emit_with_ack(
&self,
event: &str,
payload: Value,
timeout: Duration,
) -> Result<Value> {
let id = self.next_ack.fetch_add(1, Ordering::SeqCst) + 1;
let (tx, rx) = oneshot::channel();
self.acks.lock().await.insert(id, tx);
let body = serde_json::to_string(&json!([event, payload]))?;
self.send_raw(format!("42/cli,{id}{body}")).await?;
match tokio::time::timeout(timeout, rx).await {
Ok(Ok(v)) => Ok(v),
_ => {
self.acks.lock().await.remove(&id);
bail!("ack timeout/error for {event}")
}
}
}
pub async fn send_ack(&self, ack_id: u64, value: Value) -> Result<()> {
let body = serde_json::to_string(&json!([value]))?;
self.send_raw(format!("43/cli,{ack_id}{body}")).await
}
}
pub async fn connect(
ws_url: &str,
auth: Value,
) -> Result<(SioClient, mpsc::UnboundedReceiver<Inbound>, oneshot::Receiver<Result<()>>)> {
log::info!("socket.io connecting: {ws_url}");
let (ws, _resp) = connect_async(ws_url).await?;
let (sink, mut stream) = ws.split();
let sink = Arc::new(Mutex::new(sink));
let acks: Arc<Mutex<HashMap<u64, oneshot::Sender<Value>>>> =
Arc::new(Mutex::new(HashMap::new()));
let next_ack = Arc::new(AtomicU64::new(0));
while let Some(msg) = stream.next().await {
match msg {
Ok(Message::Text(t)) if t.starts_with('0') => {
log::debug!("eio open: {}", truncate(&t, 200));
break;
}
Ok(_) => { }
Err(e) => bail!("socket.io pre-open error: {e}"),
}
}
let auth_str = serde_json::to_string(&auth)?;
{
let mut g = sink.lock().await;
g.send(Message::Text(format!("40/cli,{auth_str}"))).await?;
}
let (event_tx, event_rx) = mpsc::unbounded_channel();
let (connect_tx, connect_rx) = oneshot::channel::<Result<()>>();
let connect_tx = Arc::new(Mutex::new(Some(connect_tx)));
let reader_sink = Arc::clone(&sink);
let reader_acks = Arc::clone(&acks);
tokio::spawn(async move {
let mut connect_tx = connect_tx;
while let Some(msg) = stream.next().await {
match msg {
Ok(Message::Text(t)) => {
if !handle_frame(
&t,
&reader_sink,
&reader_acks,
&event_tx,
&mut connect_tx,
)
.await
{
break;
}
}
Ok(Message::Ping(p)) => {
let _ = reader_sink.lock().await.send(Message::Pong(p)).await;
}
Ok(Message::Close(_)) => {
let _ = event_tx.send(Inbound::Closed);
break;
}
Ok(_) => {}
Err(e) => {
log::warn!("ws read error: {e}");
let _ = event_tx.send(Inbound::Closed);
break;
}
}
}
log::info!("socket.io reader exited");
});
Ok((
SioClient {
sink,
acks,
next_ack,
},
event_rx,
connect_rx,
))
}
async fn handle_frame(
t: &str,
sink: &Arc<Mutex<WsSink>>,
acks: &Arc<Mutex<HashMap<u64, oneshot::Sender<Value>>>>,
event_tx: &mpsc::UnboundedSender<Inbound>,
connect_tx: &mut Arc<Mutex<Option<oneshot::Sender<Result<()>>>>>,
) -> bool {
let mut chars = t.chars();
let Some(eio) = chars.next() else {
return true;
};
match eio {
'0' => { }
'1' => {
let _ = event_tx.send(Inbound::Closed);
return false;
}
'2' => {
let _ = sink.lock().await.send(Message::Text("3".to_string())).await;
}
'3' => { }
'4' => {
let rest = &t[1..];
if let Some(pkt) = parse_sio(rest) {
match pkt.sio_type {
'0' => {
if let Some(tx) = connect_tx.lock().await.take() {
let _ = tx.send(Ok(()));
}
let _ = event_tx.send(Inbound::ConnectOk);
}
'4' => {
let msg = if pkt.data.is_string() {
pkt.data.as_str().unwrap().to_string()
} else {
serde_json::to_string(&pkt.data).unwrap_or_default()
};
log::error!("socket.io connect_error: {msg}");
if let Some(tx) = connect_tx.lock().await.take() {
let _ = tx.send(Err(anyhow::anyhow!("connect_error: {msg}")));
}
return false;
}
'1' => {
let _ = event_tx.send(Inbound::Closed);
return false;
}
'2' => {
if let Some(arr) = pkt.data.as_array() {
let event = arr
.first()
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let payload = arr.get(1).cloned().unwrap_or(Value::Null);
let _ = event_tx.send(Inbound::Event {
event,
payload,
ack_id: pkt.ack_id,
});
}
}
'3' => {
if let Some(id) = pkt.ack_id {
let value = pkt
.data
.as_array()
.and_then(|a| a.first().cloned())
.unwrap_or(Value::Null);
if let Some(tx) = acks.lock().await.remove(&id) {
let _ = tx.send(value);
}
}
}
_ => {
log::debug!("ignoring sio type {} rest={}", pkt.sio_type, truncate(rest, 160));
}
}
}
}
_ => {
log::debug!("ignoring eio frame: {}", truncate(t, 160));
}
}
true
}
fn parse_sio(rest: &str) -> Option<SioPacket> {
let mut it = rest.chars();
let sio_type = it.next()?;
let tail = &rest[sio_type.len_utf8()..];
let (tail, _namespace) = if let Some(stripped) = tail.strip_prefix('/') {
let comma = stripped.find(',')?;
let ns = &stripped[..comma];
(&stripped[comma + 1..], Some(ns))
} else {
(tail, None)
};
let (tail, ack_id) = match tail.as_bytes().first() {
Some(b) if b.is_ascii_digit() => {
let end = tail
.find(|c: char| !c.is_ascii_digit())
.unwrap_or(tail.len());
let id: u64 = tail[..end].parse().ok()?;
(&tail[end..], Some(id))
}
_ => (tail, None),
};
let data: Value = if tail.is_empty() {
Value::Null
} else {
serde_json::from_str(tail).unwrap_or(Value::Null)
};
Some(SioPacket {
sio_type,
ack_id,
data,
})
}
fn truncate(s: &str, n: usize) -> String {
if s.len() <= n {
s.to_string()
} else {
format!("{}…", &s[..n])
}
}
#[allow(dead_code)]
type _UnusedStream = tokio_tungstenite::WebSocketStream<MaybeTlsStream<TcpStream>>;