openlimits_coinbase/client/
stream.rs1use 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
34pub 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}