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