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}