Skip to main content

stoat/
websocket.rs

1use futures::{FutureExt, SinkExt, StreamExt, future::select};
2use std::{sync::Arc, time::Duration};
3use stoat_database::events::{
4    client::{EventV1, Ping},
5    server::ClientMessage,
6};
7use tokio::{
8    sync::{
9        Mutex,
10        mpsc::{UnboundedReceiver, UnboundedSender},
11    },
12    task::AbortHandle,
13    time::sleep,
14};
15use tokio_tungstenite::connect_async_with_config;
16use tungstenite::{Message, protocol::WebSocketConfig};
17
18use crate::{Error, cache::GlobalCache};
19
20#[derive(Debug, Clone, PartialEq)]
21pub(crate) enum ProgramMessage {
22    Close,
23}
24
25#[derive(Debug)]
26pub(crate) enum EventMessage {
27    Client(ClientMessage),
28    Program(ProgramMessage),
29}
30
31async fn send(
32    ws: &Arc<Mutex<impl SinkExt<Message, Error = tungstenite::Error> + Unpin>>,
33    event: &ClientMessage,
34) -> Result<(), tungstenite::Error> {
35    let mut lock = ws.lock().await;
36
37    #[cfg(not(feature = "msgpack"))]
38    let message = Message::text(serde_json::to_string(event).unwrap());
39
40    #[cfg(feature = "msgpack")]
41    let message = Message::binary(rmp_serde::to_vec_named(event).unwrap());
42
43    lock.send(message).await
44}
45
46pub(crate) async fn run(
47    events: UnboundedSender<EventV1>,
48    client_events: Arc<Mutex<UnboundedReceiver<EventMessage>>>,
49    global_state: GlobalCache,
50    token: String,
51) -> Result<(), Error> {
52    let message_format = if cfg!(feature = "msgpack") {
53        "msgpack"
54    } else {
55        "json"
56    };
57
58    let uri = format!(
59        "{}/?token={token}&format={message_format}",
60        &global_state.api_config.ws
61    );
62
63    log::debug!("Connecting to websocket with {uri}");
64
65    let mut ws_config = WebSocketConfig::default();
66    ws_config.max_frame_size = Some(usize::MAX);
67    ws_config.max_message_size = Some(usize::MAX);
68
69    let (ws, _) = connect_async_with_config(uri, Some(ws_config), false)
70        .await
71        .inspect_err(|e| {
72            if let tungstenite::Error::Http(resp) = e
73                && let Some(body) = resp.body()
74                && let Ok(body) = std::str::from_utf8(body)
75            {
76                log::error!("Error when attempting to establish websocket connection:\n{body}");
77            };
78        })?;
79
80    let (ws_send, mut ws_receive) = ws.split();
81
82    let ws_send = Arc::new(Mutex::new(ws_send));
83
84    let server_client = {
85        let ws_send = ws_send.clone();
86
87        async move {
88            let mut heartbeat_handle: Option<AbortHandle> = None;
89
90            while let Some(msg) = ws_receive.next().await {
91                let msg = msg?;
92
93                let event = match msg {
94                    Message::Text(data) => {
95                        serde_json::from_str(data.as_str()).map_err(|e| e.to_string())
96                    }
97                    #[cfg(feature = "msgpack")]
98                    Message::Binary(data) => {
99                        rmp_serde::from_slice(&data).map_err(|e| e.to_string())
100                    }
101                    msg => {
102                        if let Ok(text) = msg.to_text() {
103                            log::error!("Unexpected WS message: {text:?}");
104                        } else {
105                            log::error!("Unexpected WS message: {:?}", msg.into_data());
106                        }
107                        continue;
108                    }
109                };
110
111                match event {
112                    Ok(event) => {
113                        log::debug!("Received event {event:?}");
114
115                        if let EventV1::Authenticated = &event {
116                            heartbeat_handle = Some(
117                                tokio::spawn({
118                                    let ws = ws_send.clone();
119                                    let mut i = 0;
120
121                                    async move {
122                                        loop {
123                                            send(
124                                                &ws,
125                                                &ClientMessage::Ping {
126                                                    data: Ping::Number(i),
127                                                    responded: None,
128                                                },
129                                            )
130                                            .await?;
131                                            i = i.wrapping_add(1);
132
133                                            sleep(Duration::from_secs(30)).await;
134                                        }
135
136                                        #[allow(unreachable_code)]
137                                        Ok::<(), Error>(())
138                                    }
139                                })
140                                .abort_handle(),
141                            );
142                        };
143
144                        events.send(event).map_err(|_| Error::InternalError)?;
145                    }
146                    Err(e) => {
147                        log::error!("Failed to deserialise event: {e:?}");
148                    }
149                }
150            }
151
152            if let Some(handle) = heartbeat_handle {
153                handle.abort();
154            };
155
156            Ok::<_, Error>(())
157        }
158    }
159    .boxed();
160
161    let client_server = {
162        let ws_send = ws_send.clone();
163
164        async move {
165            while let Some(message) = client_events.lock().await.recv().await {
166                match message {
167                    EventMessage::Client(message) => send(&ws_send, &message).await?,
168                    EventMessage::Program(ProgramMessage::Close) => return Err(Error::Close),
169                }
170            }
171
172            Ok::<_, Error>(())
173        }
174    }
175    .boxed();
176
177    select(server_client, client_server).await.into_inner().0
178}