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