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#[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
266fn 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}