Skip to main content

binance_rs_plus/
websockets.rs

1use crate::async_websocket_client::AsyncWebsocketClient;
2use crate::config::Config;
3use crate::errors::Result;
4use crate::model::{
5    AccountUpdateEvent, AggrTradesEvent, BalanceUpdateEvent, BookTickerEvent, DayTickerEvent,
6    DepthOrderBookEvent, KlineEvent, OrderBook, OrderTradeEvent, TradeEvent, WindowTickerEvent,
7};
8// New
9
10use serde::{Deserialize, Serialize};
11
12use std::future::Future;
13// New
14use std::pin::Pin;
15// New
16use std::sync::atomic::AtomicBool;
17// Ordering might be needed by user
18use std::sync::Arc;
19// New
20use tokio::sync::Mutex as TokioMutex;
21// New for handler
22
23// WebsocketAPI enum remains the same
24#[allow(clippy::all)]
25enum WebsocketAPI {
26    Default,
27    MultiStream,
28    Custom(String),
29}
30
31impl WebsocketAPI {
32    fn params(self, subscription: &str) -> String {
33        match self {
34            WebsocketAPI::Default => format!("wss://stream.binance.com/ws/{}", subscription),
35            WebsocketAPI::MultiStream => {
36                format!("wss://stream.binance.com/stream?streams={}", subscription)
37            }
38            WebsocketAPI::Custom(url) => format!("{}/{}", url, subscription),
39        }
40    }
41}
42
43// WebsocketEvent enum remains the same - this is what the user's handler will receive
44#[allow(clippy::large_enum_variant)]
45#[derive(Debug, Serialize, Deserialize, Clone)]
46pub enum WebsocketEvent {
47    AccountUpdate(AccountUpdateEvent),
48    BalanceUpdate(BalanceUpdateEvent),
49    OrderTrade(OrderTradeEvent),
50    AggrTrades(AggrTradesEvent),
51    Trade(TradeEvent),
52    OrderBook(OrderBook),
53    DayTicker(DayTickerEvent),
54    DayTickerAll(Vec<DayTickerEvent>),
55    WindowTicker(WindowTickerEvent),
56    WindowTickerAll(Vec<WindowTickerEvent>),
57    Kline(KlineEvent),
58    DepthOrderBook(DepthOrderBookEvent),
59    BookTicker(BookTickerEvent),
60}
61
62// Events enum is what AsyncWebsocketClient will deserialize into (as E)
63// This needs to be DeserializeOwned + Send + 'a
64#[derive(Serialize, Deserialize, Debug, Clone)] // Added Clone for potential use, ensure it's Send + 'a compatible
65#[serde(untagged)]
66enum Events {
67    DayTickerEventAll(Vec<DayTickerEvent>),
68    WindowTickerEventAll(Vec<WindowTickerEvent>),
69    BalanceUpdateEvent(BalanceUpdateEvent),
70    DayTickerEvent(DayTickerEvent),
71    WindowTickerEvent(WindowTickerEvent),
72    BookTickerEvent(BookTickerEvent),
73    AccountUpdateEvent(AccountUpdateEvent),
74    OrderTradeEvent(OrderTradeEvent),
75    AggrTradesEvent(AggrTradesEvent),
76    TradeEvent(TradeEvent),
77    KlineEvent(KlineEvent),
78    OrderBook(OrderBook),
79    DepthOrderBookEvent(DepthOrderBookEvent),
80}
81
82// Define the type for the adapter handler passed to AsyncWebsocketClient
83type AdapterHandler<'a> = Box<
84    dyn FnMut(Events) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>> + Send + Sync + 'a,
85>;
86
87pub struct WebSockets<'a> {
88    client: AsyncWebsocketClient<'a, Events, AdapterHandler<'a>>,
89}
90
91impl<'a> WebSockets<'a> {
92    pub fn new<Callback>(user_handler: Callback) -> WebSockets<'a>
93    where
94        Callback: FnMut(WebsocketEvent) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>>
95            + Send
96            + Sync
97            + 'a,
98    {
99        let shared_user_handler = Arc::new(TokioMutex::new(user_handler));
100
101        let adapter_handler: AdapterHandler<'a> = Box::new(move |events_obj: Events| {
102            let user_handler_clone = Arc::clone(&shared_user_handler);
103            Box::pin(async move {
104                let action = match events_obj {
105                    Events::DayTickerEventAll(v) => WebsocketEvent::DayTickerAll(v),
106                    Events::WindowTickerEventAll(v) => WebsocketEvent::WindowTickerAll(v),
107                    Events::BalanceUpdateEvent(v) => WebsocketEvent::BalanceUpdate(v),
108                    Events::DayTickerEvent(v) => WebsocketEvent::DayTicker(v),
109                    Events::WindowTickerEvent(v) => WebsocketEvent::WindowTicker(v),
110                    Events::BookTickerEvent(v) => WebsocketEvent::BookTicker(v),
111                    Events::AccountUpdateEvent(v) => WebsocketEvent::AccountUpdate(v),
112                    Events::OrderTradeEvent(v) => WebsocketEvent::OrderTrade(v),
113                    Events::AggrTradesEvent(v) => WebsocketEvent::AggrTrades(v),
114                    Events::TradeEvent(v) => WebsocketEvent::Trade(v),
115                    Events::KlineEvent(v) => WebsocketEvent::Kline(v),
116                    Events::OrderBook(v) => WebsocketEvent::OrderBook(v),
117                    Events::DepthOrderBookEvent(v) => WebsocketEvent::DepthOrderBook(v),
118                };
119                let mut handler_guard = user_handler_clone.lock().await;
120                (handler_guard)(action).await
121            })
122        });
123
124        WebSockets {
125            client: AsyncWebsocketClient::new(adapter_handler),
126        }
127    }
128
129    pub async fn connect(&mut self, subscription: &str) -> Result<()> {
130        self.client
131            .connect(&WebsocketAPI::Default.params(subscription))
132            .await
133    }
134
135    pub async fn connect_with_config(&mut self, subscription: &str, config: &Config) -> Result<()> {
136        self.client
137            .connect(&WebsocketAPI::Custom(config.ws_endpoint.clone()).params(subscription))
138            .await
139    }
140
141    pub async fn connect_multiple_streams(&mut self, endpoints: &[String]) -> Result<()> {
142        self.client
143            .connect(&WebsocketAPI::MultiStream.params(&endpoints.join("/")))
144            .await
145    }
146
147    pub async fn disconnect(&mut self) -> Result<()> {
148        self.client.disconnect().await
149    }
150
151    // event_loop now takes Arc<AtomicBool>
152    pub async fn event_loop(&mut self, running: Arc<AtomicBool>) -> Result<()> {
153        self.client.event_loop(running).await
154    }
155}