Skip to main content

tuwunel_router/
layers.rs

1#[cfg(test)]
2mod tests;
3
4use std::{
5	any::Any,
6	convert::Infallible,
7	mem::replace,
8	sync::Arc,
9	task::{Context, Poll},
10	time::Duration,
11};
12
13use axum::{
14	Extension, Router,
15	extract::{DefaultBodyLimit, MatchedPath, Request},
16	response::{IntoResponse, Response},
17};
18use futures::{FutureExt, future::Map};
19use http::{
20	HeaderValue, Method, StatusCode,
21	header::{
22		self, CONTENT_SECURITY_POLICY, CONTENT_TYPE, ETAG, HeaderName, IF_MATCH, IF_NONE_MATCH,
23		X_FRAME_OPTIONS,
24	},
25	uri::PathAndQuery,
26};
27use ipnet::IpNet;
28use tower::{
29	Layer, Service, ServiceBuilder,
30	layer::util::Identity,
31	util::{Either, MapResponseLayer, option_layer},
32};
33use tower_http::{
34	catch_panic::CatchPanicLayer,
35	cors::{AllowOrigin, CorsLayer},
36	sensitive_headers::SetSensitiveHeadersLayer,
37	set_header::SetResponseHeaderLayer,
38	timeout::{RequestBodyTimeoutLayer, ResponseBodyTimeoutLayer, TimeoutLayer},
39	trace::{DefaultOnFailure, DefaultOnRequest, DefaultOnResponse, TraceLayer},
40};
41use tracing::Level;
42use tuwunel_api::router::{ConfiguredIpSource, TrustedPeerSubnets, state::Guard};
43use tuwunel_core::{
44	Result, Server, config::IpSource, debug, error, utils::content_disposition::content_type_is,
45};
46use tuwunel_service::Services;
47
48use crate::{request, router};
49
50type Convert = fn(Result<Response, StatusCode>) -> Result<Response, Infallible>;
51
52/// Bespoke `axum::middleware::from_fn`: threading the handler's future type
53/// through `F` spares the boxes the generic middleware allocates per request.
54#[derive(Clone)]
55pub(crate) struct HandleLayer<F> {
56	pub(crate) services: Arc<Services>,
57	pub(crate) handler: F,
58}
59
60#[derive(Clone)]
61pub(crate) struct Handle<S, F> {
62	services: Arc<Services>,
63	handler: F,
64	inner: S,
65}
66
67const TUWUNEL_CSP: &[&str] = &[
68	"default-src 'none'",
69	"script-src 'self'",
70	"style-src 'self'",
71	"frame-ancestors 'none'",
72	"form-action 'self'",
73	"base-uri 'none'",
74];
75
76pub(crate) fn build(services: &Arc<Services>) -> Result<(Router, Guard)> {
77	let server = &services.server;
78	let layers = ServiceBuilder::new();
79
80	#[cfg(feature = "sentry_telemetry")]
81	let layers = layers.layer(sentry_tower::NewSentryLayer::<http::Request<_>>::new_from_top());
82
83	#[cfg(any(
84		feature = "zstd_compression",
85		feature = "gzip_compression",
86		feature = "brotli_compression"
87	))]
88	let layers = layers.layer(compression_layer(server));
89
90	let services_ = services.clone();
91	let layers = layers
92		.layer(SetSensitiveHeadersLayer::new([header::AUTHORIZATION]))
93		.layer(
94			TraceLayer::new_for_http()
95				.make_span_with(tracing_span::<_>)
96				.on_failure(DefaultOnFailure::new().level(Level::ERROR))
97				.on_request(DefaultOnRequest::new().level(Level::TRACE))
98				.on_response(DefaultOnResponse::new().level(Level::DEBUG)),
99		)
100		.layer(HandleLayer {
101			services: Arc::clone(services),
102			handler: request::handle,
103		})
104		.layer(trusted_peer_subnets_layer(&server.config.ip_source_trusted_subnets))
105		.layer(ip_source_layer(server.config.ip_source))
106		.layer(ResponseBodyTimeoutLayer::new(Duration::from_secs(
107			server.config.client_response_timeout,
108		)))
109		.layer(RequestBodyTimeoutLayer::new(Duration::from_secs(
110			server.config.client_receive_timeout,
111		)))
112		.layer(TimeoutLayer::with_status_code(
113			StatusCode::REQUEST_TIMEOUT,
114			Duration::from_secs(server.config.client_request_timeout),
115		))
116		.layer(SetResponseHeaderLayer::if_not_present(
117			header::X_CONTENT_TYPE_OPTIONS,
118			HeaderValue::from_static("nosniff"),
119		))
120		.layer(html_layer())
121		.layer(cors_layer(server))
122		.layer(body_limit_layer(server))
123		.layer(CatchPanicLayer::custom(move |panic| catch_panic(panic, services_.clone())));
124
125	let (router, guard) = router::build(services);
126	Ok((router.layer(layers), guard))
127}
128
129impl<S, F: Clone> Layer<S> for HandleLayer<F> {
130	type Service = Handle<S, F>;
131
132	fn layer(&self, inner: S) -> Self::Service {
133		Handle {
134			services: self.services.clone(),
135			handler: self.handler.clone(),
136			inner,
137		}
138	}
139}
140
141impl<S, F, Fut> Service<Request> for Handle<S, F>
142where
143	S: Service<Request, Error = Infallible> + Clone,
144	F: FnMut(Arc<Services>, Request, S) -> Fut,
145	Fut: Future<Output = Result<Response, StatusCode>>,
146{
147	type Error = Infallible;
148	type Future = Map<Fut, Convert>;
149	type Response = Response;
150
151	fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
152		self.inner.poll_ready(cx)
153	}
154
155	fn call(&mut self, req: Request) -> Self::Future {
156		let convert: Convert = |result| Ok(result.into_response());
157		let unready = self.inner.clone();
158		let inner = replace(&mut self.inner, unready);
159
160		(self.handler)(self.services.clone(), req, inner).map(convert)
161	}
162}
163
164#[cfg(any(
165	feature = "zstd_compression",
166	feature = "gzip_compression",
167	feature = "brotli_compression"
168))]
169fn compression_layer(server: &Server) -> tower_http::compression::CompressionLayer {
170	let mut compression_layer = tower_http::compression::CompressionLayer::new();
171
172	#[cfg(feature = "zstd_compression")]
173	{
174		compression_layer = if server.config.zstd_compression {
175			compression_layer.zstd(true)
176		} else {
177			compression_layer.no_zstd()
178		};
179	};
180
181	#[cfg(feature = "gzip_compression")]
182	{
183		compression_layer = if server.config.gzip_compression {
184			compression_layer.gzip(true)
185		} else {
186			compression_layer.no_gzip()
187		};
188	};
189
190	#[cfg(feature = "brotli_compression")]
191	{
192		compression_layer = if server.config.brotli_compression {
193			compression_layer.br(true)
194		} else {
195			compression_layer.no_br()
196		};
197	};
198
199	compression_layer
200}
201
202fn cors_layer(server: &Server) -> CorsLayer {
203	const METHODS: [Method; 7] = [
204		Method::DELETE,
205		Method::GET,
206		Method::HEAD,
207		Method::OPTIONS,
208		Method::PATCH,
209		Method::POST,
210		Method::PUT,
211	];
212
213	let headers: [HeaderName; 7] = [
214		header::ACCEPT,
215		header::AUTHORIZATION,
216		CONTENT_TYPE,
217		IF_MATCH,
218		IF_NONE_MATCH,
219		header::ORIGIN,
220		HeaderName::from_static("x-requested-with"),
221	];
222
223	let allow_origin_list = server
224		.config
225		.access_control_allow_origin
226		.iter()
227		.map(AsRef::as_ref)
228		.map(HeaderValue::from_str)
229		.filter_map(Result::ok);
230
231	let allow_origin = if !server
232		.config
233		.access_control_allow_origin
234		.is_empty()
235	{
236		AllowOrigin::list(allow_origin_list)
237	} else {
238		AllowOrigin::any()
239	};
240
241	CorsLayer::new()
242		.max_age(Duration::from_hours(24))
243		.allow_methods(METHODS)
244		.allow_headers(headers)
245		.expose_headers([ETAG])
246		.allow_origin(allow_origin)
247}
248
249fn body_limit_layer(server: &Server) -> DefaultBodyLimit {
250	DefaultBodyLimit::max(server.config.max_request_size)
251}
252
253fn trusted_peer_subnets_layer(
254	subnets: &[IpNet],
255) -> Either<Extension<TrustedPeerSubnets>, Identity> {
256	option_layer((!subnets.is_empty()).then(|| Extension(TrustedPeerSubnets(Arc::from(subnets)))))
257}
258
259fn ip_source_layer(source: Option<IpSource>) -> Either<Extension<ConfiguredIpSource>, Identity> {
260	option_layer(source.map(|source| Extension(ConfiguredIpSource(source))))
261}
262
263fn html_layer<T>() -> MapResponseLayer<impl Fn(http::Response<T>) -> http::Response<T> + Clone> {
264	MapResponseLayer::new(set_html_headers)
265}
266
267/// Denies framing and foreign scripts for a response that is HTML.
268///
269/// The Content-Type is echoed from whoever uploaded the media, so the media
270/// type decides on its own and regardless of case: a parameter that merely
271/// mentions HTML leaves a response that is not HTML alone.
272fn set_html_headers<T>(mut response: http::Response<T>) -> http::Response<T> {
273	let headers = response.headers_mut();
274
275	let content_type = headers
276		.get(CONTENT_TYPE)
277		.map(HeaderValue::to_str)
278		.and_then(Result::ok);
279
280	if content_type_is(content_type, "text/html") {
281		headers
282			.entry(CONTENT_SECURITY_POLICY)
283			.or_insert(HeaderValue::from_static(const_str::join!(TUWUNEL_CSP, ";")));
284
285		headers
286			.entry(X_FRAME_OPTIONS)
287			.or_insert(HeaderValue::from_static("DENY"));
288	}
289
290	response
291}
292
293#[tracing::instrument(name = "panic", level = "error", skip_all)]
294#[expect(clippy::needless_pass_by_value)]
295fn catch_panic(
296	err: Box<dyn Any + Send + 'static>,
297	services: Arc<Services>,
298) -> http::Response<http_body_util::Full<bytes::Bytes>> {
299	services
300		.server
301		.metrics
302		.requests_panic
303		.fetch_add(1, std::sync::atomic::Ordering::Release);
304
305	let details = match err.downcast_ref::<String>() {
306		| Some(s) => s.clone(),
307		| _ => match err.downcast_ref::<&str>() {
308			| Some(s) => (*s).to_owned(),
309			| _ => "Unknown internal server error occurred.".to_owned(),
310		},
311	};
312
313	error!("{details:#}");
314	let body = serde_json::json!({
315		"errcode": "M_UNKNOWN",
316		"error": "M_UNKNOWN: Internal server error occurred",
317		"details": details,
318	});
319
320	http::Response::builder()
321		.status(StatusCode::INTERNAL_SERVER_ERROR)
322		.header(CONTENT_TYPE, "application/json")
323		.body(http_body_util::Full::from(body.to_string()))
324		.expect("Failed to create response for our panic catcher?")
325}
326
327fn tracing_span<T>(request: &http::Request<T>) -> tracing::Span {
328	let path = request
329		.extensions()
330		.get::<MatchedPath>()
331		.map_or_else(|| request_path_str(request), truncated_matched_path);
332
333	tracing::span! {
334		parent: None,
335		debug::INFO_SPAN_LEVEL,
336		"router",
337		method = %request.method(),
338		%path,
339	}
340}
341
342fn request_path_str<T>(request: &http::Request<T>) -> &str {
343	request
344		.uri()
345		.path_and_query()
346		.map(PathAndQuery::as_str)
347		.unwrap_or("/")
348}
349
350fn truncated_matched_path(path: &MatchedPath) -> &str {
351	path.as_str()
352		.rsplit_once('{')
353		.map_or(path.as_str(), |path| path.0.strip_suffix('/').unwrap_or(path.0))
354}