Skip to main content

tuwunel_api/router/
handler.rs

1use std::{any::Any, fmt::Debug};
2
3use axum::{
4	Router,
5	body::Body,
6	extract::{FromRequest, FromRequestParts},
7	handler::Handler,
8	response::{IntoResponse, Response},
9	routing::{MethodFilter, on},
10};
11use http::{Method, Request};
12use ruma::api::{IncomingRequest, path_builder::PathBuilder};
13use tuwunel_core::Result;
14
15use super::{Ruma, RumaResponse, State, auth::AuthDispatch};
16
17pub(in super::super) trait RumaHandler<T> {
18	fn add_routes(&'static self, router: Router<State>) -> Router<State>;
19
20	fn add_route(&'static self, router: Router<State>, path: &str) -> Router<State>;
21
22	fn call_route(
23		handler: RouteHandler,
24		state: State,
25		request: Request<Body>,
26	) -> impl Future<Output = Response> + Send + 'static;
27}
28
29pub(in super::super) trait RouterExt {
30	fn ruma_route<H: RumaHandler<T>, T>(self, handler: &'static H) -> Self;
31}
32
33struct Route<Fut> {
34	call: RouteCall<Fut>,
35	handler: RouteHandler,
36}
37
38type RouteCall<Fut> = fn(RouteHandler, State, Request<Body>) -> Fut;
39type RouteHandler = &'static (dyn Any + Send + Sync);
40
41impl<Fut> Copy for Route<Fut> {}
42
43impl<Fut> Clone for Route<Fut> {
44	fn clone(&self) -> Self { *self }
45}
46
47impl RouterExt for Router<State> {
48	fn ruma_route<H: RumaHandler<T>, T>(self, handler: &'static H) -> Self {
49		handler.add_routes(self)
50	}
51}
52
53impl<Fut> Handler<(), State> for Route<Fut>
54where
55	Fut: Future<Output = Response> + Send + 'static,
56{
57	type Future = Fut;
58
59	fn call(self, request: Request<Body>, state: State) -> Self::Future {
60		(self.call)(self.handler, state, request)
61	}
62}
63
64macro_rules! ruma_handler {
65	( $($tx:ident),* $(,)? ) => {
66		#[allow(clippy::allow_attributes, non_snake_case)]
67		impl<Err, Req, Fut, Fun, const ADMIN: bool, $($tx,)*> RumaHandler<($($tx,)* Ruma<Req, ADMIN>,)> for Fun
68		where
69			Fun: Fn($($tx,)* Ruma<Req, ADMIN>,) -> Fut + Send + Sync + 'static,
70			Fut: Future<Output = Result<Req::OutgoingResponse, Err>> + Send + 'static,
71			Req: IncomingRequest + Debug + Send + Sync + 'static,
72			Req::Authentication: AuthDispatch,
73			Err: IntoResponse + Send,
74			<Req as IncomingRequest>::OutgoingResponse: Send,
75			$( $tx: FromRequestParts<State> + Send + Sync + 'static, )*
76		{
77			fn add_routes(&'static self, router: Router<State>) -> Router<State> {
78				Req::PATH_BUILDER
79					.all_paths()
80					.fold(router, |router, path| self.add_route(router, path))
81			}
82
83			fn add_route(&'static self, router: Router<State>, path: &str) -> Router<State> {
84				let route = Route { handler: self, call: Self::call_route };
85
86				router.route(path, on(method_to_filter(&Req::METHOD), route))
87			}
88
89			fn call_route(
90				handler: RouteHandler,
91				state: State,
92				request: Request<Body>,
93			) -> impl Future<Output = Response> + Send + 'static {
94				let handler: &'static Fun = handler
95					.downcast_ref()
96					.expect("route handler matches the type it registered with");
97
98				let response = async move {
99					#[allow(unused_mut)]
100					let (mut parts, body) = request.into_parts();
101					$(
102						let $tx = match $tx::from_request_parts(&mut parts, &state).await {
103							| Err(error) => return error.into_response(),
104							| Ok(value) => value,
105						};
106					)*
107
108					let request = Request::from_parts(parts, body);
109					let args = match Ruma::<Req, ADMIN>::from_request(request, &state).await {
110						| Err(error) => return error.into_response(),
111						| Ok(args) => args,
112					};
113
114					match handler($($tx,)* args).await {
115						| Err(error) => error.into_response(),
116						| Ok(response) => RumaResponse(response).into_response(),
117					}
118				};
119
120				cfg_select! {
121					debug_assertions => Box::pin(response),
122					_ => response,
123				}
124			}
125		}
126	}
127}
128ruma_handler!();
129ruma_handler!(T1);
130ruma_handler!(T1, T2);
131ruma_handler!(T1, T2, T3);
132ruma_handler!(T1, T2, T3, T4);
133
134fn method_to_filter(method: &Method) -> MethodFilter {
135	match method {
136		| &Method::DELETE => MethodFilter::DELETE,
137		| &Method::GET => MethodFilter::GET,
138		| &Method::HEAD => MethodFilter::HEAD,
139		| &Method::OPTIONS => MethodFilter::OPTIONS,
140		| &Method::PATCH => MethodFilter::PATCH,
141		| &Method::POST => MethodFilter::POST,
142		| &Method::PUT => MethodFilter::PUT,
143		| &Method::TRACE => MethodFilter::TRACE,
144		| _ => panic!("Unsupported HTTP method"),
145	}
146}