a3s_boot/websocket/
connection.rs1use super::context::WebSocketContext;
2use super::gateway::WebSocketGatewayDefinition;
3use super::message::{send_to_outbounds, WebSocketMessage, WebSocketOutbound};
4use super::server::WebSocketGatewayServer;
5use super::state::normalize_room;
6use crate::{
7 BootError, BootRequest, BoxFuture, CallHandler, ContextId, ContextIdFactory, Result,
8 WebSocketInterceptor,
9};
10use std::sync::atomic::{AtomicBool, Ordering};
11use std::sync::Arc;
12
13pub trait WebSocketConnection: Send + Sync {
15 fn request(&self) -> &BootRequest;
16
17 fn dispatch(
18 &self,
19 message: WebSocketMessage,
20 ) -> BoxFuture<'static, Result<Option<WebSocketMessage>>>;
21}
22
23#[derive(Clone)]
25pub struct WebSocketGatewayConnection {
26 pub(crate) gateway: WebSocketGatewayDefinition,
27 pub(crate) id: u64,
28 pub(crate) request: BootRequest,
29 pub(crate) outbound: Option<Arc<dyn WebSocketOutbound>>,
30 pub(crate) opened: Arc<AtomicBool>,
31}
32
33impl WebSocketGatewayConnection {
34 pub fn id(&self) -> u64 {
35 self.id
36 }
37
38 pub fn request(&self) -> &BootRequest {
39 &self.request
40 }
41
42 pub fn namespace(&self) -> Option<&str> {
43 self.gateway.namespace()
44 }
45
46 pub fn server(&self) -> WebSocketGatewayServer {
47 self.gateway.server()
48 }
49
50 pub fn rooms(&self) -> Result<Vec<String>> {
51 self.gateway.state.rooms_for_connection(self.id)
52 }
53
54 pub fn join(&self, room: impl Into<String>) -> Result<()> {
55 self.gateway.state.join(self.id, room)
56 }
57
58 pub fn leave(&self, room: impl Into<String>) -> Result<()> {
59 self.gateway.state.leave(self.id, room)
60 }
61
62 pub async fn emit(&self, message: WebSocketMessage) -> Result<bool> {
63 self.gateway.emit_to_connection(self.id, message).await
64 }
65
66 pub async fn broadcast(&self, message: WebSocketMessage) -> Result<usize> {
67 let outbounds = self.gateway.state.broadcast_targets(None, Some(self.id))?;
68 send_to_outbounds(outbounds, message).await
69 }
70
71 pub async fn broadcast_to_room(
72 &self,
73 room: impl Into<String>,
74 message: WebSocketMessage,
75 ) -> Result<usize> {
76 let room = normalize_room(room)?;
77 let outbounds = self
78 .gateway
79 .state
80 .broadcast_targets(Some(&room), Some(self.id))?;
81 send_to_outbounds(outbounds, message).await
82 }
83
84 pub async fn open(&self) -> Result<()> {
85 if self.opened.swap(true, Ordering::AcqRel) {
86 return Ok(());
87 }
88 self.gateway
89 .state
90 .register(self.id, self.outbound.clone())?;
91 let mut hook_result = Ok(());
92 for hook in &self.gateway.connection_hooks {
93 if let Err(error) = hook.handle_connection(self.clone()).await {
94 hook_result = Err(error);
95 break;
96 }
97 }
98 if hook_result.is_err() {
99 self.opened.store(false, Ordering::Release);
100 self.gateway.state.unregister(self.id)?;
101 }
102 hook_result
103 }
104
105 pub async fn close(&self) -> Result<()> {
106 if !self.opened.swap(false, Ordering::AcqRel) {
107 return Ok(());
108 }
109 let mut hook_result = Ok(());
110 for hook in self.gateway.disconnect_hooks.iter().rev() {
111 if let Err(error) = hook.handle_disconnect(self.clone()).await {
112 hook_result = Err(error);
113 break;
114 }
115 }
116 let unregister_result = self.gateway.state.unregister(self.id);
117 hook_result?;
118 unregister_result
119 }
120
121 pub async fn dispatch(&self, message: WebSocketMessage) -> Result<Option<WebSocketMessage>> {
122 let context_id = ContextIdFactory::create();
123 let mut connection = self.clone();
124 if let Some(module_ref) = &connection.gateway.module_ref {
125 connection.request = connection
126 .request
127 .with_module_ref(module_ref.context_scope(&context_id));
128 }
129 let context = WebSocketContext::new(
130 &connection.gateway,
131 connection.request.clone(),
132 &message.event,
133 );
134 match connection
135 .dispatch_pipeline(message, context.clone(), context_id.clone())
136 .await
137 {
138 Ok(reply) => Ok(reply),
139 Err(error) => connection.handle_error(context, error, &context_id).await,
140 }
141 }
142
143 async fn dispatch_pipeline(
144 &self,
145 message: WebSocketMessage,
146 context: WebSocketContext,
147 context_id: ContextId,
148 ) -> Result<Option<WebSocketMessage>> {
149 let event = message.event.clone();
150 let handler = self.gateway.handlers.get(&event).cloned().ok_or_else(|| {
151 BootError::NotFound(format!("websocket event {} {}", self.gateway.path, event))
152 })?;
153
154 for guard in &self.gateway.guards {
155 let can_activate = guard
156 .resolve(&context_id)?
157 .can_activate(context.clone())
158 .await?;
159 if !can_activate {
160 return Err(BootError::Forbidden(format!(
161 "websocket event {} {}",
162 self.gateway.path, message.event
163 )));
164 }
165 }
166 for guard in &handler.guards {
167 let can_activate = guard
168 .resolve(&context_id)?
169 .can_activate(context.clone())
170 .await?;
171 if !can_activate {
172 return Err(BootError::Forbidden(format!(
173 "websocket event {} {}",
174 self.gateway.path, message.event
175 )));
176 }
177 }
178
179 let interceptors = self
180 .gateway
181 .interceptors
182 .iter()
183 .chain(handler.interceptors.iter())
184 .map(|interceptor| interceptor.resolve(&context_id))
185 .collect::<Result<Vec<_>>>()?;
186
187 let connection = self.clone();
188 let replay_handler = handler.clone();
189 let replay_message = message.clone();
190 let replay_context_id = context_id.clone();
191 let terminal = CallHandler::from_fn(move || {
192 let connection = connection.clone();
193 let handler = replay_handler.clone();
194 let mut message = replay_message.clone();
195 let context_id = replay_context_id.clone();
196 async move {
197 for pipe in &connection.gateway.pipes {
198 message = pipe.resolve(&context_id)?.transform(message).await?;
199 }
200 for pipe in &handler.pipes {
201 message = pipe.resolve(&context_id)?.transform(message).await?;
202 }
203
204 if handler.validation_enabled {
205 for validator in &handler.validators {
206 message = validator(message, handler.validation_options)?;
207 }
208 }
209
210 handler.handler.call(connection, message).await
211 }
212 });
213
214 run_interceptor_chain(&interceptors, context, terminal).await
215 }
216
217 async fn handle_error(
218 &self,
219 context: WebSocketContext,
220 error: BootError,
221 context_id: &ContextId,
222 ) -> Result<Option<WebSocketMessage>> {
223 if let Some(handler) = self.gateway.handlers.get(&context.event) {
224 for filter in handler.filters.iter().rev() {
225 let filter = filter.resolve(context_id)?;
226 if let Some(response) = filter
227 .catch(context.clone(), error.clone_for_filter())
228 .await?
229 {
230 return Ok(response.into_message());
231 }
232 }
233 }
234 for filter in self.gateway.filters.iter().rev() {
235 let filter = filter.resolve(context_id)?;
236 if let Some(response) = filter
237 .catch(context.clone(), error.clone_for_filter())
238 .await?
239 {
240 return Ok(response.into_message());
241 }
242 }
243 Err(error)
244 }
245}
246
247fn run_interceptor_chain<'a>(
248 interceptors: &'a [Arc<dyn WebSocketInterceptor>],
249 context: WebSocketContext,
250 terminal: CallHandler<'a, Option<WebSocketMessage>>,
251) -> BoxFuture<'a, Result<Option<WebSocketMessage>>> {
252 let Some((interceptor, remaining)) = interceptors.split_first() else {
253 return terminal.handle();
254 };
255
256 let next_context = context.clone();
257 let next_terminal = terminal.clone();
258 let next = CallHandler::from_fn(move || {
259 run_interceptor_chain(remaining, next_context.clone(), next_terminal.clone())
260 });
261 interceptor.intercept(context, next)
262}
263
264impl WebSocketConnection for WebSocketGatewayConnection {
265 fn request(&self) -> &BootRequest {
266 self.request()
267 }
268
269 fn dispatch(
270 &self,
271 message: WebSocketMessage,
272 ) -> BoxFuture<'static, Result<Option<WebSocketMessage>>> {
273 let connection = self.clone();
274 Box::pin(async move { connection.dispatch(message).await })
275 }
276}