1use std::{
2 convert::Infallible,
3 fmt::Debug,
4 sync::{Arc, atomic::Ordering},
5 time::Duration,
6};
7
8use axum::{
9 extract::{MatchedPath, Request},
10 response::{IntoResponse, Response},
11};
12use futures::FutureExt;
13use http::{Method, StatusCode, Uri};
14use ruma::api::error::ErrorKind;
15use tokio::{sync::Notify, task, time::sleep};
16use tower::{Service, ServiceExt};
17use tracing::{Span, field::Empty};
18use tuwunel_core::{
19 Error, Result, debug, debug_error, debug_warn, defer, error, trace, utils::SanitizedUri,
20};
21use tuwunel_service::Services;
22
23#[tracing::instrument(
24 name = "request",
25 level = "debug",
26 skip_all,
27 err(Debug, level = "debug")
28 fields(
29 task = %task::id(),
30 id = %services
31 .server
32 .metrics
33 .requests_count
34 .fetch_add(1, Ordering::Relaxed),
35 origin = Empty,
36 user_id = Empty,
37 device_id = Empty,
38 )
39)]
40pub(crate) async fn handle<S>(
41 services: Arc<Services>,
42 mut req: Request,
43 inner: S,
44) -> Result<Response, StatusCode>
45where
46 S: Service<Request, Error = Infallible> + Send + 'static,
47 S::Response: IntoResponse,
48 S::Future: Send + 'static,
49{
50 let matched_path = req.extensions().get::<MatchedPath>().cloned();
51
52 if !services.server.is_running() {
53 let uri = matched_path.as_ref().map_or_else(
54 || SanitizedUri::new(req.uri()),
55 |path| SanitizedUri::with_path(req.uri(), path.as_str()),
56 );
57
58 debug_warn!(
59 method = %req.method(),
60 %uri,
61 "unavailable pending shutdown"
62 );
63
64 return Err(StatusCode::SERVICE_UNAVAILABLE);
65 }
66
67 let uri = req.uri().clone();
68 let method = req.method().clone();
69 let parent = Span::current();
70 req.extensions_mut().insert(parent.clone());
71
72 let response = match method {
73 | Method::PUT | Method::POST | Method::DELETE | Method::PATCH =>
74 spawn_execute(services, req, inner, parent).await?,
75 | _ => execute(&services, req, inner, &parent).await,
76 };
77
78 handle_result(&method, &uri, matched_path.as_ref(), response)
79}
80
81async fn spawn_execute<S>(
82 services: Arc<Services>,
83 mut req: Request,
84 inner: S,
85 parent: Span,
86) -> Result<Response, StatusCode>
87where
88 S: Service<Request, Error = Infallible> + Send + 'static,
89 S::Response: IntoResponse,
90 S::Future: Send + 'static,
91{
92 let detached = Arc::new(Notify::new());
93 req.extensions_mut().insert(detached.clone());
94
95 let task = services
96 .clone()
97 .server
98 .runtime()
99 .spawn(async move {
100 tokio::select! {
101 response = execute(&services, req, inner, &parent) => response,
102 response = services.server.until_shutdown()
103 .then(|()| {
104 let timeout = services.config.client_shutdown_timeout;
105 sleep(Duration::from_secs(timeout))
106 })
107 .map(|()| StatusCode::SERVICE_UNAVAILABLE)
108 .map(IntoResponse::into_response) => response,
109 }
110 });
111
112 let abort = task.abort_handle();
113 defer! {{
114 if !abort.is_finished() {
115 debug_warn!(
116 task = ?abort.id(),
117 "Client disconnected; detached request."
118 );
119
120 detached.notify_one();
121 }
122 }};
123
124 task.await.map_err(unhandled)
125}
126
127#[tracing::instrument(
128 name = "handle",
129 level = "debug",
130 parent = parent,
131 skip_all,
132 ret(level = "trace"),
133 fields(
134 task = %task::id(),
135 )
136)]
137#[cfg_attr(not(debug_assertions), expect(unused_variables))]
138async fn execute<S>(
139 services: &Arc<Services>,
143 req: Request,
144 inner: S,
145 parent: &Span,
146) -> Response
147where
148 S: Service<Request, Error = Infallible>,
149 S::Response: IntoResponse,
150{
151 #[cfg(debug_assertions)]
152 services
153 .server
154 .metrics
155 .requests_handle_active
156 .fetch_add(1, Ordering::Relaxed);
157
158 #[cfg(debug_assertions)]
159 defer! {{
160 _ = services.server
161 .metrics
162 .requests_handle_finished
163 .fetch_add(1, Ordering::Relaxed);
164 _ = services.server
165 .metrics
166 .requests_handle_active
167 .fetch_sub(1, Ordering::Relaxed);
168 }};
169
170 inner
171 .oneshot(req)
172 .map(IntoResponse::into_response)
173 .await
174}
175
176fn handle_result(
177 method: &Method,
178 uri: &Uri,
179 matched_path: Option<&MatchedPath>,
180 result: Response,
181) -> Result<Response, StatusCode> {
182 let status = result.status();
183 let code = status.as_u16();
184 let reason = status
185 .canonical_reason()
186 .unwrap_or("Unknown Reason");
187
188 let uri = matched_path.map_or_else(
189 || SanitizedUri::new(uri),
190 |path| SanitizedUri::with_path(uri, path.as_str()),
191 );
192
193 match status {
194 | status if status.is_redirection() =>
195 debug!(method = ?method, %uri, status = code, %reason, "request complete"),
196 | status if status.is_server_error() =>
197 error!(method = ?method, %uri, status = code, %reason, "request complete"),
198 | status if status.is_client_error() => {
199 debug_error!(method = ?method, %uri, status = code, %reason, "request complete");
200 },
201 | _ => trace!(method = ?method, %uri, status = code, %reason, "request complete"),
202 }
203
204 if status == StatusCode::METHOD_NOT_ALLOWED {
205 return Ok(Error::Request(
206 ErrorKind::Unrecognized,
207 "Method Not Allowed".into(),
208 StatusCode::METHOD_NOT_ALLOWED,
209 )
210 .into_response());
211 }
212
213 Ok(result)
214}
215
216#[cold]
217fn unhandled<Error: Debug>(e: Error) -> StatusCode {
218 error!(error = ?e, "unhandled error or panic during request");
219
220 StatusCode::INTERNAL_SERVER_ERROR
221}