Skip to main content

openlimits_coinbase/client/
stream.rs

1use async_trait::async_trait;
2use std::{collections::HashMap, pin::Pin, task::Poll};
3use futures::{
4    stream::{SplitStream, Stream},
5    SinkExt, StreamExt,
6};
7use serde::{Deserialize, Serialize};
8use tokio::net::TcpStream;
9use tokio_tungstenite::tungstenite::Message;
10use tokio_tungstenite::{connect_async, MaybeTlsStream, WebSocketStream};
11use crate::model::websocket::{Channel, CoinbaseSubscription, CoinbaseWebsocketMessage, Subscribe, SubscribeCmd};
12use openlimits_exchange::errors::OpenLimitsError;
13use crate::model::websocket::ChannelType;
14use crate::CoinbaseParameters;
15use openlimits_exchange::traits::stream::{ExchangeStream, Subscriptions};
16use futures::stream::BoxStream;
17use std::sync::Mutex;
18use tokio::sync::mpsc::{unbounded_channel, UnboundedSender};
19use super::shared::Result;
20use openlimits_exchange::exchange::Environment;
21
22const WS_URL_PROD: &str = "wss://ws-feed.exchange.coinbase.com";
23const WS_URL_SANDBOX: &str = "wss://ws-feed-public.sandbox.exchange.coinbase.com";
24
25#[derive(Debug, Clone, Deserialize, Serialize)]
26#[serde(untagged)]
27enum Either<L, R> {
28    Left(L),
29    Right(R),
30}
31
32type WSStream = WebSocketStream<MaybeTlsStream<TcpStream>>;
33
34/// A websocket connection to Coinbase
35pub struct CoinbaseWebsocket {
36    pub subscriptions: HashMap<CoinbaseSubscription, SplitStream<WSStream>>,
37    pub parameters: CoinbaseParameters,
38    disconnection_senders: Mutex<Vec<UnboundedSender<()>>>,
39}
40
41impl CoinbaseWebsocket {
42    pub async fn subscribe_(&mut self, subscription: CoinbaseSubscription) -> Result<()> {
43        let (channels, product_ids) = match &subscription {
44            CoinbaseSubscription::Level2(product_id) => (
45                vec![Channel::Name(ChannelType::Level2)],
46                vec![product_id.clone()],
47            ),
48            CoinbaseSubscription::Heartbeat(product_id) => (
49                vec![Channel::Name(ChannelType::Heartbeat)],
50                vec![product_id.clone()],
51            ),
52            CoinbaseSubscription::Matches(product_id) => (
53                vec![Channel::Name(ChannelType::Matches)],
54                vec![product_id.clone()]
55            )
56        };
57        let subscribe = Subscribe {
58            _type: SubscribeCmd::Subscribe,
59            auth: None,
60            channels,
61            product_ids,
62        };
63
64        let stream = self.connect(subscribe).await?;
65        self.subscriptions.insert(subscription, stream);
66        Ok(())
67    }
68
69    pub async fn connect(&self, subscribe: Subscribe) -> Result<SplitStream<WSStream>> {
70        let ws_url = if self.parameters.environment == Environment::Sandbox {
71            WS_URL_SANDBOX
72        } else {
73            WS_URL_PROD
74        };
75        let url = url::Url::parse(ws_url).expect("Couldn't parse url.");
76        let (ws_stream, _) = connect_async(&url).await?;
77        let (mut sink, stream) = ws_stream.split();
78        let subscribe = serde_json::to_string(&subscribe)?;
79
80        sink.send(Message::Text(subscribe)).await?;
81        Ok(stream)
82    }
83}
84
85impl Stream for CoinbaseWebsocket {
86    type Item = Result<CoinbaseWebsocketMessage>;
87
88    fn poll_next(
89        mut self: std::pin::Pin<&mut Self>,
90        cx: &mut std::task::Context<'_>,
91    ) -> Poll<Option<Self::Item>> {
92        for (_sub, stream) in &mut self.subscriptions.iter_mut() {
93            if let Poll::Ready(Some(message)) = Pin::new(stream).poll_next(cx) {
94                let message = parse_message(message?);
95                return Poll::Ready(Some(message));
96            }
97        }
98
99        std::task::Poll::Pending
100    }
101}
102
103fn parse_message(ws_message: Message) -> Result<CoinbaseWebsocketMessage> {
104    let msg = match ws_message {
105        Message::Text(m) => m,
106        _ => return Err(OpenLimitsError::SocketError()),
107    };
108    Ok(serde_json::from_str(&msg)?)
109}
110
111#[async_trait]
112impl ExchangeStream for CoinbaseWebsocket {
113    type InitParams = CoinbaseParameters;
114    type Subscription = CoinbaseSubscription;
115    type Response = CoinbaseWebsocketMessage;
116
117    async fn new(parameters: Self::InitParams) -> Result<Self> {
118        Ok(Self {
119            subscriptions: Default::default(),
120            parameters,
121            disconnection_senders: Default::default(),
122        })
123    }
124
125    async fn disconnect(&self) {
126        if let Ok(mut senders) = self.disconnection_senders.lock() {
127            for sender in senders.iter() {
128                sender.send(()).ok();
129            }
130            senders.clear();
131        }
132    }
133
134    async fn create_stream_specific(
135        &self,
136        subscription: Subscriptions<Self::Subscription>,
137    ) -> Result<BoxStream<'static, Result<Self::Response>>> {
138        let ws_url = if self.parameters.environment == Environment::Sandbox {
139            WS_URL_SANDBOX
140        } else {
141            WS_URL_PROD
142        };
143        let endpoint = url::Url::parse(ws_url).expect("Couldn't parse url.");
144        let (ws_stream, _) = connect_async(endpoint).await?;
145
146        let (channel_name, product_ids) = match &subscription.as_slice()[0] {
147            CoinbaseSubscription::Level2(product_id) => (
148                ChannelType::Level2,
149                vec![product_id.clone()],
150            ),
151            CoinbaseSubscription::Heartbeat(product_id) => (
152                ChannelType::Heartbeat,
153                vec![product_id.clone()],
154            ),
155            CoinbaseSubscription::Matches(product_id) => (
156                ChannelType::Matches,
157                vec![product_id.clone()]
158            )
159        };
160        let channels = vec![Channel::Name(channel_name.clone())];
161        let subscribe = Subscribe {
162            _type: SubscribeCmd::Subscribe,
163            auth: None,
164            channels,
165            product_ids: product_ids.clone(),
166        };
167        let subscribe = serde_json::to_string(&subscribe)?;
168        let (mut sink, stream) = ws_stream.split();
169        let (disconnection_sender, mut disconnection_receiver) = unbounded_channel();
170        sink.send(Message::Text(subscribe)).await?;
171        tokio::spawn(async move {
172            if disconnection_receiver.recv().await.is_some() {
173                sink.close().await.ok();
174            }
175        });
176
177        if let Ok(mut senders) = self.disconnection_senders.lock() {
178            senders.push(disconnection_sender);
179        }
180        let mut s = stream.map(|message| parse_message(message?));
181
182        let name = channel_name;
183        let product = Channel::WithProduct { name, product_ids };
184        let channels = vec![product];
185        let expected_response = CoinbaseWebsocketMessage::Subscriptions { channels };
186
187        let response = s.next().await;
188        if let Some(Ok(response)) = response {
189            if response == expected_response {
190                Ok(s.boxed())
191            } else {
192                Err(OpenLimitsError::UnkownResponse(format!("Response: {:#?}, expected response: {:#?}", response, expected_response)))
193            }
194        } else {
195            Err(OpenLimitsError::UnkownResponse(format!("No response")))
196        }
197    }
198}