Skip to main content

a3s_boot/websocket/
connection.rs

1use 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
13/// Adapter-neutral WebSocket connection.
14pub 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/// In-process WebSocket gateway connection used by adapters and tests.
24#[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}