rosu_render/websocket/
mod.rs1#![cfg(any(
2 feature = "native-tls",
3 feature = "rustls-native-roots",
4 feature = "rustls-webpki-roots"
5))]
6
7use bytes::Bytes;
8
9use crate::WebsocketError;
10
11use self::{
12 engineio::{
13 packet::{Packet as EnginePacket, PacketId as EnginePacketId},
14 EngineIo,
15 },
16 event::RawEvent,
17 packet::{Packet, PacketKind},
18 reconnect::Reconnect,
19};
20
21mod engineio;
22mod packet;
23mod reconnect;
24
25pub mod error;
26pub mod event;
27
28pub struct OrdrWebsocket {
38 engineio: EngineIo,
39 reconnect: Reconnect,
40 auth: Option<Box<str>>,
41}
42
43impl OrdrWebsocket {
44 pub async fn connect() -> Result<Self, WebsocketError> {
46 let engineio = EngineIo::connect().await?;
47
48 let mut this = Self {
49 engineio,
50 reconnect: Reconnect::default(),
51 auth: None,
52 };
53
54 this.open().await?;
55
56 Ok(this)
57 }
58
59 pub async fn next_event(&mut self) -> Result<RawEvent, WebsocketError> {
65 loop {
66 let Some(bytes) = self.engineio.next_message().await? else {
67 self.reconnect().await?;
68
69 continue;
70 };
71
72 let packet = Packet::from_bytes(&bytes)?;
73
74 match packet.kind {
75 PacketKind::Event => {}
76 PacketKind::Ack => self.ack(&packet).await?,
77 PacketKind::Connect => continue,
78 PacketKind::Disconnect | PacketKind::ConnectError => {
79 self.reconnect().await?;
80
81 continue;
82 }
83 }
84
85 if let Some(data) = packet.data {
86 if data.starts_with(b"[\"bot_auth\"") {
89 continue;
90 }
91
92 return RawEvent::from_bytes(data);
93 }
94 }
95 }
96
97 pub async fn authenticate(&mut self, key: &str) -> Result<(), WebsocketError> {
107 self.auth = Some(key.into());
108
109 self.emit_auth(key).await?;
110
111 loop {
112 let Some(bytes) = self.engineio.next_message().await? else {
113 self.reconnect().await?;
114
115 continue;
116 };
117
118 let packet = Packet::from_bytes(&bytes)?;
119
120 match packet.kind {
121 PacketKind::Event => {
122 let Some(data) = packet.data else {
125 continue;
126 };
127
128 let Some(message) = bot_auth_message(&data) else {
129 continue;
130 };
131
132 if message.starts_with("Authentication successful") {
133 return Ok(());
134 }
135
136 return Err(WebsocketError::BotAuth { message });
137 }
138 PacketKind::Ack => self.ack(&packet).await?,
139 PacketKind::Connect => {}
140 PacketKind::Disconnect | PacketKind::ConnectError => self.reconnect().await?,
141 }
142 }
143 }
144
145 pub async fn disconnect(self) -> Result<(), WebsocketError> {
147 self.engineio
148 .disconnect()
149 .await
150 .map_err(WebsocketError::EngineIo)
151 }
152
153 async fn reconnect(&mut self) -> Result<(), WebsocketError> {
154 if let Some(delay) = self.reconnect.delay() {
155 trace!(?delay, "Delaying reconnect...");
156 tokio::time::sleep(delay).await;
157 }
158
159 let err = match self.engineio.reconnect().await {
160 Ok(()) => match self.open().await {
161 Ok(()) => {
162 self.reconnect.reset();
163
164 match self.auth.clone() {
170 Some(key) => match self.emit_auth(&key).await {
171 Ok(()) => return Ok(()),
172 Err(err) => err,
173 },
174 None => return Ok(()),
175 }
176 }
177 Err(err) => err,
178 },
179 Err(err) => WebsocketError::EngineIo(err),
180 };
181
182 self.reconnect.backoff();
183
184 Err(err)
185 }
186
187 async fn emit_auth(&mut self, key: &str) -> Result<(), WebsocketError> {
188 let payload = serde_json::to_string(&["bot_auth", key])
189 .expect("a &str always serializes to a JSON string");
190
191 self.emit(Packet::new_event(Bytes::from(payload))).await
192 }
193
194 async fn emit(&mut self, packet: Packet) -> Result<(), WebsocketError> {
195 let msg = EnginePacket::new(EnginePacketId::Message, packet.to_bytes());
196
197 self.engineio
198 .emit(msg)
199 .await
200 .map_err(WebsocketError::EngineIo)
201 }
202
203 async fn open(&mut self) -> Result<(), WebsocketError> {
204 self.emit(Packet::new(PacketKind::Connect, None)).await
205 }
206
207 async fn ack(&mut self, packet: &Packet) -> Result<(), WebsocketError> {
208 let Some(id) = packet.id else { return Ok(()) };
209
210 self.emit(Packet::new_ack(id)).await
211 }
212}
213
214fn bot_auth_message(data: &[u8]) -> Option<Box<str>> {
219 let (name, message) = serde_json::from_slice::<(String, Box<str>)>(data).ok()?;
220
221 (name == "bot_auth").then_some(message)
222}
223
224#[cfg(test)]
225mod tests {
226 use bytes::Bytes;
227
228 use super::{bot_auth_message, Packet};
229
230 #[test]
231 fn bot_auth_emit_frame() {
232 let payload = serde_json::to_string(&["bot_auth", "secret-key"]).unwrap();
233
234 let frame = Packet::new_event(Bytes::from(payload)).to_bytes();
235
236 assert_eq!(frame.as_ref(), &br#"2["bot_auth","secret-key"]"#[..]);
237 }
238
239 #[test]
240 fn bot_auth_reply_frame() {
241 let frame = br#"2["bot_auth","Authentication successful for bathbot"]"#;
242
243 let packet = Packet::from_bytes(&Bytes::from_static(frame)).unwrap();
244 let data = packet.data.unwrap();
245
246 assert_eq!(
247 bot_auth_message(&data),
248 Some("Authentication successful for bathbot".into()),
249 );
250 }
251
252 #[test]
253 fn bot_auth_failure_message_is_captured() {
254 assert_eq!(
255 bot_auth_message(br#"["bot_auth","Invalid API key"]"#),
256 Some("Invalid API key".into()),
257 );
258 }
259
260 #[test]
261 fn bot_auth_ignores_other_events() {
262 let data = br#"["render_done_json",{"renderID":1}]"#;
263
264 assert_eq!(bot_auth_message(data), None);
265 }
266}