1use std::future::Future;
31
32use axum::extract::{FromRequest, FromRequestParts, Query, RawPathParams, Request};
33use axum::handler::Handler;
34use axum::middleware::Next;
35use axum::response::{IntoResponse, Response};
36use axum::routing::{MethodFilter, MethodRouter};
37use http::header::{HeaderName, HeaderValue};
38use http::HeaderMap;
39use net_backend_protocol::http_call::placeholders;
40use net_backend_protocol::routes::HttpMethod;
41use net_backend_protocol::{HttpCall, PathParams, PayloadKind};
42use serde_json::Value;
43use utoipa::openapi::path::{HttpMethod as DocMethod, Paths};
44use utoipa_axum::router::{OpenApiRouter, UtoipaMethodRouter};
45
46use super::ApiJson;
47use crate::auth::{AuthContext, AuthFailure};
48use crate::error::AppError;
49use crate::state::AppState;
50
51#[derive(Debug)]
55pub struct Call<C>(pub C);
56
57impl<S, C> FromRequest<S> for Call<C>
58where
59 S: Send + Sync,
60 C: HttpCall + Send,
61 C::Payload: Send,
62{
63 type Rejection = AppError;
64
65 async fn from_request(request: Request, state: &S) -> Result<Self, Self::Rejection> {
66 let (mut parts, body) = request.into_parts();
67 let mut params = PathParams::new();
68 if !placeholders(C::ROUTE.path).is_empty() {
69 let raw =
70 RawPathParams::from_request_parts(&mut parts, state).await.map_err(|_| AppError::bad_request("the path parameters are not valid UTF-8"))?;
71 for (name, value) in &raw {
72 params.insert(name, value);
73 }
74 }
75 let payload: C::Payload = match C::PAYLOAD {
76 PayloadKind::Json => ApiJson::<C::Payload>::from_request(Request::from_parts(parts, body), state).await?.0,
77 PayloadKind::Query => {
78 Query::<C::Payload>::from_request_parts(&mut parts, state).await.map_err(|rejection| AppError::bad_request(rejection.body_text()))?.0
79 }
80 _ => serde_json::from_value(Value::Null).map_err(AppError::internal)?,
82 };
83 C::from_parts(¶ms, payload).map(Call).map_err(AppError::from)
84 }
85}
86
87#[derive(Debug)]
90pub struct Reply<C: HttpCall> {
91 data: C::Response,
92 headers: HeaderMap,
93}
94
95impl<C: HttpCall> Reply<C> {
96 pub fn new(data: C::Response) -> Self {
98 Self { data, headers: HeaderMap::new() }
99 }
100
101 pub fn with_header(mut self, name: HeaderName, value: HeaderValue) -> Self {
103 self.headers.insert(name, value);
104 self
105 }
106
107 pub fn data(&self) -> &C::Response {
109 &self.data
110 }
111}
112
113impl<C: HttpCall> IntoResponse for Reply<C> {
114 fn into_response(self) -> Response {
115 (self.headers, ApiJson(self.data)).into_response()
116 }
117}
118
119pub type CallResult<C> = Result<Reply<C>, AppError>;
121
122pub trait CallHandler<C, T>: Clone + Send + Sync + Sized + 'static {}
126
127macro_rules! call_handler {
128 ($($ty:ident),*) => {
129 impl<F, Fut, C, $($ty,)*> CallHandler<C, ($($ty,)*)> for F
130 where
131 F: Fn($($ty,)* Call<C>) -> Fut + Clone + Send + Sync + 'static,
132 Fut: Future<Output = CallResult<C>> + Send,
133 C: HttpCall,
134 {
135 }
136 };
137}
138
139call_handler!();
140call_handler!(T1);
141call_handler!(T1, T2);
142call_handler!(T1, T2, T3);
143call_handler!(T1, T2, T3, T4);
144call_handler!(T1, T2, T3, T4, T5);
145call_handler!(T1, T2, T3, T4, T5, T6);
146call_handler!(T1, T2, T3, T4, T5, T6, T7);
147call_handler!(T1, T2, T3, T4, T5, T6, T7, T8);
148
149fn method_filter(method: HttpMethod) -> Option<MethodFilter> {
152 match method {
153 HttpMethod::Get => Some(MethodFilter::GET),
154 HttpMethod::Post => Some(MethodFilter::POST),
155 HttpMethod::Put => Some(MethodFilter::PUT),
156 HttpMethod::Patch => Some(MethodFilter::PATCH),
157 HttpMethod::Delete => Some(MethodFilter::DELETE),
158 _ => None,
159 }
160}
161
162pub(crate) fn supported_method(method: HttpMethod) -> bool {
164 method_filter(method).is_some()
165}
166
167async fn require_auth(request: Request, next: Next) -> Response {
169 if request.extensions().get::<AuthContext>().is_some() {
170 return next.run(request).await;
171 }
172 request.extensions().get::<AuthFailure>().map_or_else(AppError::unauthorized, AuthFailure::to_error).into_response()
173}
174
175fn doc_method(method: HttpMethod) -> DocMethod {
176 match method {
177 HttpMethod::Post => DocMethod::Post,
178 HttpMethod::Put => DocMethod::Put,
179 HttpMethod::Patch => DocMethod::Patch,
180 HttpMethod::Delete => DocMethod::Delete,
181 _ => DocMethod::Get,
182 }
183}
184
185pub fn method_router<C, H, T, M>(handler: H) -> MethodRouter<AppState>
190where
191 C: HttpCall,
192 H: CallHandler<C, T> + Handler<M, AppState>,
193 M: 'static,
194{
195 let Some(filter) = method_filter(C::ROUTE.method) else {
196 tracing::error!(method = %C::ROUTE.method, path = C::ROUTE.path, "this server version cannot serve this method: the route answers 405");
197 return MethodRouter::new();
198 };
199 let router = axum::routing::on(filter, handler);
200 if C::ROUTE.auth {
201 router.route_layer(axum::middleware::from_fn(require_auth))
202 } else {
203 router
204 }
205}
206
207pub fn documented<C, H, T, M>(doc: UtoipaMethodRouter<AppState>, handler: H) -> UtoipaMethodRouter<AppState>
211where
212 C: HttpCall,
213 H: CallHandler<C, T> + Handler<M, AppState>,
214 M: 'static,
215{
216 let (schemas, doc_paths, _) = doc;
217 let operation = doc_paths.paths.into_values().find_map(|item| item.get.or(item.put).or(item.post).or(item.delete).or(item.patch));
218 let mut paths = Paths::new();
219 if let Some(operation) = operation {
220 paths.add_path_operation(C::ROUTE.path, vec![doc_method(C::ROUTE.method)], operation);
221 }
222 (schemas, paths, method_router::<C, H, T, M>(handler))
223}
224
225pub fn undocumented<C, H, T, M>(router: OpenApiRouter<AppState>, handler: H) -> OpenApiRouter<AppState>
227where
228 C: HttpCall,
229 H: CallHandler<C, T> + Handler<M, AppState>,
230 M: 'static,
231{
232 router.route(C::ROUTE.path, method_router::<C, H, T, M>(handler))
233}
234
235#[macro_export]
240macro_rules! call_route {
241 ($call:ty, $handler:path) => {
242 $crate::http::call::documented::<$call, _, _, _>($crate::utoipa_axum::routes!($handler), $handler)
243 };
244}
245
246#[cfg(test)]
247mod tests {
248 use super::*;
249
250 #[test]
251 fn methods() {
252 for (method, filter, doc) in [
253 (HttpMethod::Get, MethodFilter::GET, DocMethod::Get),
254 (HttpMethod::Post, MethodFilter::POST, DocMethod::Post),
255 (HttpMethod::Put, MethodFilter::PUT, DocMethod::Put),
256 (HttpMethod::Patch, MethodFilter::PATCH, DocMethod::Patch),
257 (HttpMethod::Delete, MethodFilter::DELETE, DocMethod::Delete),
258 ] {
259 assert_eq!(method_filter(method), Some(filter));
260 assert!(supported_method(method));
261 assert!(doc_method(method) == doc, "{method}");
262 }
263 }
264}