rivetkit_core/
websocket.rs1use 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
13pub(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 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#[cfg(test)]
232#[path = "../tests/websocket.rs"]
233mod tests;