Skip to main content

tuwunel_router/
request.rs

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	// we made a safety contract that Services will not go out of scope
140	// during the request; this ensures a reference is accounted for at
141	// the base frame of the task regardless of its detachment.
142	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}