1use std::ops::{Deref, DerefMut};
4
5use axum::extract::FromRequestParts;
6use axum::extract::path::ErrorKind;
7use axum::extract::rejection::PathRejection;
8use axum::http::request::Parts;
9use serde::de::DeserializeOwned;
10
11use crate::Error;
12
13#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
24pub struct Path<T>(pub T);
25
26impl<T> Deref for Path<T> {
27 type Target = T;
28
29 fn deref(&self) -> &T {
30 &self.0
31 }
32}
33
34impl<T> DerefMut for Path<T> {
35 fn deref_mut(&mut self) -> &mut T {
36 &mut self.0
37 }
38}
39
40impl<T, S> FromRequestParts<S> for Path<T>
41where
42 T: DeserializeOwned + Send,
43 S: Send + Sync,
44{
45 type Rejection = Error;
46
47 async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Error> {
48 match axum::extract::Path::<T>::from_request_parts(parts, state).await {
49 Ok(axum::extract::Path(value)) => Ok(Path(value)),
50 Err(PathRejection::FailedToDeserializePathParams(err)) => match err.kind() {
53 ErrorKind::WrongNumberOfParameters { .. } | ErrorKind::UnsupportedType { .. } => {
54 Err(anyhow::anyhow!("{err}").into())
55 }
56 _ => Err(Error::NotFound),
57 },
58 Err(other) => Err(anyhow::anyhow!("{other}").into()),
60 }
61 }
62}
63
64#[derive(Debug, Clone, Default, PartialEq)]
85pub struct Found<M>(pub M);
86
87impl<M> Deref for Found<M> {
88 type Target = M;
89
90 fn deref(&self) -> &M {
91 &self.0
92 }
93}
94
95impl<M> DerefMut for Found<M> {
96 fn deref_mut(&mut self) -> &mut M {
97 &mut self.0
98 }
99}
100
101impl<M, S> FromRequestParts<S> for Found<M>
102where
103 M: crate::db::Model,
104 S: Send + Sync,
105{
106 type Rejection = Error;
107
108 async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Error> {
109 let params = axum::extract::RawPathParams::from_request_parts(parts, state)
110 .await
111 .map_err(|err| anyhow::anyhow!("{err}"))?;
112 let params: Vec<(String, String)> = params
113 .iter()
114 .map(|(name, value)| (name.to_owned(), value.to_owned()))
115 .collect();
116 let (name, value) = match params.iter().find(|(name, _)| name == M::TABLE) {
117 Some(param) => param.clone(),
118 None if params.len() == 1 => params[0].clone(),
119 None => {
120 return Err(anyhow::anyhow!(
121 "Found<{}> needs a route parameter named `{}` (the route has {})",
122 std::any::type_name::<M>(),
123 M::TABLE,
124 params
125 .iter()
126 .map(|(name, _)| format!("`{name}`"))
127 .collect::<Vec<_>>()
128 .join(", ")
129 )
130 .into());
131 }
132 };
133 let app = parts
134 .extensions
135 .get::<crate::AppState>()
136 .cloned()
137 .ok_or_else(|| anyhow::anyhow!("Found<T> needs Renox's request layers"))?;
138 let by_column = name != M::TABLE && name != "id" && M::COLUMNS.contains(&name.as_str());
139 let found = if by_column {
140 M::query().where_eq(&name, value).first(&app.db).await?
141 } else {
142 let key = value
143 .parse::<M::Key>()
144 .ok()
145 .filter(|key| !crate::db::ModelKey::is_unsaved(key));
146 match key {
147 Some(key) => M::find(&app.db, key).await?,
148 None => None,
149 }
150 };
151 found.map(Found).ok_or(Error::NotFound)
152 }
153}