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, WebSocketGatewayDefinition,
7};
8use std::sync::Arc;
9
10/// Compiled test module with an in-process [`BootApplication`].
11pub struct TestingModule {
12    app: BootApplication,
13}
14
15impl std::fmt::Debug for TestingModule {
16    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
17        f.debug_struct("TestingModule")
18            .field("modules", &self.app.module_names())
19            .field("routes", &self.app.routes().len())
20            .field("gateways", &self.app.gateways().len())
21            .field("message_patterns", &self.app.message_patterns().len())
22            .finish()
23    }
24}
25
26impl TestingModule {
27    pub fn builder() -> TestingModuleBuilder {
28        TestingModuleBuilder::new()
29    }
30
31    pub fn app(&self) -> &BootApplication {
32        &self.app
33    }
34
35    pub fn into_app(self) -> BootApplication {
36        self.app
37    }
38
39    pub fn module_ref(&self) -> &ModuleRef {
40        self.app.module_ref()
41    }
42
43    pub fn get<T>(&self) -> Result<Arc<T>>
44    where
45        T: Send + Sync + 'static,
46    {
47        self.get_optional::<T>()?
48            .ok_or_else(|| BootError::MissingProvider(ProviderToken::of::<T>().to_string()))
49    }
50
51    pub fn get_named<T>(&self, token: &str) -> Result<Arc<T>>
52    where
53        T: Send + Sync + 'static,
54    {
55        self.get_optional_named::<T>(token)?
56            .ok_or_else(|| BootError::MissingProvider(ProviderToken::named(token).to_string()))
57    }
58
59    pub fn get_optional<T>(&self) -> Result<Option<Arc<T>>>
60    where
61        T: Send + Sync + 'static,
62    {
63        if let Some(value) = self.app.get_optional::<T>()? {
64            return Ok(Some(value));
65        }
66
67        for instance in &self.app.module_instances {
68            if let Some(value) = instance.module_ref.get_optional::<T>()? {
69                return Ok(Some(value));
70            }
71        }
72
73        Ok(None)
74    }
75
76    pub fn get_optional_named<T>(&self, token: &str) -> Result<Option<Arc<T>>>
77    where
78        T: Send + Sync + 'static,
79    {
80        if let Some(value) = self.app.get_optional_named::<T>(token)? {
81            return Ok(Some(value));
82        }
83
84        for instance in &self.app.module_instances {
85            if let Some(value) = instance.module_ref.get_optional_named::<T>(token)? {
86                return Ok(Some(value));
87            }
88        }
89
90        Ok(None)
91    }
92
93    pub async fn call(&self, request: BootRequest) -> Result<BootResponse> {
94        self.app.call(request).await
95    }
96}
97
98/// Builder for Nest-style test modules.
99pub struct TestingModuleBuilder {
100    app: BootApplicationBuilder,
101    module: DynamicModule,
102    pipeline_overrides: PipelineOverrides,
103}
104
105impl Default for TestingModuleBuilder {
106    fn default() -> Self {
107        Self {
108            app: BootApplication::builder(),
109            module: DynamicModule::new("TestingModule"),
110            pipeline_overrides: PipelineOverrides::default(),
111        }
112    }
113}
114
115impl TestingModuleBuilder {
116    pub fn new() -> Self {
117        Self::default()
118    }
119
120    pub fn import<M>(mut self, module: M) -> Self
121    where
122        M: Module,
123    {
124        self.module = self.module.import(module);
125        self
126    }
127
128    pub fn import_arc(mut self, module: Arc<dyn Module>) -> Self {
129        self.module = self.module.import_arc(module);
130        self
131    }
132
133    pub fn provider(mut self, provider: ProviderDefinition) -> Self {
134        self.module = self.module.provider(provider);
135        self
136    }
137
138    pub fn controller(mut self, controller: ControllerDefinition) -> Self {
139        self.module = self.module.controller(controller);
140        self
141    }
142
143    pub fn route(mut self, route: RouteDefinition) -> Self {
144        self.module = self.module.route(route);
145        self
146    }
147
148    pub fn gateway(mut self, gateway: WebSocketGatewayDefinition) -> Self {
149        self.module = self.module.gateway(gateway);
150        self
151    }
152
153    pub fn message_pattern(mut self, pattern: MessagePatternDefinition) -> Self {
154        self.module = self.module.message_pattern(pattern);
155        self
156    }
157
158    pub fn override_provider(mut self, provider: ProviderDefinition) -> Self {
159        self.app = self.app.override_provider(provider);
160        self
161    }
162
163    pub fn override_guard<T, G>(mut self, guard: G) -> Self
164    where
165        T: Guard,
166        G: Guard,
167    {
168        self.pipeline_overrides.override_guard::<T, G>(guard);
169        self
170    }
171
172    pub fn override_interceptor<T, I>(mut self, interceptor: I) -> Self
173    where
174        T: Interceptor,
175        I: Interceptor,
176    {
177        self.pipeline_overrides
178            .override_interceptor::<T, I>(interceptor);
179        self
180    }
181
182    pub fn override_filter<T, F>(mut self, filter: F) -> Self
183    where
184        T: ExceptionFilter,
185        F: ExceptionFilter,
186    {
187        self.pipeline_overrides.override_filter::<T, F>(filter);
188        self
189    }
190
191    pub fn override_pipe<T, P>(mut self, pipe: P) -> Self
192    where
193        T: Pipe,
194        P: Pipe,
195    {
196        self.pipeline_overrides.override_pipe::<T, P>(pipe);
197        self
198    }
199
200    pub fn compile(self) -> Result<TestingModule> {
201        let mut app = self.app.import(self.module).build()?;
202        if !self.pipeline_overrides.is_empty() {
203            app.routes = app
204                .routes
205                .into_iter()
206                .map(|route| route.with_pipeline_overrides(&self.pipeline_overrides))
207                .collect();
208        }
209        Ok(TestingModule { app })
210    }
211
212    pub async fn compile_async(self) -> Result<TestingModule> {
213        let mut app = self.app.import(self.module).build_async().await?;
214        if !self.pipeline_overrides.is_empty() {
215            app.routes = app
216                .routes
217                .into_iter()
218                .map(|route| route.with_pipeline_overrides(&self.pipeline_overrides))
219                .collect();
220        }
221        Ok(TestingModule { app })
222    }
223}