1use crate::types::{
2 Event, EventFilter, EventTopic, SubscriptionErrorResponse, SubscriptionResponse,
3 SubscriptionStatus, UnsubscribeResponse, WebSocketRequest,
4};
5use crate::{SubscriptionRequest, UnsubscribeRequest};
6use futures::future::{BoxFuture, FutureExt};
7use futures::{SinkExt, StreamExt};
8use serde::Deserialize;
9use std::collections::HashMap;
10use std::future::Future;
11use std::panic::{self, AssertUnwindSafe};
12use std::sync::Arc;
13use std::time::Duration;
14use tokio::net::TcpStream;
15use tokio::sync::{mpsc, Mutex, RwLock};
16use tokio_tungstenite::tungstenite::Bytes;
17use tokio_tungstenite::{connect_async, tungstenite::protocol::Message, WebSocketStream};
18use tracing::{debug, error, info, warn};
19
20#[derive(Debug, Clone, thiserror::Error)]
22pub enum WebSocketError {
23 #[error("Failed to connect to the server: {0}")]
25 ConnectionFailed(String),
26 #[error("Failed to send a message: {0}")]
28 SendFailed(String),
29 #[error("Failed to parse a response: {0}")]
31 ParseError(String),
32 #[error("Failed to subscribe: {0}")]
34 SubscriptionFailed(String),
35 #[error("Failed to unsubscribe: {0}")]
37 UnsubscriptionFailed(String),
38 #[error("Failed to read from the WebSocket: {0}")]
40 ReadFailed(String),
41 #[error("Other error: {0}")]
43 Other(String),
44}
45
46pub type EventCallback = Box<dyn Fn(Event) + Send + Sync + 'static>;
48
49pub type ConnectionCallback = Box<dyn Fn(bool) + Send + Sync + 'static>;
51
52pub type AsyncEventCallback = Box<dyn Fn(Event) -> BoxFuture<'static, ()> + Send + Sync + 'static>;
54
55struct SubscriptionHandler {
57 topic: EventTopic,
59 filter: EventFilter,
61 pending: bool,
63}
64
65#[derive(Debug, Deserialize)]
67#[serde(untagged)]
68pub enum WebSocketMessage {
69 Event(Event),
71 SubscriptionResponse(SubscriptionResponse),
73 UnsubscribeResponse(UnsubscribeResponse),
75 ErrorResponse(SubscriptionErrorResponse),
77}
78
79#[derive(Clone, Debug)]
81pub enum BackoffStrategy {
82 Constant(Duration),
84
85 Linear { initial: Duration, step: Duration },
87
88 Exponential {
91 initial: Duration,
92 factor: f32,
93 max_delay: Duration,
94 jitter: f32, },
96}
97
98impl BackoffStrategy {
99 pub fn default_exponential() -> Self {
101 Self::Exponential {
102 initial: Duration::from_secs(1),
103 factor: 2.0,
104 max_delay: Duration::from_secs(5),
105 jitter: 0.1,
106 }
107 }
108
109 pub fn next_delay(&self, attempt: usize) -> Duration {
111 match self {
112 Self::Constant(duration) => *duration,
113
114 Self::Linear { initial, step } => *initial + (*step * attempt as u32),
115
116 Self::Exponential {
117 initial,
118 factor,
119 max_delay,
120 jitter,
121 } => {
122 let base_ms = initial.as_millis() as f32 * factor.powi(attempt as i32);
124
125 let jitter_factor = 1.0 - jitter + rand::random::<f32>() * jitter * 2.0;
127 let jittered_ms = base_ms * jitter_factor;
128
129 let capped_ms = jittered_ms.min(max_delay.as_millis() as f32);
131
132 Duration::from_millis(capped_ms as u64)
133 }
134 }
135 }
136}
137
138pub struct WebSocketClient {
140 sender: mpsc::Sender<Message>,
142 subscriptions: Arc<RwLock<HashMap<String, SubscriptionHandler>>>,
144 event_callbacks: Arc<RwLock<HashMap<EventTopic, Vec<EventCallback>>>>,
146 connection_callbacks: Arc<RwLock<Vec<ConnectionCallback>>>,
148 connected: Arc<Mutex<bool>>,
150 server_url: Arc<String>,
152 auto_reconnect: Arc<Mutex<bool>>,
154 backoff_strategy: Arc<Mutex<BackoffStrategy>>,
156 max_reconnect_attempts: Arc<Mutex<usize>>,
158 running: Arc<Mutex<bool>>,
160 cancel_tx: Option<mpsc::Sender<()>>,
162 async_event_callbacks: Arc<RwLock<HashMap<EventTopic, Vec<AsyncEventCallback>>>>,
164 state_change_lock: Arc<Mutex<()>>,
166 keep_alive_interval: Arc<Mutex<Option<Duration>>>,
168 keep_alive_handle: Arc<Mutex<Option<tokio::task::JoinHandle<()>>>>,
170}
171
172impl WebSocketClient {
173 pub fn new(url: &str) -> Self {
175 let subscriptions = Arc::new(RwLock::new(HashMap::new()));
177 let event_callbacks = Arc::new(RwLock::new(HashMap::new()));
178 let connection_callbacks = Arc::new(RwLock::new(Vec::new()));
179 let connected = Arc::new(Mutex::new(false));
180 let server_url = Arc::new(url.to_string());
181 let auto_reconnect = Arc::new(Mutex::new(false));
182 let backoff_strategy = Arc::new(Mutex::new(BackoffStrategy::default_exponential()));
183 let max_reconnect_attempts = Arc::new(Mutex::new(0));
184 let running = Arc::new(Mutex::new(true));
185 let async_event_callbacks = Arc::new(RwLock::new(HashMap::new()));
186 let state_change_lock = Arc::new(Mutex::new(()));
187 let keep_alive_interval = Arc::new(Mutex::new(None));
188 let keep_alive_handle = Arc::new(Mutex::new(None));
189
190 let (sender, _) = mpsc::channel::<Message>(100);
192
193 WebSocketClient {
194 sender,
195 subscriptions,
196 event_callbacks,
197 connection_callbacks,
198 connected,
199 server_url,
200 auto_reconnect,
201 backoff_strategy,
202 max_reconnect_attempts,
203 running,
204 cancel_tx: None,
205 async_event_callbacks,
206 state_change_lock,
207 keep_alive_interval,
208 keep_alive_handle,
209 }
210 }
211
212 pub async fn connect(&mut self) -> Result<(), WebSocketError> {
214 let _connection_lock = self.state_change_lock.lock().await;
216
217 if *self.connected.lock().await {
219 return Ok(());
220 }
221
222 if let Some(cancel_tx) = self.cancel_tx.take() {
224 let _ = cancel_tx.send(()).await;
226 tokio::time::sleep(Duration::from_millis(50)).await;
227 }
228
229 let (cancel_tx, cancel_rx) = mpsc::channel::<()>(1);
231 self.cancel_tx = Some(cancel_tx);
232
233 let connect_result = Self::establish_new_connection(
235 &self.server_url,
236 &self.connected,
237 &self.connection_callbacks,
238 &self.keep_alive_handle,
239 &self.keep_alive_interval,
240 &self.running,
241 )
242 .await;
243
244 if let Err(e) = &connect_result {
246 self.cancel_tx = None;
248 return Err(e.clone());
249 }
250
251 let (read, sender) = connect_result.unwrap();
253
254 self.sender = sender.clone();
256
257 let _task_handle = tokio::spawn(Self::message_processor(
259 read,
260 sender,
261 self.subscriptions.clone(),
262 self.event_callbacks.clone(),
263 self.async_event_callbacks.clone(),
264 self.connection_callbacks.clone(),
265 self.connected.clone(),
266 self.keep_alive_handle.clone(),
267 self.keep_alive_interval.clone(),
268 self.running.clone(),
269 self.server_url.clone(),
270 self.auto_reconnect.clone(),
271 self.backoff_strategy.clone(),
272 self.max_reconnect_attempts.clone(),
273 cancel_rx,
274 ));
275
276 Ok(())
277 }
278
279 pub async fn connect_static(url: &str) -> Result<Self, WebSocketError> {
281 let mut client = Self::new(url);
282 client.connect().await?;
283 Ok(client)
284 }
285
286 async fn establish_new_connection(
289 server_url: &Arc<String>,
290 connected: &Arc<Mutex<bool>>,
291 connection_callbacks: &Arc<RwLock<Vec<ConnectionCallback>>>,
292 keep_alive_handle: &Arc<Mutex<Option<tokio::task::JoinHandle<()>>>>,
293 keep_alive_interval: &Arc<Mutex<Option<Duration>>>,
294 running: &Arc<Mutex<bool>>,
295 ) -> Result<
296 (
297 futures::stream::SplitStream<
298 WebSocketStream<tokio_tungstenite::MaybeTlsStream<TcpStream>>,
299 >,
300 mpsc::Sender<Message>,
301 ),
302 WebSocketError,
303 > {
304 let ws_stream = Self::establish_connection(server_url).await?;
306
307 let (write, read) = ws_stream.split();
309
310 let sender = Self::spawn_writer_task(write);
312
313 if let Some(interval) = *keep_alive_interval.lock().await {
315 Self::restart_keep_alive(keep_alive_handle, &sender, running, connected, interval)
316 .await?;
317 }
318
319 Self::update_connection_status(connected, true, connection_callbacks).await;
321
322 Ok((read, sender))
323 }
324
325 #[allow(clippy::too_many_arguments)]
327 async fn attempt_reconnection(
328 server_url: &Arc<String>,
329 connected: &Arc<Mutex<bool>>,
330 connection_callbacks: &Arc<RwLock<Vec<ConnectionCallback>>>,
331 subscriptions: &Arc<RwLock<HashMap<String, SubscriptionHandler>>>,
332 keep_alive_handle: &Arc<Mutex<Option<tokio::task::JoinHandle<()>>>>,
333 keep_alive_interval: &Arc<Mutex<Option<Duration>>>,
334 running: &Arc<Mutex<bool>>,
335 backoff_strategy: &Arc<Mutex<BackoffStrategy>>,
336 max_reconnect_attempts: &Arc<Mutex<usize>>,
337 reconnect_attempts: &mut usize,
338 cancel_rx: &mut mpsc::Receiver<()>,
339 ) -> Result<
340 (
341 futures::stream::SplitStream<
342 WebSocketStream<tokio_tungstenite::MaybeTlsStream<TcpStream>>,
343 >,
344 mpsc::Sender<Message>,
345 ),
346 WebSocketError,
347 > {
348 loop {
349 let max_attempts = *max_reconnect_attempts.lock().await;
351 if max_attempts > 0 && *reconnect_attempts >= max_attempts {
352 error!("Maximum reconnection attempts reached ({})", max_attempts);
353 return Err(WebSocketError::ConnectionFailed(
354 "Maximum reconnection attempts reached".to_string(),
355 ));
356 }
357
358 let delay = backoff_strategy
360 .lock()
361 .await
362 .next_delay(*reconnect_attempts);
363 tokio::time::sleep(delay).await;
364
365 *reconnect_attempts += 1;
366 info!(
367 "Attempting to reconnect (attempt {}, delay: {:?})...",
368 reconnect_attempts, delay
369 );
370
371 match Self::establish_new_connection(
373 server_url,
374 connected,
375 connection_callbacks,
376 keep_alive_handle,
377 keep_alive_interval,
378 running,
379 )
380 .await
381 {
382 Ok((read, sender)) => {
383 if let Err(e) = Self::resubscribe_all(&sender, subscriptions).await {
385 error!("Failed to re-subscribe: {}", e);
386 }
387
388 return Ok((read, sender));
389 }
390 Err(e) => {
391 error!("Reconnection failed: {}", e);
392 }
394 }
395
396 if let Ok(Some(())) =
398 tokio::time::timeout(tokio::time::Duration::from_millis(10), cancel_rx.recv()).await
399 {
400 return Err(WebSocketError::Other(
401 "Cancelled during reconnection".to_string(),
402 ));
403 }
404 }
405 }
406
407 async fn notify_connection_status(
409 connected: bool,
410 callbacks: &Arc<RwLock<Vec<ConnectionCallback>>>,
411 ) {
412 let callbacks_guard = callbacks.read().await;
413 for callback in callbacks_guard.iter() {
414 callback(connected);
415 }
416 }
417
418 async fn handle_subscription_response(
420 response: SubscriptionResponse,
421 subscriptions: &Arc<RwLock<HashMap<String, SubscriptionHandler>>>,
422 ) {
423 info!("Received subscription response: {:?}", response);
424
425 let mut subs = subscriptions.write().await;
427
428 let Some(lookup_key) = &response.request_id else {
430 warn!("Received subscription response for unknown subscription");
431 return;
432 };
433
434 if let Some(handler) = subs.remove(lookup_key) {
435 if matches!(response.status, SubscriptionStatus::Subscribed) {
436 subs.insert(
438 response.subscription_id.clone(),
439 SubscriptionHandler {
440 topic: handler.topic,
441 filter: handler.filter,
442 pending: false,
443 },
444 );
445 }
446 } else {
447 warn!(
448 "Received subscription response for unknown subscription: {:?}",
449 response
450 );
451 }
452 }
453
454 async fn resubscribe_all(
456 sender: &mpsc::Sender<Message>,
457 subscriptions: &Arc<RwLock<HashMap<String, SubscriptionHandler>>>,
458 ) -> Result<(), WebSocketError> {
459 let resubscribe_list = {
461 let mut subs = subscriptions.write().await;
462
463 let to_resubscribe: Vec<_> = subs
465 .iter()
466 .map(|(id, handler)| (id.clone(), handler.topic.clone(), handler.filter.clone()))
467 .collect();
468
469 subs.clear();
471
472 to_resubscribe
473 };
474
475 for (_, topic, filter) in resubscribe_list {
477 let pending_id = format!("pending-{}-{}", topic, uuid::Uuid::new_v4());
479
480 {
482 let mut subs = subscriptions.write().await;
483 subs.insert(
484 pending_id.clone(),
485 SubscriptionHandler {
486 topic: topic.clone(),
487 filter: filter.clone(),
488 pending: true,
489 },
490 );
491 }
492
493 let request = WebSocketRequest::Subscribe(SubscriptionRequest {
495 topic,
496 filter,
497 request_id: Some(pending_id),
498 });
499
500 let message = serde_json::to_string(&request).map_err(|e| {
501 WebSocketError::Other(format!("Failed to serialize request: {}", e))
502 })?;
503
504 sender
505 .send(Message::Text(message.into()))
506 .await
507 .map_err(|e| WebSocketError::SendFailed(e.to_string()))?;
508 }
509
510 Ok(())
511 }
512
513 async fn process_message(
515 message: Message,
516 sender: &mpsc::Sender<Message>,
517 subscriptions: &Arc<RwLock<HashMap<String, SubscriptionHandler>>>,
518 event_callbacks: &Arc<RwLock<HashMap<EventTopic, Vec<EventCallback>>>>,
519 async_event_callbacks: &Arc<RwLock<HashMap<EventTopic, Vec<AsyncEventCallback>>>>,
520 ) -> bool {
521 match message {
523 Message::Text(text) => {
524 match serde_json::from_str::<WebSocketMessage>(&text) {
526 Ok(WebSocketMessage::Event(event)) => {
527 Self::handle_event(event, event_callbacks, async_event_callbacks).await;
528 }
529 Ok(WebSocketMessage::SubscriptionResponse(response)) => {
530 Self::handle_subscription_response(response, subscriptions).await;
531 }
532 Ok(WebSocketMessage::UnsubscribeResponse(response)) => {
533 Self::handle_unsubscribe_response(response, subscriptions).await;
534 }
535 Ok(WebSocketMessage::ErrorResponse(error)) => {
536 warn!("Subscription error: {}", error.error);
537 }
538 Err(e) => {
539 error!("Failed to parse WebSocket message: {}", e);
540 debug!("Message content: {}", text);
541 }
542 }
543 false
544 }
545 Message::Binary(_) => {
546 debug!("Received binary message");
547 false
548 }
549 Message::Ping(data) => {
550 if let Err(e) = sender.send(Message::Pong(data)).await {
552 warn!("Failed to send pong: {}", e);
553 }
554 false
555 }
556 Message::Pong(_) => false, Message::Frame(_) => false, Message::Close(_) => true, }
560 }
561
562 async fn handle_event(
564 event: Event,
565 event_callbacks: &Arc<RwLock<HashMap<EventTopic, Vec<EventCallback>>>>,
566 async_event_callbacks: &Arc<RwLock<HashMap<EventTopic, Vec<AsyncEventCallback>>>>,
567 ) {
568 let topic = event.topic();
569
570 {
572 let callbacks = event_callbacks.read().await;
573 if let Some(handlers) = callbacks.get(&topic) {
574 for handler in handlers {
575 match panic::catch_unwind(AssertUnwindSafe(|| {
577 handler(event.clone());
578 })) {
579 Ok(_) => {}
580 Err(e) => {
581 let panic_msg = if let Some(s) = e.downcast_ref::<&str>() {
583 s
584 } else if let Some(s) = e.downcast_ref::<String>() {
585 s.as_str()
586 } else {
587 "Unknown panic"
588 };
589 error!("Event handler panicked: {}", panic_msg);
590 }
591 }
592 }
593 }
594 }
595
596 {
598 let async_callbacks = async_event_callbacks.read().await;
599 if let Some(handlers) = async_callbacks.get(&topic) {
600 for handler in handlers {
601 let event_clone = event.clone();
604 let future = match panic::catch_unwind(AssertUnwindSafe(|| {
605 handler(event_clone.clone())
606 })) {
607 Ok(future) => future,
608 Err(e) => {
609 let panic_msg = if let Some(s) = e.downcast_ref::<&str>() {
611 s
612 } else if let Some(s) = e.downcast_ref::<String>() {
613 s.as_str()
614 } else {
615 "Unknown panic"
616 };
617 error!("Async event handler panicked during setup: {}", panic_msg);
618 continue;
619 }
620 };
621
622 tokio::spawn(async move {
624 match panic::catch_unwind(AssertUnwindSafe(|| async {
625 future.await;
626 })) {
627 Ok(f) => {
628 f.await;
629 }
630 Err(e) => {
631 let panic_msg = if let Some(s) = e.downcast_ref::<&str>() {
633 s
634 } else if let Some(s) = e.downcast_ref::<String>() {
635 s.as_str()
636 } else {
637 "Unknown panic"
638 };
639 error!(
640 "Async event handler panicked during execution: {}",
641 panic_msg
642 );
643 }
644 }
645 });
646 }
647 }
648 }
649 }
650
651 async fn handle_unsubscribe_response(
653 response: UnsubscribeResponse,
654 subscriptions: &Arc<RwLock<HashMap<String, SubscriptionHandler>>>,
655 ) {
656 if matches!(response.status, SubscriptionStatus::Unsubscribed) {
657 subscriptions
658 .write()
659 .await
660 .remove(&response.subscription_id);
661 }
662 }
663
664 async fn establish_connection(
666 server_url: &str,
667 ) -> Result<WebSocketStream<tokio_tungstenite::MaybeTlsStream<TcpStream>>, WebSocketError> {
668 match connect_async(server_url).await {
669 Ok((ws_stream, _)) => Ok(ws_stream),
670 Err(e) => Err(WebSocketError::ConnectionFailed(e.to_string())),
671 }
672 }
673
674 fn spawn_writer_task(
676 write: futures::stream::SplitSink<
677 WebSocketStream<tokio_tungstenite::MaybeTlsStream<TcpStream>>,
678 Message,
679 >,
680 ) -> mpsc::Sender<Message> {
681 let (sender, mut new_receiver) = mpsc::channel::<Message>(1000);
683
684 tokio::spawn(async move {
686 let mut writer = write;
687 while let Some(msg) = new_receiver.recv().await {
688 if let Err(e) = writer.send(msg).await {
689 error!("Failed to send message: {}", e);
690 break;
691 }
692 }
693 });
694
695 sender
696 }
697
698 #[allow(clippy::too_many_arguments)]
700 async fn message_processor(
701 initial_stream: futures::stream::SplitStream<
702 WebSocketStream<tokio_tungstenite::MaybeTlsStream<TcpStream>>,
703 >,
704 initial_sender: mpsc::Sender<Message>,
705 subscriptions: Arc<RwLock<HashMap<String, SubscriptionHandler>>>,
706 event_callbacks: Arc<RwLock<HashMap<EventTopic, Vec<EventCallback>>>>,
707 async_event_callbacks: Arc<RwLock<HashMap<EventTopic, Vec<AsyncEventCallback>>>>,
708 connection_callbacks: Arc<RwLock<Vec<ConnectionCallback>>>,
709 connected: Arc<Mutex<bool>>,
710 keep_alive_handle: Arc<Mutex<Option<tokio::task::JoinHandle<()>>>>,
711 keep_alive_interval: Arc<Mutex<Option<Duration>>>,
712 running: Arc<Mutex<bool>>,
713 server_url: Arc<String>,
714 auto_reconnect: Arc<Mutex<bool>>,
715 backoff_strategy: Arc<Mutex<BackoffStrategy>>,
716 max_reconnect_attempts: Arc<Mutex<usize>>,
717 mut cancel_rx: mpsc::Receiver<()>,
718 ) {
719 let mut reconnect_attempts: usize = 0;
720 let mut read = initial_stream;
721 let mut sender = initial_sender;
722
723 while *running.lock().await {
725 let disconnected = Self::process_messages(
726 &mut read,
727 &sender,
728 &subscriptions,
729 &event_callbacks,
730 &async_event_callbacks,
731 &connection_callbacks,
732 &connected,
733 &mut cancel_rx,
734 )
735 .await;
736
737 if disconnected {
738 if !*running.lock().await || !*auto_reconnect.lock().await {
740 return; }
742
743 match Self::attempt_reconnection(
745 &server_url,
746 &connected,
747 &connection_callbacks,
748 &subscriptions,
749 &keep_alive_handle,
750 &keep_alive_interval,
751 &running,
752 &backoff_strategy,
753 &max_reconnect_attempts,
754 &mut reconnect_attempts,
755 &mut cancel_rx,
756 )
757 .await
758 {
759 Ok((new_read, new_sender)) => {
760 read = new_read;
762 sender = new_sender;
763 reconnect_attempts = 0; }
765 Err(_) => return, }
767 }
768 }
769 }
770
771 #[allow(clippy::too_many_arguments)]
773 async fn process_messages(
774 read: &mut futures::stream::SplitStream<
775 WebSocketStream<tokio_tungstenite::MaybeTlsStream<TcpStream>>,
776 >,
777 sender: &mpsc::Sender<Message>,
778 subscriptions: &Arc<RwLock<HashMap<String, SubscriptionHandler>>>,
779 event_callbacks: &Arc<RwLock<HashMap<EventTopic, Vec<EventCallback>>>>,
780 async_event_callbacks: &Arc<RwLock<HashMap<EventTopic, Vec<AsyncEventCallback>>>>,
781 connection_callbacks: &Arc<RwLock<Vec<ConnectionCallback>>>,
782 connected: &Arc<Mutex<bool>>,
783 cancel_rx: &mut mpsc::Receiver<()>,
784 ) -> bool {
785 loop {
786 tokio::select! {
787 _ = cancel_rx.recv() => {
789 return false; }
791
792 message = read.next() => {
794 match message {
795 Some(Ok(msg)) => {
796 if Self::process_message(msg, sender, subscriptions, event_callbacks, async_event_callbacks).await {
797 Self::update_connection_status(connected, false, connection_callbacks).await;
799 return true; }
801 }
802 Some(Err(e)) => {
803 error!("WebSocket read error: {}", e);
805 Self::update_connection_status(connected, false, connection_callbacks).await;
806 return true; }
808 None => {
809 debug!("WebSocket stream ended");
811 Self::update_connection_status(connected, false, connection_callbacks).await;
812 return true; }
814 }
815 }
816 }
817 }
818 }
819
820 pub async fn subscribe(
822 &self,
823 topic: EventTopic,
824 filter: EventFilter,
825 ) -> Result<(), WebSocketError> {
826 let pending_id = format!("pending-{}-{}", topic, uuid::Uuid::new_v4());
828
829 let subs = self.subscriptions.read().await;
831 for (_, handler) in subs.iter() {
832 if !handler.pending && handler.topic == topic && handler.filter == filter {
834 return Ok(());
835 }
836 }
837 drop(subs);
838
839 let mut subs = self.subscriptions.write().await;
841 subs.insert(
842 pending_id.clone(), SubscriptionHandler {
844 topic: topic.clone(),
845 filter: filter.clone(),
846 pending: true,
847 },
848 );
849 drop(subs);
850
851 let request = WebSocketRequest::Subscribe(SubscriptionRequest {
853 topic: topic.clone(),
854 filter: filter.clone(),
855 request_id: Some(pending_id.clone()),
856 });
857
858 let message = serde_json::to_string(&request)
859 .map_err(|e| WebSocketError::Other(format!("Failed to serialize request: {}", e)))?;
860
861 self.sender
862 .send(Message::Text(message.into()))
863 .await
864 .map_err(|e| WebSocketError::SendFailed(e.to_string()))?;
865
866 Ok(())
867 }
868
869 pub async fn on_event<F>(
871 &self,
872 topic: EventTopic,
873 filter: Option<EventFilter>,
874 callback: F,
875 ) -> Result<(), WebSocketError>
876 where
877 F: Fn(Event) + Send + Sync + 'static,
878 {
879 let subs = self.subscriptions.read().await;
881 let has_topic_subscription = subs.values().any(|s| s.topic == topic);
882 drop(subs);
883
884 if !has_topic_subscription {
885 self.subscribe(topic.clone(), filter.unwrap_or_default())
887 .await?;
888 }
889
890 let mut callbacks = self.event_callbacks.write().await;
892 callbacks
893 .entry(topic)
894 .or_insert_with(Vec::new)
895 .push(Box::new(callback));
896
897 Ok(())
898 }
899
900 pub async fn on_event_async<F, Fut>(
902 &self,
903 topic: EventTopic,
904 filter: Option<EventFilter>,
905 callback: F,
906 ) -> Result<(), WebSocketError>
907 where
908 F: Fn(Event) -> Fut + Send + Sync + 'static,
909 Fut: Future<Output = ()> + Send + 'static,
910 {
911 let subs = self.subscriptions.read().await;
913 let has_topic_subscription = subs.values().any(|s| s.topic == topic);
914 drop(subs);
915
916 if !has_topic_subscription {
917 self.subscribe(topic.clone(), filter.unwrap_or_default())
919 .await?;
920 }
921
922 let boxed_callback =
924 move |event: Event| -> BoxFuture<'static, ()> { callback(event).boxed() };
925
926 let mut callbacks = self.async_event_callbacks.write().await;
928 callbacks
929 .entry(topic)
930 .or_insert_with(Vec::new)
931 .push(Box::new(boxed_callback));
932
933 Ok(())
934 }
935
936 pub async fn on_connection_change<F>(&self, callback: F)
938 where
939 F: Fn(bool) + Send + Sync + 'static,
940 {
941 self.connection_callbacks
942 .write()
943 .await
944 .push(Box::new(callback));
945 }
946
947 pub async fn unsubscribe(&self, subscription_id: &str) -> Result<(), WebSocketError> {
949 let topic = {
950 let subs = self.subscriptions.read().await;
952 match subs.get(subscription_id) {
953 Some(sub) => sub.topic.clone(),
954 None => {
955 return Err(WebSocketError::UnsubscriptionFailed(
956 "Subscription not found".to_string(),
957 ))
958 }
959 }
960 };
961
962 let request = WebSocketRequest::Unsubscribe(UnsubscribeRequest {
963 topic,
964 subscription_id: subscription_id.to_string(),
965 });
966
967 let request_json =
968 serde_json::to_string(&request).map_err(|e| WebSocketError::Other(e.to_string()))?;
969
970 self.sender
972 .send(Message::Text(request_json.into()))
973 .await
974 .map_err(|e| WebSocketError::SendFailed(e.to_string()))?;
975
976 self.subscriptions.write().await.remove(subscription_id);
978
979 Ok(())
980 }
981
982 pub async fn unsubscribe_topic(&self, topic: &EventTopic) -> Result<(), WebSocketError> {
984 let subscription_ids: Vec<String> = {
985 let subs = self.subscriptions.read().await;
986 subs.iter()
987 .filter(|(_, handler)| handler.topic == *topic)
988 .map(|(id, _)| id.clone())
989 .collect()
990 };
991
992 let mut result = Ok(());
993 for id in subscription_ids {
994 if let Err(e) = self.unsubscribe(&id).await {
995 result = Err(e);
996 }
997 }
998
999 result
1000 }
1001
1002 pub async fn remove_event_listeners(&self, topic: &EventTopic) {
1004 let mut callbacks = self.event_callbacks.write().await;
1006 callbacks.remove(topic);
1007
1008 let mut async_callbacks = self.async_event_callbacks.write().await;
1010 async_callbacks.remove(topic);
1011 }
1012
1013 pub async fn set_auto_reconnect(
1015 &self,
1016 enabled: bool,
1017 interval: std::time::Duration,
1018 max_attempts: usize,
1019 ) {
1020 *self.auto_reconnect.lock().await = enabled;
1021 *self.backoff_strategy.lock().await = BackoffStrategy::Constant(interval);
1022 *self.max_reconnect_attempts.lock().await = max_attempts;
1023 }
1024
1025 pub async fn is_connected(&self) -> bool {
1027 *self.connected.lock().await
1028 }
1029
1030 pub async fn close(&self) -> Result<(), WebSocketError> {
1035 *self.auto_reconnect.lock().await = false;
1037
1038 *self.running.lock().await = false;
1040
1041 if let Some(cancel_tx) = &self.cancel_tx {
1043 let _ = cancel_tx.send(()).await;
1044 }
1045
1046 if let Some(handle) = self.keep_alive_handle.lock().await.take() {
1048 handle.abort();
1049 }
1050
1051 let _ = self.sender.send(Message::Close(None)).await;
1053
1054 Self::update_connection_status(&self.connected, false, &self.connection_callbacks).await;
1056
1057 Ok(())
1058 }
1059
1060 pub async fn set_reconnect_options(
1062 &self,
1063 enabled: bool,
1064 strategy: BackoffStrategy,
1065 max_attempts: usize,
1066 ) {
1067 *self.auto_reconnect.lock().await = enabled;
1068 *self.backoff_strategy.lock().await = strategy;
1069 *self.max_reconnect_attempts.lock().await = max_attempts;
1070 }
1071
1072 async fn update_connection_status(
1074 connected: &Arc<Mutex<bool>>,
1075 new_state: bool,
1076 connection_callbacks: &Arc<RwLock<Vec<ConnectionCallback>>>,
1077 ) -> bool {
1078 let mut connected_guard = connected.lock().await;
1080 let changed = *connected_guard != new_state;
1081 *connected_guard = new_state;
1082 drop(connected_guard);
1083
1084 if changed {
1086 Self::notify_connection_status(new_state, connection_callbacks).await;
1087 }
1088
1089 changed }
1091
1092 pub async fn enable_keep_alive(&self, interval: Duration) -> Result<(), WebSocketError> {
1094 *self.keep_alive_interval.lock().await = Some(interval);
1096
1097 Self::restart_keep_alive(
1099 &self.keep_alive_handle,
1100 &self.sender,
1101 &self.running,
1102 &self.connected,
1103 interval,
1104 )
1105 .await
1106 }
1107
1108 async fn restart_keep_alive(
1110 keep_alive_handle: &Arc<Mutex<Option<tokio::task::JoinHandle<()>>>>,
1111 sender: &mpsc::Sender<Message>,
1112 running: &Arc<Mutex<bool>>,
1113 connected: &Arc<Mutex<bool>>,
1114 interval: Duration,
1115 ) -> Result<(), WebSocketError> {
1116 if let Some(handle) = keep_alive_handle.lock().await.take() {
1118 handle.abort();
1119 }
1120
1121 let sender = sender.clone();
1122 let running = running.clone();
1123 let connected = connected.clone();
1124
1125 let handle = tokio::spawn(async move {
1127 let mut interval_timer = tokio::time::interval(interval);
1128
1129 while *running.lock().await {
1130 interval_timer.tick().await;
1131
1132 if *connected.lock().await {
1134 if let Err(e) = sender.send(Message::Ping(Bytes::from_static(&[]))).await {
1135 error!("Failed to send ping: {}", e);
1136 }
1138 }
1139 }
1140 });
1141
1142 *keep_alive_handle.lock().await = Some(handle);
1144
1145 Ok(())
1146 }
1147}
1148
1149impl Drop for WebSocketClient {
1150 fn drop(&mut self) {
1151 if let Some(running) = Arc::get_mut(&mut self.running) {
1153 if let Ok(mut guard) = running.try_lock() {
1154 *guard = false;
1155 }
1156 }
1157
1158 if let Some(cancel_tx) = self.cancel_tx.take() {
1160 let _ = cancel_tx.try_send(());
1162 }
1163
1164 if let Some(keep_alive_handle) = Arc::get_mut(&mut self.keep_alive_handle) {
1166 if let Ok(mut guard) = keep_alive_handle.try_lock() {
1167 if let Some(handle) = guard.take() {
1168 handle.abort();
1169 }
1170 }
1171 }
1172 }
1173}