Skip to main content

rivetkit_core/
websocket.rs

1use std::fmt;
2use std::sync::Arc;
3
4use anyhow::Result;
5use futures::future::BoxFuture;
6use parking_lot::RwLock;
7use rivet_envoy_client::config::WebSocketSender;
8
9use crate::actor::context::WebSocketCallbackRegion;
10use crate::error::ActorRuntime;
11use crate::types::WsMessage;
12
13// Rivet supports a non-standard async close-listener extension for actor
14// WebSockets. Core tracks close-event delivery with the websocket callback
15// region instead of reusing disconnect callbacks because close listeners are
16// WebSocket event work, while `onDisconnect` is connection lifecycle work.
17pub(crate) type WebSocketSendCallback = Arc<dyn Fn(WsMessage) -> Result<()> + Send + Sync>;
18pub(crate) type WebSocketCloseCallback =
19	Arc<dyn Fn(Option<u16>, Option<String>) -> BoxFuture<'static, Result<()>> + Send + Sync>;
20pub(crate) type WebSocketMessageEventCallback =
21	Arc<dyn Fn(WsMessage, Option<u16>) -> Result<()> + Send + Sync>;
22pub(crate) type WebSocketCloseEventCallback =
23	Arc<dyn Fn(u16, String, bool) -> BoxFuture<'static, Result<()>> + Send + Sync>;
24pub(crate) type WebSocketCallbackRegionFactory =
25	Arc<dyn Fn() -> WebSocketCallbackRegion + Send + Sync>;
26
27#[derive(Clone)]
28pub struct WebSocket(Arc<WebSocketInner>);
29
30struct WebSocketInner {
31	// Forced-sync: WebSocket configuration and event dispatch are synchronous
32	// public APIs, so callbacks are cloned out before any async close work.
33	send_callback: RwLock<Option<WebSocketSendCallback>>,
34	close_callback: RwLock<Option<WebSocketCloseCallback>>,
35	message_event_callback: RwLock<Option<WebSocketMessageEventCallback>>,
36	close_event_callback: RwLock<Option<WebSocketCloseEventCallback>>,
37	close_event_callback_region: RwLock<Option<WebSocketCallbackRegionFactory>>,
38}
39
40impl WebSocket {
41	pub fn new() -> Self {
42		Self(Arc::new(WebSocketInner {
43			send_callback: RwLock::new(None),
44			close_callback: RwLock::new(None),
45			message_event_callback: RwLock::new(None),
46			close_event_callback: RwLock::new(None),
47			close_event_callback_region: RwLock::new(None),
48		}))
49	}
50
51	pub fn from_sender(sender: WebSocketSender) -> Self {
52		let websocket = Self::new();
53		websocket.configure_sender(sender);
54		websocket
55	}
56
57	pub fn send(&self, msg: WsMessage) {
58		if let Err(error) = self.try_send(msg) {
59			tracing::error!(?error, "failed to send websocket message");
60		}
61	}
62
63	pub async fn close(&self, code: Option<u16>, reason: Option<String>) {
64		if let Err(error) = self.try_close(code, reason).await {
65			tracing::error!(?error, "failed to close websocket");
66		}
67	}
68
69	pub fn dispatch_message_event(&self, msg: WsMessage, message_index: Option<u16>) {
70		if let Err(error) = self.try_dispatch_message_event(msg, message_index) {
71			tracing::error!(?error, "failed to dispatch websocket message event");
72		}
73	}
74
75	pub async fn dispatch_close_event(&self, code: u16, reason: String, was_clean: bool) {
76		if let Err(error) = self.try_dispatch_close_event(code, reason, was_clean).await {
77			tracing::error!(?error, "failed to dispatch websocket close event");
78		}
79	}
80
81	pub fn configure_sender(&self, sender: WebSocketSender) {
82		let send_sender = sender.clone();
83		let close_sender = sender;
84		self.configure_send_callback(Some(Arc::new(move |message| {
85			match message {
86				WsMessage::Text(text) => send_sender.send_text(&text),
87				WsMessage::Binary(bytes) => send_sender.send(bytes, true),
88			}
89			Ok(())
90		})));
91		self.configure_close_callback(Some(Arc::new(move |code, reason| {
92			let close_sender = close_sender.clone();
93			Box::pin(async move {
94				close_sender.close(code, reason);
95				Ok(())
96			})
97		})));
98	}
99
100	pub(crate) fn configure_send_callback(&self, send_callback: Option<WebSocketSendCallback>) {
101		*self.0.send_callback.write() = send_callback;
102	}
103
104	pub(crate) fn configure_close_callback(&self, close_callback: Option<WebSocketCloseCallback>) {
105		*self.0.close_callback.write() = close_callback;
106	}
107
108	pub fn configure_message_event_callback(
109		&self,
110		message_event_callback: Option<WebSocketMessageEventCallback>,
111	) {
112		*self.0.message_event_callback.write() = message_event_callback;
113	}
114
115	pub fn configure_close_event_callback(
116		&self,
117		close_event_callback: Option<WebSocketCloseEventCallback>,
118	) {
119		*self.0.close_event_callback.write() = close_event_callback;
120	}
121
122	pub(crate) fn configure_close_event_callback_region(
123		&self,
124		close_event_callback_region: Option<WebSocketCallbackRegionFactory>,
125	) {
126		*self.0.close_event_callback_region.write() = close_event_callback_region;
127	}
128
129	pub(crate) fn try_send(&self, msg: WsMessage) -> Result<()> {
130		let callback = self.send_callback()?;
131		callback(msg)
132	}
133
134	pub(crate) async fn try_close(&self, code: Option<u16>, reason: Option<String>) -> Result<()> {
135		let callback = self.close_callback()?;
136		callback(code, reason).await
137	}
138
139	pub(crate) fn try_dispatch_message_event(
140		&self,
141		msg: WsMessage,
142		message_index: Option<u16>,
143	) -> Result<()> {
144		let callback = self.message_event_callback()?;
145		callback(msg, message_index)
146	}
147
148	pub(crate) async fn try_dispatch_close_event(
149		&self,
150		code: u16,
151		reason: String,
152		was_clean: bool,
153	) -> Result<()> {
154		let callback = self.close_event_callback()?;
155		let _region = self.close_event_callback_region().map(|create| create());
156		callback(code, reason, was_clean).await
157	}
158
159	fn send_callback(&self) -> Result<WebSocketSendCallback> {
160		self.0
161			.send_callback
162			.read()
163			.clone()
164			.ok_or_else(|| websocket_not_configured("send callback"))
165	}
166
167	fn close_callback(&self) -> Result<WebSocketCloseCallback> {
168		self.0
169			.close_callback
170			.read()
171			.clone()
172			.ok_or_else(|| websocket_not_configured("close callback"))
173	}
174
175	fn message_event_callback(&self) -> Result<WebSocketMessageEventCallback> {
176		self.0
177			.message_event_callback
178			.read()
179			.clone()
180			.ok_or_else(|| websocket_not_configured("message event callback"))
181	}
182
183	fn close_event_callback(&self) -> Result<WebSocketCloseEventCallback> {
184		self.0
185			.close_event_callback
186			.read()
187			.clone()
188			.ok_or_else(|| websocket_not_configured("close event callback"))
189	}
190
191	fn close_event_callback_region(&self) -> Option<WebSocketCallbackRegionFactory> {
192		self.0.close_event_callback_region.read().clone()
193	}
194}
195
196fn websocket_not_configured(component: &str) -> anyhow::Error {
197	ActorRuntime::NotConfigured {
198		component: format!("websocket {component}"),
199	}
200	.build()
201}
202
203impl Default for WebSocket {
204	fn default() -> Self {
205		Self::new()
206	}
207}
208
209impl fmt::Debug for WebSocket {
210	fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
211		f.debug_struct("WebSocket")
212			.field("send_configured", &self.0.send_callback.read().is_some())
213			.field("close_configured", &self.0.close_callback.read().is_some())
214			.field(
215				"message_event_configured",
216				&self.0.message_event_callback.read().is_some(),
217			)
218			.field(
219				"close_event_configured",
220				&self.0.close_event_callback.read().is_some(),
221			)
222			.field(
223				"close_event_region_configured",
224				&self.0.close_event_callback_region.read().is_some(),
225			)
226			.finish()
227	}
228}
229
230// Test shim keeps moved tests in crate-root tests/ with private-module access.
231#[cfg(test)]
232#[path = "../tests/websocket.rs"]
233mod tests;