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::{BootError, BootRequest, BoxFuture, Result};
7use std::sync::atomic::{AtomicBool, Ordering};
8use std::sync::Arc;
9
10/// Adapter-neutral WebSocket connection.
11pub trait WebSocketConnection: Send + Sync {
12    fn request(&self) -> &BootRequest;
13
14    fn dispatch(
15        &self,
16        message: WebSocketMessage,
17    ) -> BoxFuture<'static, Result<Option<WebSocketMessage>>>;
18}
19
20/// In-process WebSocket gateway connection used by adapters and tests.
21#[derive(Clone)]
22pub struct WebSocketGatewayConnection {
23    pub(crate) gateway: WebSocketGatewayDefinition,
24    pub(crate) id: u64,
25    pub(crate) request: BootRequest,
26    pub(crate) outbound: Option<Arc<dyn WebSocketOutbound>>,
27    pub(crate) opened: Arc<AtomicBool>,
28}
29
30impl WebSocketGatewayConnection {
31    pub fn id(&self) -> u64 {
32        self.id
33    }
34
35    pub fn request(&self) -> &BootRequest {
36        &self.request
37    }
38
39    pub fn namespace(&self) -> Option<&str> {
40        self.gateway.namespace()
41    }
42
43    pub fn server(&self) -> WebSocketGatewayServer {
44        self.gateway.server()
45    }
46
47    pub fn rooms(&self) -> Result<Vec<String>> {
48        self.gateway.state.rooms_for_connection(self.id)
49    }
50
51    pub fn join(&self, room: impl Into<String>) -> Result<()> {
52        self.gateway.state.join(self.id, room)
53    }
54
55    pub fn leave(&self, room: impl Into<String>) -> Result<()> {
56        self.gateway.state.leave(self.id, room)
57    }
58
59    pub async fn emit(&self, message: WebSocketMessage) -> Result<bool> {
60        self.gateway.emit_to_connection(self.id, message).await
61    }
62
63    pub async fn broadcast(&self, message: WebSocketMessage) -> Result<usize> {
64        let outbounds = self.gateway.state.broadcast_targets(None, Some(self.id))?;
65        send_to_outbounds(outbounds, message).await
66    }
67
68    pub async fn broadcast_to_room(
69        &self,
70        room: impl Into<String>,
71        message: WebSocketMessage,
72    ) -> Result<usize> {
73        let room = normalize_room(room)?;
74        let outbounds = self
75            .gateway
76            .state
77            .broadcast_targets(Some(&room), Some(self.id))?;
78        send_to_outbounds(outbounds, message).await
79    }
80
81    pub async fn open(&self) -> Result<()> {
82        if self.opened.swap(true, Ordering::AcqRel) {
83            return Ok(());
84        }
85        self.gateway
86            .state
87            .register(self.id, self.outbound.clone())?;
88        let mut hook_result = Ok(());
89        for hook in &self.gateway.connection_hooks {
90            if let Err(error) = hook.handle_connection(self.clone()).await {
91                hook_result = Err(error);
92                break;
93            }
94        }
95        if hook_result.is_err() {
96            self.opened.store(false, Ordering::Release);
97            self.gateway.state.unregister(self.id)?;
98        }
99        hook_result
100    }
101
102    pub async fn close(&self) -> Result<()> {
103        if !self.opened.swap(false, Ordering::AcqRel) {
104            return Ok(());
105        }
106        let mut hook_result = Ok(());
107        for hook in self.gateway.disconnect_hooks.iter().rev() {
108            if let Err(error) = hook.handle_disconnect(self.clone()).await {
109                hook_result = Err(error);
110                break;
111            }
112        }
113        let unregister_result = self.gateway.state.unregister(self.id);
114        hook_result?;
115        unregister_result
116    }
117
118    pub async fn dispatch(&self, message: WebSocketMessage) -> Result<Option<WebSocketMessage>> {
119        let context = WebSocketContext::new(&self.gateway, self.request.clone(), &message.event);
120        match self.dispatch_pipeline(message, context.clone()).await {
121            Ok(reply) => Ok(reply),
122            Err(error) => self.handle_error(context, error).await,
123        }
124    }
125
126    async fn dispatch_pipeline(
127        &self,
128        mut message: WebSocketMessage,
129        context: WebSocketContext,
130    ) -> Result<Option<WebSocketMessage>> {
131        let event = message.event.clone();
132        let handler = self.gateway.handlers.get(&event).cloned().ok_or_else(|| {
133            BootError::NotFound(format!("websocket event {} {}", self.gateway.path, event))
134        })?;
135
136        for guard in &self.gateway.guards {
137            let can_activate = guard.inner().can_activate(context.clone()).await?;
138            if !can_activate {
139                return Err(BootError::Forbidden(format!(
140                    "websocket event {} {}",
141                    self.gateway.path, message.event
142                )));
143            }
144        }
145        for guard in &handler.guards {
146            let can_activate = guard.inner().can_activate(context.clone()).await?;
147            if !can_activate {
148                return Err(BootError::Forbidden(format!(
149                    "websocket event {} {}",
150                    self.gateway.path, message.event
151                )));
152            }
153        }
154
155        for interceptor in &self.gateway.interceptors {
156            interceptor.inner().before(context.clone()).await?;
157        }
158        for interceptor in &handler.interceptors {
159            interceptor.inner().before(context.clone()).await?;
160        }
161
162        for pipe in &self.gateway.pipes {
163            message = pipe.inner().transform(message).await?;
164        }
165        for pipe in &handler.pipes {
166            message = pipe.inner().transform(message).await?;
167        }
168
169        if handler.validation_enabled {
170            for validator in &handler.validators {
171                message = validator(message, handler.validation_options)?;
172            }
173        }
174
175        let mut reply = handler.handler.call(self.clone(), message).await?;
176        for interceptor in handler.interceptors.iter().rev() {
177            reply = interceptor.inner().after(context.clone(), reply).await?;
178        }
179        for interceptor in self.gateway.interceptors.iter().rev() {
180            reply = interceptor.inner().after(context.clone(), reply).await?;
181        }
182        Ok(reply)
183    }
184
185    async fn handle_error(
186        &self,
187        context: WebSocketContext,
188        error: BootError,
189    ) -> Result<Option<WebSocketMessage>> {
190        if let Some(handler) = self.gateway.handlers.get(&context.event) {
191            for filter in handler.filters.iter().rev() {
192                if let Some(response) = filter
193                    .inner()
194                    .catch(context.clone(), error.clone_for_filter())
195                    .await?
196                {
197                    return Ok(response.into_message());
198                }
199            }
200        }
201        for filter in self.gateway.filters.iter().rev() {
202            if let Some(response) = filter
203                .inner()
204                .catch(context.clone(), error.clone_for_filter())
205                .await?
206            {
207                return Ok(response.into_message());
208            }
209        }
210        Err(error)
211    }
212}
213
214impl WebSocketConnection for WebSocketGatewayConnection {
215    fn request(&self) -> &BootRequest {
216        self.request()
217    }
218
219    fn dispatch(
220        &self,
221        message: WebSocketMessage,
222    ) -> BoxFuture<'static, Result<Option<WebSocketMessage>>> {
223        let connection = self.clone();
224        Box::pin(async move { connection.dispatch(message).await })
225    }
226}