1#![warn(missing_docs)]
14
15use std::collections::{VecDeque, HashMap};
16use std::time::{Duration, Instant};
17
18#[derive(Debug, Clone)]
22pub enum WsMessage {
23 Text(String),
25 Binary(Vec<u8>),
27 Close {
29 code: u16,
31 reason: String,
33 },
34 Ping(Vec<u8>),
36 Pong(Vec<u8>),
38}
39
40impl WsMessage {
41 pub fn text(s: impl Into<String>) -> Self { Self::Text(s.into()) }
43 pub fn binary(v: Vec<u8>) -> Self { Self::Binary(v) }
45 pub fn close_normal() -> Self { Self::Close { code: 1000, reason: "Normal closure".into() } }
47
48 pub fn is_data(&self) -> bool {
50 matches!(self, Self::Text(_) | Self::Binary(_))
51 }
52
53 pub fn as_text(&self) -> Option<&str> {
55 if let Self::Text(s) = self { Some(s) } else { None }
56 }
57
58 pub fn len(&self) -> usize {
60 match self {
61 Self::Text(s) => s.len(),
62 Self::Binary(v) => v.len(),
63 _ => 0,
64 }
65 }
66
67 pub fn is_empty(&self) -> bool { self.len() == 0 }
69}
70
71#[derive(Debug, Clone, Copy, PartialEq, Eq)]
75pub enum WsState {
76 Disconnected,
78 Connecting,
80 Handshaking,
82 Connected,
84 ReconnectBackoff,
86 Closed,
88}
89
90impl WsState {
91 pub fn is_connected(self) -> bool { self == Self::Connected }
93 pub fn is_live(self) -> bool { matches!(self, Self::Connected | Self::Handshaking) }
95}
96
97#[derive(Debug, Clone)]
101pub enum WsEvent {
102 Connected {
104 url: String,
106 },
107 Disconnected {
109 url: String,
111 code: u16,
113 reason: String,
115 will_reconnect: bool,
117 },
118 Message {
120 message: WsMessage,
122 channel: Option<String>,
124 },
125 Reconnecting {
127 url: String,
129 attempt: u32,
131 backoff_ms: u64,
133 },
134 Error {
136 description: String,
138 },
139 PingRtt {
141 millis: f64,
143 },
144}
145
146#[derive(Debug, Clone)]
149struct OutboundMessage {
150 message: WsMessage,
151 channel: Option<String>,
152 queued_at: Instant,
153}
154
155enum WorkerEvent {
159 Connected,
160 Message(WsMessage),
161 Pong,
162 Closed { code: u16, reason: String },
163 Failed(String),
164}
165
166struct Worker {
168 outbound: std::sync::mpsc::Sender<WsMessage>,
169 inbound: std::sync::mpsc::Receiver<WorkerEvent>,
170}
171
172pub struct WsClient {
184 pub url: String,
186 pub state: WsState,
188 pub auto_reconnect: bool,
190 pub max_reconnects: u32,
192 pub ping_interval: Duration,
194 pub pong_timeout: Duration,
196 pub max_queue_size: usize,
198
199 reconnect_attempt: u32,
200 reconnect_timer: f32,
201 reconnect_backoff: f32,
202 last_ping: Option<Instant>,
203 awaiting_pong: bool,
204 ping_payload: Vec<u8>,
205 outbound: VecDeque<OutboundMessage>,
206 events: VecDeque<WsEvent>,
207 channels: HashMap<String, ChannelConfig>,
209 connect_time: Option<Instant>,
210 messages_sent: u64,
211 messages_received: u64,
212 bytes_sent: u64,
213 bytes_received: u64,
214 worker: Option<Worker>,
215}
216
217#[derive(Debug, Clone)]
219pub struct ChannelConfig {
220 pub name: String,
222 pub filter: Option<String>,
224 pub active: bool,
226}
227
228impl WsClient {
229 pub fn new(url: impl Into<String>) -> Self {
231 Self {
232 url: url.into(),
233 state: WsState::Disconnected,
234 auto_reconnect: true,
235 max_reconnects: 10,
236 ping_interval: Duration::from_secs(30),
237 pong_timeout: Duration::from_secs(10),
238 max_queue_size: 1024,
239 reconnect_attempt: 0,
240 reconnect_timer: 0.0,
241 reconnect_backoff: 1.0,
242 last_ping: None,
243 awaiting_pong: false,
244 ping_payload: vec![1, 2, 3, 4],
245 outbound: VecDeque::new(),
246 events: VecDeque::new(),
247 channels: HashMap::new(),
248 connect_time: None,
249 messages_sent: 0,
250 messages_received: 0,
251 bytes_sent: 0,
252 bytes_received: 0,
253 worker: None,
254 }
255 }
256
257 pub fn connect(&mut self) {
261 if matches!(self.state, WsState::Disconnected | WsState::Closed) {
262 self.state = WsState::Connecting;
263 self.reconnect_attempt = 0;
264 }
265 }
266
267 pub fn close(&mut self) {
269 if self.state != WsState::Closed {
270 if let Some(w) = &self.worker {
271 let _ = w.outbound.send(WsMessage::close_normal());
272 }
273 self.worker = None;
274 self.state = WsState::Closed;
275 }
276 }
277
278 pub fn disconnect(&mut self) {
281 self.worker = None;
282 self.state = WsState::Disconnected;
283 self.reconnect_attempt = self.reconnect_attempt.max(1);
284 self.reconnect_timer = 0.0;
285 }
286
287 pub fn send(&mut self, message: WsMessage) -> bool {
291 self.send_raw(message, None)
292 }
293
294 pub fn send_on_channel(&mut self, channel: &str, message: WsMessage) -> bool {
296 self.send_raw(message, Some(channel.to_owned()))
297 }
298
299 fn send_raw(&mut self, message: WsMessage, channel: Option<String>) -> bool {
300 if self.outbound.len() >= self.max_queue_size {
301 self.events.push_back(WsEvent::Error {
302 description: "Outbound queue full, message dropped".into(),
303 });
304 return false;
305 }
306 self.outbound.push_back(OutboundMessage {
307 message,
308 channel,
309 queued_at: Instant::now(),
310 });
311 true
312 }
313
314 pub fn subscribe(&mut self, channel: impl Into<String>) {
318 let name = channel.into();
319 self.channels.insert(name.clone(), ChannelConfig {
320 name,
321 filter: None,
322 active: true,
323 });
324 }
325
326 pub fn unsubscribe(&mut self, channel: &str) {
328 self.channels.remove(channel);
329 }
330
331 pub fn tick(&mut self, dt: f32) {
335 match self.state {
336 WsState::Disconnected => {
337 if self.auto_reconnect
338 && self.reconnect_attempt > 0
339 && (self.max_reconnects == 0 || self.reconnect_attempt < self.max_reconnects)
340 {
341 self.reconnect_timer -= dt;
342 if self.reconnect_timer <= 0.0 {
343 self.state = WsState::Connecting;
344 }
345 }
346 }
347
348 WsState::Connecting => match spawn_worker(&self.url) {
349 Ok(worker) => {
350 self.worker = Some(worker);
351 self.state = WsState::Handshaking;
352 }
353 Err(description) => {
354 self.events.push_back(WsEvent::Error { description: description.clone() });
355 if cfg!(feature = "websocket") {
356 self.handle_disconnect(1006, description);
357 } else {
358 self.state = WsState::Closed;
360 }
361 }
362 },
363
364 WsState::Handshaking | WsState::Connected => {
365 self.poll_worker();
366 if self.state != WsState::Connected {
367 return;
368 }
369 if let Some(w) = &self.worker {
371 while let Some(msg) = self.outbound.pop_front() {
372 let len = msg.message.len() as u64;
373 if w.outbound.send(msg.message).is_err() {
374 break;
375 }
376 self.messages_sent += 1;
377 self.bytes_sent += len;
378 }
379 }
380 let due = self.last_ping.is_none_or(|t| t.elapsed() >= self.ping_interval);
382 if due && !self.awaiting_pong {
383 if let Some(w) = &self.worker {
384 let _ = w.outbound.send(WsMessage::Ping(self.ping_payload.clone()));
385 }
386 self.last_ping = Some(Instant::now());
387 self.awaiting_pong = true;
388 }
389 if self.awaiting_pong
390 && self.last_ping.is_some_and(|t| t.elapsed() > self.pong_timeout)
391 {
392 self.handle_disconnect(1001, "Pong timeout".into());
393 }
394 }
395
396 WsState::ReconnectBackoff => {
397 self.reconnect_timer -= dt;
398 if self.reconnect_timer <= 0.0 {
399 self.events.push_back(WsEvent::Reconnecting {
400 url: self.url.clone(),
401 attempt: self.reconnect_attempt,
402 backoff_ms: (self.reconnect_backoff * 1000.0) as u64,
403 });
404 self.state = WsState::Connecting;
405 }
406 }
407
408 WsState::Closed => {}
409 }
410 }
411
412 fn poll_worker(&mut self) {
413 loop {
414 let Some(w) = &self.worker else { return };
415 let ev = match w.inbound.try_recv() {
416 Ok(ev) => ev,
417 Err(std::sync::mpsc::TryRecvError::Empty) => return,
418 Err(std::sync::mpsc::TryRecvError::Disconnected) => {
419 WorkerEvent::Failed("connection thread ended".into())
420 }
421 };
422 match ev {
423 WorkerEvent::Connected => {
424 self.state = WsState::Connected;
425 self.connect_time = Some(Instant::now());
426 self.reconnect_attempt = 0;
427 self.reconnect_backoff = 1.0;
428 self.last_ping = Some(Instant::now());
429 self.awaiting_pong = false;
430 self.events.push_back(WsEvent::Connected { url: self.url.clone() });
431 }
432 WorkerEvent::Message(message) => {
433 self.messages_received += 1;
434 self.bytes_received += message.len() as u64;
435 self.events.push_back(WsEvent::Message { message, channel: None });
436 }
437 WorkerEvent::Pong => {
438 if let Some(t) = self.last_ping {
439 self.events.push_back(WsEvent::PingRtt {
440 millis: t.elapsed().as_secs_f64() * 1000.0,
441 });
442 }
443 self.awaiting_pong = false;
444 }
445 WorkerEvent::Closed { code, reason } => {
446 self.handle_disconnect(code, reason);
447 return;
448 }
449 WorkerEvent::Failed(reason) => {
450 self.events.push_back(WsEvent::Error { description: reason.clone() });
451 self.handle_disconnect(1006, reason);
452 return;
453 }
454 }
455 }
456 }
457
458 fn handle_disconnect(&mut self, code: u16, reason: String) {
459 self.worker = None;
460 self.connect_time = None;
461 self.awaiting_pong = false;
462 let will_reconnect = self.auto_reconnect
463 && (self.max_reconnects == 0 || self.reconnect_attempt < self.max_reconnects);
464
465 self.events.push_back(WsEvent::Disconnected {
466 url: self.url.clone(),
467 code,
468 reason,
469 will_reconnect,
470 });
471
472 if will_reconnect {
473 self.reconnect_attempt += 1;
474 self.reconnect_timer = self.reconnect_backoff;
475 self.reconnect_backoff = (self.reconnect_backoff * 2.0).min(60.0);
477 self.state = WsState::ReconnectBackoff;
478 } else {
479 self.state = WsState::Closed;
480 }
481 }
482
483 pub fn drain_events(&mut self) -> impl Iterator<Item = WsEvent> + '_ {
487 self.events.drain(..)
488 }
489
490 pub fn is_connected(&self) -> bool { self.state.is_connected() }
492 pub fn messages_sent(&self) -> u64 { self.messages_sent }
494 pub fn messages_received(&self) -> u64 { self.messages_received }
496 pub fn bytes_sent(&self) -> u64 { self.bytes_sent }
498 pub fn bytes_received(&self) -> u64 { self.bytes_received }
500 pub fn uptime(&self) -> Option<Duration> { self.connect_time.map(|t| t.elapsed()) }
502 pub fn pending_outbound(&self) -> usize { self.outbound.len() }
504}
505
506#[cfg(feature = "websocket")]
509fn spawn_worker(url: &str) -> Result<Worker, String> {
510 use std::sync::mpsc::{channel, TryRecvError};
511 use tungstenite::{stream::MaybeTlsStream, Message};
512
513 let (out_tx, out_rx) = channel::<WsMessage>();
514 let (in_tx, in_rx) = channel::<WorkerEvent>();
515 let url = url.to_owned();
516 std::thread::Builder::new()
517 .name("proof-websocket".into())
518 .spawn(move || {
519 let (mut socket, _resp) = match tungstenite::connect(url.as_str()) {
520 Ok(s) => s,
521 Err(e) => {
522 let _ = in_tx.send(WorkerEvent::Failed(format!("connect {url}: {e}")));
523 return;
524 }
525 };
526 let timeout = Some(Duration::from_millis(15));
528 match socket.get_mut() {
529 MaybeTlsStream::Plain(s) => { let _ = s.set_read_timeout(timeout); }
530 MaybeTlsStream::Rustls(s) => { let _ = s.get_mut().set_read_timeout(timeout); }
531 _ => {}
532 }
533 if in_tx.send(WorkerEvent::Connected).is_err() {
534 return;
535 }
536 loop {
537 loop {
539 match out_rx.try_recv() {
540 Ok(msg) => {
541 let m = match msg {
542 WsMessage::Text(t) => Message::text(t),
543 WsMessage::Binary(b) => Message::binary(b),
544 WsMessage::Ping(p) => Message::Ping(p.into()),
545 WsMessage::Pong(p) => Message::Pong(p.into()),
546 WsMessage::Close { code, reason } => {
547 let _ = socket.close(Some(tungstenite::protocol::CloseFrame {
548 code: code.into(),
549 reason: reason.into(),
550 }));
551 let _ = socket.flush();
552 return;
553 }
554 };
555 if let Err(e) = socket.send(m) {
556 let _ = in_tx.send(WorkerEvent::Failed(format!("send: {e}")));
557 return;
558 }
559 }
560 Err(TryRecvError::Empty) => break,
561 Err(TryRecvError::Disconnected) => {
563 let _ = socket.close(None);
564 let _ = socket.flush();
565 return;
566 }
567 }
568 }
569 match socket.read() {
571 Ok(Message::Text(t)) => {
572 let _ = in_tx.send(WorkerEvent::Message(WsMessage::Text(t.to_string())));
573 }
574 Ok(Message::Binary(b)) => {
575 let _ = in_tx.send(WorkerEvent::Message(WsMessage::Binary(b.to_vec())));
576 }
577 Ok(Message::Pong(_)) => {
578 let _ = in_tx.send(WorkerEvent::Pong);
579 }
580 Ok(Message::Ping(_)) | Ok(Message::Frame(_)) => {} Ok(Message::Close(frame)) => {
582 let (code, reason) = frame
583 .map(|f| (u16::from(f.code), f.reason.to_string()))
584 .unwrap_or((1005, String::new()));
585 let _ = in_tx.send(WorkerEvent::Closed { code, reason });
586 return;
587 }
588 Err(tungstenite::Error::Io(e))
589 if matches!(e.kind(), std::io::ErrorKind::WouldBlock | std::io::ErrorKind::TimedOut) => {}
590 Err(tungstenite::Error::ConnectionClosed) | Err(tungstenite::Error::AlreadyClosed) => {
591 let _ = in_tx.send(WorkerEvent::Closed { code: 1000, reason: String::new() });
592 return;
593 }
594 Err(e) => {
595 let _ = in_tx.send(WorkerEvent::Failed(format!("read: {e}")));
596 return;
597 }
598 }
599 }
600 })
601 .map_err(|e| format!("could not start the websocket thread: {e}"))?;
602 Ok(Worker { outbound: out_tx, inbound: in_rx })
603}
604
605#[cfg(not(feature = "websocket"))]
606fn spawn_worker(_url: &str) -> Result<Worker, String> {
607 Err("proof-engine was built without the `websocket` feature, so it cannot open connections".into())
608}
609
610#[cfg(test)]
611mod tests {
612 use super::*;
613
614 fn pump(c: &mut WsClient, until: impl Fn(&[WsEvent]) -> bool) -> Vec<WsEvent> {
615 let deadline = Instant::now() + Duration::from_secs(10);
616 let mut seen = Vec::new();
617 while Instant::now() < deadline {
618 c.tick(0.016);
619 seen.extend(c.drain_events());
620 if until(&seen) {
621 break;
622 }
623 std::thread::sleep(Duration::from_millis(5));
624 }
625 seen
626 }
627
628 #[cfg(not(feature = "websocket"))]
629 #[test]
630 fn without_feature_connect_reports_an_error_instead_of_pretending() {
631 let mut c = WsClient::new("ws://127.0.0.1:9/");
632 c.connect();
633 let ev = pump(&mut c, |e| !e.is_empty());
634 assert!(matches!(ev.first(), Some(WsEvent::Error { .. })), "{ev:?}");
635 assert_eq!(c.state, WsState::Closed);
636 assert!(!c.is_connected());
637 }
638
639 #[test]
640 fn new_client_does_not_connect_by_itself() {
641 let mut c = WsClient::new("ws://127.0.0.1:9/");
642 for _ in 0..5 {
643 c.tick(1.0);
644 }
645 assert_eq!(c.state, WsState::Disconnected);
646 }
647
648 #[cfg(feature = "websocket")]
649 #[test]
650 fn echo_round_trip_against_a_local_server() {
651 let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
652 let addr = listener.local_addr().unwrap();
653 std::thread::spawn(move || {
654 let (stream, _) = listener.accept().unwrap();
655 let mut ws = tungstenite::accept(stream).unwrap();
656 loop {
657 match ws.read() {
658 Ok(m) if m.is_text() || m.is_binary() => {
659 ws.send(m).unwrap();
660 }
661 Ok(tungstenite::Message::Close(_)) | Err(_) => break,
662 Ok(_) => {}
663 }
664 }
665 });
666
667 let mut c = WsClient::new(format!("ws://{addr}/"));
668 c.ping_interval = Duration::from_millis(50);
669 c.connect();
670 c.send(WsMessage::text("hello"));
671 c.send(WsMessage::binary(vec![1, 2, 3]));
672 let ev = pump(&mut c, |e| {
673 e.iter().filter(|x| matches!(x, WsEvent::Message { .. })).count() >= 2
674 && e.iter().any(|x| matches!(x, WsEvent::PingRtt { .. }))
675 });
676 assert!(matches!(ev.first(), Some(WsEvent::Connected { .. })), "{ev:?}");
677 let texts: Vec<_> = ev
678 .iter()
679 .filter_map(|x| match x {
680 WsEvent::Message { message, .. } => Some(format!("{message:?}")),
681 _ => None,
682 })
683 .collect();
684 assert_eq!(texts, vec!["Text(\"hello\")".to_string(), "Binary([1, 2, 3])".to_string()]);
685 assert!(ev.iter().any(|x| matches!(x, WsEvent::PingRtt { .. })), "pong received");
686 assert_eq!(c.messages_sent(), 2);
687 assert_eq!(c.messages_received(), 2);
688 c.close();
689 assert_eq!(c.state, WsState::Closed);
690 }
691
692 #[cfg(feature = "websocket")]
693 #[test]
694 fn refused_connection_backs_off_and_gives_up() {
695 let port = std::net::TcpListener::bind("127.0.0.1:0").unwrap().local_addr().unwrap().port();
696 let mut c = WsClient::new(format!("ws://127.0.0.1:{port}/"));
697 c.max_reconnects = 1;
698 c.connect();
699 let ev = pump(&mut c, |e| {
700 e.iter().any(|x| matches!(x, WsEvent::Disconnected { will_reconnect: true, .. }))
701 });
702 assert!(ev.iter().any(|x| matches!(x, WsEvent::Error { .. })), "{ev:?}");
703 assert_eq!(c.state, WsState::ReconnectBackoff);
704 }
705}