Skip to main content

a3s_boot/websocket/
handler.rs

1use super::connection::WebSocketGatewayConnection;
2use super::message::{IntoWebSocketReply, WebSocketMessage};
3use super::pipeline::{WebSocketGuard, WebSocketInterceptor, WebSocketPipe};
4use super::server::WebSocketGatewayServer;
5use crate::pipeline::{PipelineComponent, PipelineOverrides};
6use crate::{
7    catch_errors, validate_json_value_with_options, BootError, BootErrorKind, BoxFuture, Result,
8    Validate, ValidationOptions, ValidationSchema, WebSocketExceptionFilter,
9};
10use serde::de::DeserializeOwned;
11use serde::Serialize;
12use serde_json::Value;
13use std::collections::BTreeMap;
14use std::future::Future;
15use std::sync::Arc;
16
17pub(crate) type WebSocketHandlerFuture = BoxFuture<'static, Result<Option<WebSocketMessage>>>;
18type WebSocketMessageValidator =
19    Arc<dyn Fn(WebSocketMessage, ValidationOptions) -> Result<WebSocketMessage> + Send + Sync>;
20
21pub(crate) trait WebSocketMessageHandler: Send + Sync + 'static {
22    fn call(
23        &self,
24        connection: WebSocketGatewayConnection,
25        message: WebSocketMessage,
26    ) -> WebSocketHandlerFuture;
27}
28
29pub(crate) struct WebSocketHandlerAdapter<H> {
30    pub(crate) handler: H,
31}
32
33impl<H, Fut, R> WebSocketMessageHandler for WebSocketHandlerAdapter<H>
34where
35    H: Fn(WebSocketMessage) -> Fut + Send + Sync + 'static,
36    Fut: Future<Output = Result<R>> + Send + 'static,
37    R: IntoWebSocketReply + Send + 'static,
38{
39    fn call(
40        &self,
41        _connection: WebSocketGatewayConnection,
42        message: WebSocketMessage,
43    ) -> WebSocketHandlerFuture {
44        let future = (self.handler)(message);
45        Box::pin(async move { Ok(future.await?.into_websocket_reply()) })
46    }
47}
48
49pub(crate) struct WebSocketConnectionHandlerAdapter<H> {
50    pub(crate) handler: H,
51}
52
53impl<H, Fut, R> WebSocketMessageHandler for WebSocketConnectionHandlerAdapter<H>
54where
55    H: Fn(WebSocketGatewayConnection, WebSocketMessage) -> Fut + Send + Sync + 'static,
56    Fut: Future<Output = Result<R>> + Send + 'static,
57    R: IntoWebSocketReply + Send + 'static,
58{
59    fn call(
60        &self,
61        connection: WebSocketGatewayConnection,
62        message: WebSocketMessage,
63    ) -> WebSocketHandlerFuture {
64        let future = (self.handler)(connection, message);
65        Box::pin(async move { Ok(future.await?.into_websocket_reply()) })
66    }
67}
68
69pub(crate) struct WebSocketServerHandlerAdapter<H> {
70    pub(crate) handler: H,
71}
72
73impl<H, Fut, R> WebSocketMessageHandler for WebSocketServerHandlerAdapter<H>
74where
75    H: Fn(WebSocketGatewayServer, WebSocketMessage) -> Fut + Send + Sync + 'static,
76    Fut: Future<Output = Result<R>> + Send + 'static,
77    R: IntoWebSocketReply + Send + 'static,
78{
79    fn call(
80        &self,
81        connection: WebSocketGatewayConnection,
82        message: WebSocketMessage,
83    ) -> WebSocketHandlerFuture {
84        let future = (self.handler)(connection.server(), message);
85        Box::pin(async move { Ok(future.await?.into_websocket_reply()) })
86    }
87}
88
89/// Handler definition for one WebSocket subscription.
90#[derive(Clone)]
91pub struct WebSocketSubscriptionDefinition {
92    pub(crate) handler: Arc<dyn WebSocketMessageHandler>,
93    pub(crate) pipes: Vec<PipelineComponent<dyn WebSocketPipe>>,
94    pub(crate) guards: Vec<PipelineComponent<dyn WebSocketGuard>>,
95    pub(crate) interceptors: Vec<PipelineComponent<dyn WebSocketInterceptor>>,
96    pub(crate) filters: Vec<PipelineComponent<dyn WebSocketExceptionFilter>>,
97    pub(crate) validators: Vec<WebSocketMessageValidator>,
98    pub(crate) validation_enabled: bool,
99    pub(crate) validation_disabled: bool,
100    pub(crate) validation_options: ValidationOptions,
101    pub(crate) metadata: BTreeMap<String, Value>,
102}
103
104impl WebSocketSubscriptionDefinition {
105    pub fn new<H, Fut, R>(handler: H) -> Self
106    where
107        H: Fn(WebSocketMessage) -> Fut + Send + Sync + 'static,
108        Fut: Future<Output = Result<R>> + Send + 'static,
109        R: IntoWebSocketReply + Send + 'static,
110    {
111        Self {
112            handler: Arc::new(WebSocketHandlerAdapter { handler }),
113            pipes: Vec::new(),
114            guards: Vec::new(),
115            interceptors: Vec::new(),
116            filters: Vec::new(),
117            validators: Vec::new(),
118            validation_enabled: false,
119            validation_disabled: false,
120            validation_options: ValidationOptions::default(),
121            metadata: BTreeMap::new(),
122        }
123    }
124
125    pub fn new_with_connection<H, Fut, R>(handler: H) -> Self
126    where
127        H: Fn(WebSocketGatewayConnection, WebSocketMessage) -> Fut + Send + Sync + 'static,
128        Fut: Future<Output = Result<R>> + Send + 'static,
129        R: IntoWebSocketReply + Send + 'static,
130    {
131        Self {
132            handler: Arc::new(WebSocketConnectionHandlerAdapter { handler }),
133            pipes: Vec::new(),
134            guards: Vec::new(),
135            interceptors: Vec::new(),
136            filters: Vec::new(),
137            validators: Vec::new(),
138            validation_enabled: false,
139            validation_disabled: false,
140            validation_options: ValidationOptions::default(),
141            metadata: BTreeMap::new(),
142        }
143    }
144
145    pub fn new_with_server<H, Fut, R>(handler: H) -> Self
146    where
147        H: Fn(WebSocketGatewayServer, WebSocketMessage) -> Fut + Send + Sync + 'static,
148        Fut: Future<Output = Result<R>> + Send + 'static,
149        R: IntoWebSocketReply + Send + 'static,
150    {
151        Self {
152            handler: Arc::new(WebSocketServerHandlerAdapter { handler }),
153            pipes: Vec::new(),
154            guards: Vec::new(),
155            interceptors: Vec::new(),
156            filters: Vec::new(),
157            validators: Vec::new(),
158            validation_enabled: false,
159            validation_disabled: false,
160            validation_options: ValidationOptions::default(),
161            metadata: BTreeMap::new(),
162        }
163    }
164
165    pub fn metadata(&self) -> &BTreeMap<String, Value> {
166        &self.metadata
167    }
168
169    pub fn metadata_value(&self, key: &str) -> Option<&Value> {
170        self.metadata.get(key)
171    }
172
173    pub fn with_metadata<V>(self, key: impl Into<String>, value: V) -> Result<Self>
174    where
175        V: Serialize,
176    {
177        let key = key.into();
178        let value = serde_json::to_value(value).map_err(|error| {
179            BootError::Internal(format!(
180                "failed to serialize websocket subscription metadata `{key}`: {error}"
181            ))
182        })?;
183        Ok(self.with_metadata_value(key, value))
184    }
185
186    pub fn with_metadata_value(mut self, key: impl Into<String>, value: Value) -> Self {
187        self.metadata.insert(key.into(), value);
188        self
189    }
190
191    pub(crate) fn with_metadata_defaults(mut self, metadata: &BTreeMap<String, Value>) -> Self {
192        for (key, value) in metadata {
193            self.metadata
194                .entry(key.clone())
195                .or_insert_with(|| value.clone());
196        }
197        self
198    }
199
200    pub(crate) fn with_metadata_default_value(
201        mut self,
202        key: impl Into<String>,
203        value: Value,
204    ) -> Self {
205        self.metadata.entry(key.into()).or_insert(value);
206        self
207    }
208
209    pub fn with_pipe<P>(mut self, pipe: P) -> Self
210    where
211        P: WebSocketPipe,
212    {
213        self.pipes
214            .push(PipelineComponent::<dyn WebSocketPipe>::new(pipe));
215        self
216    }
217
218    pub fn with_guard<G>(mut self, guard: G) -> Self
219    where
220        G: WebSocketGuard,
221    {
222        self.guards
223            .push(PipelineComponent::<dyn WebSocketGuard>::new(guard));
224        self
225    }
226
227    pub fn with_interceptor<I>(mut self, interceptor: I) -> Self
228    where
229        I: WebSocketInterceptor,
230    {
231        self.interceptors
232            .push(PipelineComponent::<dyn WebSocketInterceptor>::new(
233                interceptor,
234            ));
235        self
236    }
237
238    pub fn with_filter<F>(mut self, filter: F) -> Self
239    where
240        F: WebSocketExceptionFilter,
241    {
242        self.filters
243            .push(PipelineComponent::<dyn WebSocketExceptionFilter>::new(
244                filter,
245            ));
246        self
247    }
248
249    pub fn with_catch_filter<I, F>(self, kinds: I, filter: F) -> Self
250    where
251        I: IntoIterator<Item = BootErrorKind>,
252        F: WebSocketExceptionFilter,
253    {
254        self.with_filter(catch_errors(kinds, filter))
255    }
256
257    pub(crate) fn with_pipeline_overrides(mut self, overrides: &PipelineOverrides) -> Self {
258        overrides.apply_to_websocket_pipes(&mut self.pipes);
259        overrides.apply_to_websocket_guards(&mut self.guards);
260        overrides.apply_to_websocket_interceptors(&mut self.interceptors);
261        overrides.apply_to_websocket_filters(&mut self.filters);
262        self
263    }
264
265    pub fn with_validation(mut self) -> Self {
266        self.validation_enabled = true;
267        self.validation_disabled = false;
268        self
269    }
270
271    pub fn with_validation_options(mut self, options: ValidationOptions) -> Self {
272        self.validation_enabled = true;
273        self.validation_disabled = false;
274        self.validation_options = self.validation_options.merge(options);
275        self
276    }
277
278    pub fn without_validation(mut self) -> Self {
279        self.validation_enabled = false;
280        self.validation_disabled = true;
281        self
282    }
283
284    pub(crate) fn with_validation_prefix(
285        mut self,
286        validation_enabled: bool,
287        validation_options: ValidationOptions,
288    ) -> Self {
289        if !self.validation_disabled {
290            self.validation_enabled = validation_enabled || self.validation_enabled;
291            self.validation_options = validation_options.merge(self.validation_options);
292        }
293        self
294    }
295
296    pub fn with_payload_validation<T>(mut self) -> Self
297    where
298        T: DeserializeOwned + Validate + 'static,
299    {
300        self.validators.push(Arc::new(|message, _| {
301            message.validated_data::<T>().map(|_| message)
302        }));
303        self.with_validation()
304    }
305
306    pub fn with_payload_validation_options<T>(mut self, options: ValidationOptions) -> Self
307    where
308        T: DeserializeOwned + Serialize + Validate + ValidationSchema + 'static,
309    {
310        self.validators
311            .push(Arc::new(move |mut message, inherited_options| {
312                let options = inherited_options.merge(options);
313                let data = validate_json_value_with_options::<T>(
314                    message.data.clone(),
315                    options,
316                    "message data",
317                )?;
318                if options.transform || options.whitelist {
319                    message.data = data;
320                }
321                Ok(message)
322            }));
323        self.with_validation()
324    }
325}