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::{BootError, BootRequest, BoxFuture, Result};
7use std::sync::atomic::{AtomicBool, Ordering};
8use std::sync::Arc;
9
10pub 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#[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}