Skip to main content

binance_rs_plus/
async_websocket_client.rs

1use 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
14/// A generic asynchronous WebSocket client.
15///
16/// E: The type of event deserialized from messages.
17/// H: The type of the handler function.
18pub struct AsyncWebsocketClient<'a, E, H>
19where
20    E: DeserializeOwned + Send + std::fmt::Debug + 'a, // Added Debug
21    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, // Added Debug
31    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()) // Convert Url to &str
44            .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        // This parsing logic might need to be customized based on how
65        // Binance wraps multi-stream data or other specific message formats.
66        // For now, assuming direct deserialization or a simple 'data' field check.
67
68        // Attempt direct deserialization
69        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 direct deserialization fails, check for a common "data" wrapper
75            // This is a simplified example; real-world scenarios might be more complex.
76            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 not a "data" wrapper, or if that also fails to parse as E
90                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            // If all attempts fail, return original direct deserialization error
108            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); // Release lock before handling message
121                        match message {
122                            Message::Text(text) => {
123                                if let Err(e) = self.handle_message_text(text).await {
124                                    // Log error or propagate? For now, let's propagate critical parsing/handling errors.
125                                    // Specific errors like pings being unhandled by user might be logged and continued.
126                                    eprintln!("Error handling message: {:?}", e); // Temporary logging
127                                    // Depending on severity, may want to break or continue.
128                                    // For now, if handle_message_text returns an error, we propagate it.
129                                    return Err(e);
130                                }
131                            }
132                            Message::Binary(_) => { /* Handle binary data if necessary */ }
133                            Message::Ping(payload) => {
134                                // Re-acquire lock to send Pong
135                                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(_) => { /* Pong received */ }
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(_) => { /* Low-level frame, usually not handled directly */
152                            }
153                        }
154                    }
155                    Some(Err(e)) => {
156                        // WebSocket stream error
157                        return Err(Error::WebSocket(e));
158                    }
159                    None => {
160                        // Stream ended (disconnected)
161                        return Err(Error::Custom("WebSocket stream ended".to_string()));
162                    }
163                }
164            } else {
165                // Socket not connected, maybe wait and retry or break
166                drop(socket_guard);
167                tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
168            }
169        }
170        Ok(())
171    }
172}