binance_rs_plus/
async_websocket_client.rs1use crate::errors::{Error, Result};
2use futures_util::{StreamExt, SinkExt};
3use tokio_tungstenite::{
4 connect_async, tungstenite::protocol::Message, MaybeTlsStream, WebSocketStream,
5};
6use tokio::net::TcpStream;
7use url::Url;
8use std::sync::Arc;
9use tokio::sync::Mutex;
10use serde::de::DeserializeOwned;
11use std::future::Future;
12use std::pin::Pin;
13
14pub struct AsyncWebsocketClient<'a, E, H>
19where
20 E: DeserializeOwned + Send + std::fmt::Debug + 'a, H: FnMut(E) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>> + Send + Sync + 'a,
22{
23 socket: Arc<Mutex<Option<WebSocketStream<MaybeTlsStream<TcpStream>>>>>,
24 handler: Arc<Mutex<H>>,
25 phantom: std::marker::PhantomData<&'a E>,
26}
27
28impl<'a, E, H> AsyncWebsocketClient<'a, E, H>
29where
30 E: DeserializeOwned + Send + std::fmt::Debug + 'a, H: FnMut(E) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>> + Send + Sync + 'a,
32{
33 pub fn new(handler: H) -> Self {
34 AsyncWebsocketClient {
35 socket: Arc::new(Mutex::new(None)),
36 handler: Arc::new(Mutex::new(handler)),
37 phantom: std::marker::PhantomData,
38 }
39 }
40
41 pub async fn connect(&self, wss_url: &str) -> Result<()> {
42 let url_obj = Url::parse(wss_url).map_err(Error::UrlParser)?;
43 let (ws_stream, _response) = connect_async(url_obj.as_str()) .await
45 .map_err(Error::WebSocket)?;
46
47 let mut socket_guard = self.socket.lock().await;
48 *socket_guard = Some(ws_stream);
49 Ok(())
50 }
51
52 pub async fn disconnect(&self) -> Result<()> {
53 let mut socket_guard = self.socket.lock().await;
54 if let Some(stream) = socket_guard.as_mut() {
55 stream.close(None).await.map_err(Error::WebSocket)?;
56 *socket_guard = None;
57 Ok(())
58 } else {
59 Err(Error::Custom("Not connected".to_string()))
60 }
61 }
62
63 async fn handle_message_text(&self, msg_text: String) -> Result<()> {
64 if let Ok(event) = serde_json::from_str::<E>(&msg_text) {
70 let mut handler_guard = self.handler.lock().await;
71 (handler_guard)(event).await?;
72 Ok(())
73 } else {
74 if let Ok(value) = serde_json::from_str::<serde_json::Value>(&msg_text) {
77 if let Some(data_val) = value.get("data") {
78 match serde_json::from_value::<E>(data_val.clone()) {
79 Ok(event) => {
80 let mut handler_guard = self.handler.lock().await;
81 (handler_guard)(event).await?;
82 return Ok(());
83 }
84 Err(e_inner) => {
85 return Err(Error::Json(e_inner));
86 }
87 }
88 }
89 if let Ok(stream_value) = serde_json::from_str::<serde_json::Value>(&msg_text) {
91 if let Some(_stream_name) = stream_value.get("stream") {
92 if let Some(data_val_stream) = stream_value.get("data") {
93 match serde_json::from_value::<E>(data_val_stream.clone()) {
94 Ok(event) => {
95 let mut handler_guard = self.handler.lock().await;
96 (handler_guard)(event).await?;
97 return Ok(());
98 }
99 Err(e_inner_stream) => {
100 return Err(Error::Json(e_inner_stream));
101 }
102 }
103 }
104 }
105 }
106 }
107 Err(Error::Json(
109 serde_json::from_str::<E>(&msg_text).unwrap_err(),
110 ))
111 }
112 }
113
114 pub async fn event_loop(&self, running: Arc<std::sync::atomic::AtomicBool>) -> Result<()> {
115 while running.load(std::sync::atomic::Ordering::Relaxed) {
116 let mut socket_guard = self.socket.lock().await;
117 if let Some(stream) = socket_guard.as_mut() {
118 match stream.next().await {
119 Some(Ok(message)) => {
120 drop(socket_guard); match message {
122 Message::Text(text) => {
123 if let Err(e) = self.handle_message_text(text).await {
124 eprintln!("Error handling message: {:?}", e); return Err(e);
130 }
131 }
132 Message::Binary(_) => { }
133 Message::Ping(payload) => {
134 let mut new_socket_guard = self.socket.lock().await;
136 if let Some(s) = new_socket_guard.as_mut() {
137 s.send(Message::Pong(payload))
138 .await
139 .map_err(Error::WebSocket)?;
140 }
141 drop(new_socket_guard);
142 }
143 Message::Pong(_) => { }
144 Message::Close(close_frame) => {
145 eprintln!("WebSocket closed by server: {:?}", close_frame);
146 return Err(Error::Custom(format!(
147 "WebSocket closed by server: {:?}",
148 close_frame
149 )));
150 }
151 Message::Frame(_) => { }
153 }
154 }
155 Some(Err(e)) => {
156 return Err(Error::WebSocket(e));
158 }
159 None => {
160 return Err(Error::Custom("WebSocket stream ended".to_string()));
162 }
163 }
164 } else {
165 drop(socket_guard);
167 tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
168 }
169 }
170 Ok(())
171 }
172}