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
10pub 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
98pub 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}