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
12pub 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
100pub 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}