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#[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}