Skip to main content

a3s_boot/routing/route/
execution.rs

1use super::definition::RouteDefinition;
2use crate::routing::path::{join_paths, route_shape_key, route_specificity};
3use crate::{
4    BootError, BootRequest, BootResponse, CallHandler, ContextId, ContextIdFactory,
5    ExecutionContext, MiddlewareOutcome, Result,
6};
7use std::sync::{Arc, Mutex};
8
9#[derive(Clone)]
10struct PipelineErrorContext {
11    context: Arc<Mutex<ExecutionContext>>,
12}
13
14impl PipelineErrorContext {
15    fn new(context: ExecutionContext) -> Self {
16        Self {
17            context: Arc::new(Mutex::new(context)),
18        }
19    }
20
21    fn replace(&self, context: ExecutionContext) {
22        *self
23            .context
24            .lock()
25            .unwrap_or_else(std::sync::PoisonError::into_inner) = context;
26    }
27
28    fn snapshot(&self) -> ExecutionContext {
29        self.context
30            .lock()
31            .unwrap_or_else(std::sync::PoisonError::into_inner)
32            .clone()
33    }
34}
35
36impl RouteDefinition {
37    pub async fn call(&self, mut request: BootRequest) -> Result<BootResponse> {
38        let context_id = ContextIdFactory::create();
39        request = self.attach_request_context(request, &context_id);
40        if !self.method.matches(request.method) {
41            let message = format!("{} {}", request.method.as_str(), request.path);
42            return self
43                .handle_error(
44                    self.execution_context(request),
45                    BootError::MethodNotAllowed(message),
46                    &context_id,
47                )
48                .await;
49        }
50
51        let params = match self.path_params(&request.path) {
52            Ok(Some(params)) => params,
53            Ok(None) => {
54                let message = format!("{} {}", request.method.as_str(), request.path);
55                return self
56                    .handle_error(
57                        self.execution_context(request),
58                        BootError::NotFound(message),
59                        &context_id,
60                    )
61                    .await;
62            }
63            Err(error) => {
64                return self
65                    .handle_error(self.execution_context(request), error, &context_id)
66                    .await;
67            }
68        };
69        request = request.with_path_params(params);
70        let host_params = match self.host_params(request.host()) {
71            Ok(Some(params)) => params,
72            Ok(None) => {
73                let message = format!("{} {}", request.method.as_str(), request.path);
74                return self
75                    .handle_error(
76                        self.execution_context(request),
77                        BootError::NotFound(message),
78                        &context_id,
79                    )
80                    .await;
81            }
82            Err(error) => {
83                return self
84                    .handle_error(self.execution_context(request), error, &context_id)
85                    .await;
86            }
87        };
88        request = request.with_host_params(host_params);
89        #[cfg(feature = "request-context")]
90        {
91            let context = crate::RequestContext::from_route_request(
92                &request,
93                self.path.clone(),
94                self.module_name.clone(),
95                self.controller_prefix.clone(),
96                self.metadata.clone(),
97            );
98            return crate::RequestContext::scope(context, self.call_pipeline(request, context_id))
99                .await;
100        }
101
102        #[cfg(not(feature = "request-context"))]
103        {
104            self.call_pipeline(request, context_id).await
105        }
106    }
107
108    async fn call_pipeline(
109        &self,
110        mut request: BootRequest,
111        context_id: ContextId,
112    ) -> Result<BootResponse> {
113        for middleware in &self.middleware {
114            let context_request = request.clone();
115            request = match middleware.handle(request).await {
116                Ok(MiddlewareOutcome::Continue(request)) => {
117                    self.attach_request_context(request, &context_id)
118                }
119                Ok(MiddlewareOutcome::Respond(response)) => return Ok(response),
120                Err(error) => {
121                    return self
122                        .handle_error(self.execution_context(context_request), error, &context_id)
123                        .await;
124                }
125            };
126        }
127
128        let context = self.execution_context(request.clone());
129
130        for guard in &self.guards {
131            let guard = match guard.resolve(&context_id) {
132                Ok(guard) => guard,
133                Err(error) => {
134                    return self.handle_error(context.clone(), error, &context_id).await;
135                }
136            };
137            let can_activate = match guard.can_activate(context.clone()).await {
138                Ok(can_activate) => can_activate,
139                Err(error) => {
140                    return self.handle_error(context.clone(), error, &context_id).await;
141                }
142            };
143
144            if !can_activate {
145                let message = format!("{} {}", context.method.as_str(), context.request_path);
146                return self
147                    .handle_error(context, BootError::Forbidden(message), &context_id)
148                    .await;
149            }
150        }
151
152        let mut resolved_interceptors = Vec::with_capacity(self.interceptors.len());
153        for interceptor in &self.interceptors {
154            match interceptor.resolve(&context_id) {
155                Ok(interceptor) => resolved_interceptors.push(interceptor),
156                Err(error) => {
157                    return self.handle_error(context.clone(), error, &context_id).await;
158                }
159            }
160        }
161
162        let error_context = PipelineErrorContext::new(context.clone());
163        let terminal_context = context.clone();
164        let terminal_error_context = error_context.clone();
165        let handler_context_id = context_id.clone();
166        let mut next = CallHandler::from_fn(move || {
167            terminal_error_context.replace(terminal_context.clone());
168            self.call_handler_pipeline(
169                request.clone(),
170                terminal_error_context.clone(),
171                handler_context_id.clone(),
172            )
173        });
174        for interceptor in resolved_interceptors.iter().rev() {
175            let interceptor_context = context.clone();
176            let success_context = context.clone();
177            let interceptor_error_context = error_context.clone();
178            let downstream = next.clone();
179            next = CallHandler::from_fn(move || {
180                interceptor_error_context.replace(interceptor_context.clone());
181                let future = interceptor.intercept(interceptor_context.clone(), downstream.clone());
182                let success_context = success_context.clone();
183                let interceptor_error_context = interceptor_error_context.clone();
184                async move {
185                    let result = future.await;
186                    if result.is_ok() {
187                        interceptor_error_context.replace(success_context);
188                    }
189                    result
190                }
191            });
192        }
193
194        match next.handle().await {
195            Ok(response) => Ok(response),
196            Err(error) => {
197                self.handle_error(error_context.snapshot(), error, &context_id)
198                    .await
199            }
200        }
201    }
202
203    async fn call_handler_pipeline(
204        &self,
205        mut request: BootRequest,
206        error_context: PipelineErrorContext,
207        context_id: ContextId,
208    ) -> Result<BootResponse> {
209        for pipe in &self.pipes {
210            let context_request = request.clone();
211            let pipe = pipe.resolve(&context_id)?;
212            request = match pipe.transform(request).await {
213                Ok(request) => self.attach_request_context(request, &context_id),
214                Err(error) => {
215                    error_context.replace(self.execution_context(context_request));
216                    return Err(error);
217                }
218            };
219        }
220
221        if self.validation_enabled {
222            for validator in &self.validators {
223                let context_request = request.clone();
224                request = match validator(request, self.validation_options) {
225                    Ok(request) => request,
226                    Err(error) => {
227                        error_context.replace(self.execution_context(context_request));
228                        return Err(error);
229                    }
230                };
231            }
232        }
233
234        self.handler.call(request).await
235    }
236
237    /// Dispatch a request through this route and convert unhandled errors into Boot HTTP responses.
238    pub async fn handle(&self, request: BootRequest) -> BootResponse {
239        match self.call(request).await {
240            Ok(response) => response,
241            Err(error) => BootResponse::from_error(&error),
242        }
243    }
244
245    pub(crate) fn matches_path_shape(&self, path: &str) -> bool {
246        self.matches_path(path)
247    }
248
249    pub(crate) fn path_shape_key(&self) -> String {
250        route_shape_key(&self.path)
251    }
252
253    pub(crate) fn path_specificity(&self) -> Vec<u8> {
254        route_specificity(&self.path)
255    }
256
257    fn execution_context(&self, request: BootRequest) -> ExecutionContext {
258        ExecutionContext::new(
259            request,
260            self.path.clone(),
261            self.module_name.clone(),
262            self.controller_prefix.clone(),
263            self.serialization.clone(),
264            self.metadata.clone(),
265        )
266    }
267
268    pub(crate) fn with_prefix(mut self, prefix: &str) -> Result<Self> {
269        self.path = join_paths(prefix, &self.path)?;
270        self.controller_prefix = Some(prefix.trim_end_matches('/').to_string());
271        Ok(self)
272    }
273
274    pub(crate) fn with_path_prefix(mut self, prefix: &str) -> Result<Self> {
275        self.path = join_paths(prefix, &self.path)?;
276        Ok(self)
277    }
278
279    pub(crate) fn with_module_name(mut self, module_name: &str) -> Self {
280        self.module_name = Some(module_name.to_string());
281        self
282    }
283
284    pub(crate) fn with_module_ref(mut self, module_ref: crate::ModuleRef) -> Self {
285        self.module_ref = Some(module_ref);
286        self
287    }
288
289    pub(crate) fn with_default_module_ref(mut self, module_ref: crate::ModuleRef) -> Self {
290        if self.module_ref.is_none() {
291            self.module_ref = Some(module_ref);
292        }
293        self
294    }
295
296    async fn handle_error(
297        &self,
298        context: ExecutionContext,
299        error: BootError,
300        context_id: &ContextId,
301    ) -> Result<BootResponse> {
302        for filter in self.filters.iter().rev() {
303            let filter = filter.resolve(context_id)?;
304            if let Some(response) = filter
305                .catch(context.clone(), error.clone_for_filter())
306                .await?
307            {
308                return Ok(response);
309            }
310        }
311        Err(error)
312    }
313
314    fn attach_request_context(&self, request: BootRequest, context_id: &ContextId) -> BootRequest {
315        match &self.module_ref {
316            Some(module_ref) => request.with_module_ref(module_ref.context_scope(context_id)),
317            None => request,
318        }
319    }
320}