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}