Skip to main content

a3s_boot/
testing.rs

1use crate::pipeline::PipelineOverrides;
2use crate::{
3    BootApplication, BootApplicationBuilder, BootError, BootRequest, BootResponse,
4    ControllerDefinition, DynamicModule, ExceptionFilter, Guard, Interceptor,
5    MessagePatternDefinition, Module, ModuleRef, Pipe, ProviderDefinition, ProviderToken, Result,
6    RouteDefinition, TransportExceptionFilter, TransportGuard, TransportInterceptor, TransportPipe,
7    WebSocketExceptionFilter, WebSocketGatewayDefinition, WebSocketGuard, WebSocketInterceptor,
8    WebSocketPipe,
9};
10use std::sync::Arc;
11
12/// Compiled test module with an in-process [`BootApplication`].
13pub struct TestingModule {
14    app: BootApplication,
15}
16
17impl std::fmt::Debug for TestingModule {
18    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
19        f.debug_struct("TestingModule")
20            .field("modules", &self.app.module_names())
21            .field("routes", &self.app.routes().len())
22            .field("gateways", &self.app.gateways().len())
23            .field("message_patterns", &self.app.message_patterns().len())
24            .finish()
25    }
26}
27
28impl TestingModule {
29    pub fn builder() -> TestingModuleBuilder {
30        TestingModuleBuilder::new()
31    }
32
33    pub fn app(&self) -> &BootApplication {
34        &self.app
35    }
36
37    pub fn into_app(self) -> BootApplication {
38        self.app
39    }
40
41    pub fn module_ref(&self) -> &ModuleRef {
42        self.app.module_ref()
43    }
44
45    pub fn get<T>(&self) -> Result<Arc<T>>
46    where
47        T: Send + Sync + 'static,
48    {
49        self.get_optional::<T>()?
50            .ok_or_else(|| BootError::MissingProvider(ProviderToken::of::<T>().to_string()))
51    }
52
53    pub fn get_named<T>(&self, token: &str) -> Result<Arc<T>>
54    where
55        T: Send + Sync + 'static,
56    {
57        self.get_optional_named::<T>(token)?
58            .ok_or_else(|| BootError::MissingProvider(ProviderToken::named(token).to_string()))
59    }
60
61    pub fn get_optional<T>(&self) -> Result<Option<Arc<T>>>
62    where
63        T: Send + Sync + 'static,
64    {
65        if let Some(value) = self.app.get_optional::<T>()? {
66            return Ok(Some(value));
67        }
68
69        for instance in &self.app.module_instances {
70            if let Some(value) = instance.module_ref.get_optional::<T>()? {
71                return Ok(Some(value));
72            }
73        }
74
75        Ok(None)
76    }
77
78    pub fn get_optional_named<T>(&self, token: &str) -> Result<Option<Arc<T>>>
79    where
80        T: Send + Sync + 'static,
81    {
82        if let Some(value) = self.app.get_optional_named::<T>(token)? {
83            return Ok(Some(value));
84        }
85
86        for instance in &self.app.module_instances {
87            if let Some(value) = instance.module_ref.get_optional_named::<T>(token)? {
88                return Ok(Some(value));
89            }
90        }
91
92        Ok(None)
93    }
94
95    pub async fn call(&self, request: BootRequest) -> Result<BootResponse> {
96        self.app.call(request).await
97    }
98}
99
100/// Builder for Nest-style test modules.
101pub struct TestingModuleBuilder {
102    app: BootApplicationBuilder,
103    module: DynamicModule,
104    pipeline_overrides: PipelineOverrides,
105}
106
107impl Default for TestingModuleBuilder {
108    fn default() -> Self {
109        Self {
110            app: BootApplication::builder(),
111            module: DynamicModule::new("TestingModule"),
112            pipeline_overrides: PipelineOverrides::default(),
113        }
114    }
115}
116
117impl TestingModuleBuilder {
118    pub fn new() -> Self {
119        Self::default()
120    }
121
122    pub fn import<M>(mut self, module: M) -> Self
123    where
124        M: Module,
125    {
126        self.module = self.module.import(module);
127        self
128    }
129
130    pub fn import_arc(mut self, module: Arc<dyn Module>) -> Self {
131        self.module = self.module.import_arc(module);
132        self
133    }
134
135    pub fn provider(mut self, provider: ProviderDefinition) -> Self {
136        self.module = self.module.provider(provider);
137        self
138    }
139
140    pub fn controller(mut self, controller: ControllerDefinition) -> Self {
141        self.module = self.module.controller(controller);
142        self
143    }
144
145    pub fn route(mut self, route: RouteDefinition) -> Self {
146        self.module = self.module.route(route);
147        self
148    }
149
150    pub fn gateway(mut self, gateway: WebSocketGatewayDefinition) -> Self {
151        self.module = self.module.gateway(gateway);
152        self
153    }
154
155    pub fn message_pattern(mut self, pattern: MessagePatternDefinition) -> Self {
156        self.module = self.module.message_pattern(pattern);
157        self
158    }
159
160    pub fn override_provider(mut self, provider: ProviderDefinition) -> Self {
161        self.app = self.app.override_provider(provider);
162        self
163    }
164
165    pub fn override_module<M>(mut self, target_name: impl Into<String>, module: M) -> Self
166    where
167        M: Module,
168    {
169        self.app = self.app.override_module(target_name, module);
170        self
171    }
172
173    pub fn override_module_arc(
174        mut self,
175        target_name: impl Into<String>,
176        module: Arc<dyn Module>,
177    ) -> Self {
178        self.app = self.app.override_module_arc(target_name, module);
179        self
180    }
181
182    pub fn override_guard<T, G>(mut self, guard: G) -> Self
183    where
184        T: Guard,
185        G: Guard,
186    {
187        self.pipeline_overrides.override_guard::<T, G>(guard);
188        self
189    }
190
191    pub fn override_interceptor<T, I>(mut self, interceptor: I) -> Self
192    where
193        T: Interceptor,
194        I: Interceptor,
195    {
196        self.pipeline_overrides
197            .override_interceptor::<T, I>(interceptor);
198        self
199    }
200
201    pub fn override_filter<T, F>(mut self, filter: F) -> Self
202    where
203        T: ExceptionFilter,
204        F: ExceptionFilter,
205    {
206        self.pipeline_overrides.override_filter::<T, F>(filter);
207        self
208    }
209
210    pub fn override_pipe<T, P>(mut self, pipe: P) -> Self
211    where
212        T: Pipe,
213        P: Pipe,
214    {
215        self.pipeline_overrides.override_pipe::<T, P>(pipe);
216        self
217    }
218
219    pub fn override_websocket_pipe<T, P>(mut self, pipe: P) -> Self
220    where
221        T: WebSocketPipe,
222        P: WebSocketPipe,
223    {
224        self.pipeline_overrides
225            .override_websocket_pipe::<T, P>(pipe);
226        self
227    }
228
229    pub fn override_websocket_guard<T, G>(mut self, guard: G) -> Self
230    where
231        T: WebSocketGuard,
232        G: WebSocketGuard,
233    {
234        self.pipeline_overrides
235            .override_websocket_guard::<T, G>(guard);
236        self
237    }
238
239    pub fn override_websocket_interceptor<T, I>(mut self, interceptor: I) -> Self
240    where
241        T: WebSocketInterceptor,
242        I: WebSocketInterceptor,
243    {
244        self.pipeline_overrides
245            .override_websocket_interceptor::<T, I>(interceptor);
246        self
247    }
248
249    pub fn override_websocket_filter<T, F>(mut self, filter: F) -> Self
250    where
251        T: WebSocketExceptionFilter,
252        F: WebSocketExceptionFilter,
253    {
254        self.pipeline_overrides
255            .override_websocket_filter::<T, F>(filter);
256        self
257    }
258
259    pub fn override_transport_pipe<T, P>(mut self, pipe: P) -> Self
260    where
261        T: TransportPipe,
262        P: TransportPipe,
263    {
264        self.pipeline_overrides
265            .override_transport_pipe::<T, P>(pipe);
266        self
267    }
268
269    pub fn override_transport_guard<T, G>(mut self, guard: G) -> Self
270    where
271        T: TransportGuard,
272        G: TransportGuard,
273    {
274        self.pipeline_overrides
275            .override_transport_guard::<T, G>(guard);
276        self
277    }
278
279    pub fn override_transport_interceptor<T, I>(mut self, interceptor: I) -> Self
280    where
281        T: TransportInterceptor,
282        I: TransportInterceptor,
283    {
284        self.pipeline_overrides
285            .override_transport_interceptor::<T, I>(interceptor);
286        self
287    }
288
289    pub fn override_transport_filter<T, F>(mut self, filter: F) -> Self
290    where
291        T: TransportExceptionFilter,
292        F: TransportExceptionFilter,
293    {
294        self.pipeline_overrides
295            .override_transport_filter::<T, F>(filter);
296        self
297    }
298
299    pub fn compile(self) -> Result<TestingModule> {
300        let mut app = self.app.import(self.module).build()?;
301        if !self.pipeline_overrides.is_empty() {
302            app.routes = app
303                .routes
304                .into_iter()
305                .map(|route| route.with_pipeline_overrides(&self.pipeline_overrides))
306                .collect();
307            app.gateways = app
308                .gateways
309                .into_iter()
310                .map(|gateway| gateway.with_pipeline_overrides(&self.pipeline_overrides))
311                .collect();
312            app.message_patterns = app
313                .message_patterns
314                .into_iter()
315                .map(|pattern| pattern.with_pipeline_overrides(&self.pipeline_overrides))
316                .collect();
317        }
318        Ok(TestingModule { app })
319    }
320
321    pub async fn compile_async(self) -> Result<TestingModule> {
322        let mut app = self.app.import(self.module).build_async().await?;
323        if !self.pipeline_overrides.is_empty() {
324            app.routes = app
325                .routes
326                .into_iter()
327                .map(|route| route.with_pipeline_overrides(&self.pipeline_overrides))
328                .collect();
329            app.gateways = app
330                .gateways
331                .into_iter()
332                .map(|gateway| gateway.with_pipeline_overrides(&self.pipeline_overrides))
333                .collect();
334            app.message_patterns = app
335                .message_patterns
336                .into_iter()
337                .map(|pattern| pattern.with_pipeline_overrides(&self.pipeline_overrides))
338                .collect();
339        }
340        Ok(TestingModule { app })
341    }
342}