1use hyper::Request;
2use serde::de::value::{Error as ValueError, MapDeserializer};
3use serde::de::{self, DeserializeOwned, Deserializer, IntoDeserializer, Visitor};
4
5use crate::error::ServeError;
6use crate::router::PathParams;
7
8pub fn path_params<T: DeserializeOwned, B>(req: &Request<B>) -> Result<T, ServeError> {
30 let params = req
31 .extensions()
32 .get::<PathParams>()
33 .ok_or_else(|| ServeError::new(500, "no path params in request extensions"))?;
34
35 let pairs = params.0.iter().map(|(k, v)| (k.clone(), ParamValue(v.clone())));
36 let deserializer = MapDeserializer::<_, ValueError>::new(pairs);
37 T::deserialize(deserializer)
38 .map_err(|_| ServeError::new(400, "invalid path parameters"))
39}
40
41struct ParamValue(String);
45
46macro_rules! deserialize_parsed {
47 ($($method:ident => $visit:ident : $ty:ty),* $(,)?) => {
48 $(
49 fn $method<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, Self::Error> {
50 match self.0.parse::<$ty>() {
51 Ok(v) => visitor.$visit(v),
52 Err(_) => Err(de::Error::invalid_value(
53 de::Unexpected::Str(&self.0),
54 &stringify!($ty),
55 )),
56 }
57 }
58 )*
59 };
60}
61
62impl<'de> Deserializer<'de> for ParamValue {
63 type Error = ValueError;
64
65 fn deserialize_any<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, Self::Error> {
66 visitor.visit_string(self.0)
67 }
68
69 fn deserialize_option<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, Self::Error> {
70 visitor.visit_some(self)
71 }
72
73 deserialize_parsed! {
74 deserialize_bool => visit_bool: bool,
75 deserialize_i8 => visit_i8: i8,
76 deserialize_i16 => visit_i16: i16,
77 deserialize_i32 => visit_i32: i32,
78 deserialize_i64 => visit_i64: i64,
79 deserialize_i128 => visit_i128: i128,
80 deserialize_u8 => visit_u8: u8,
81 deserialize_u16 => visit_u16: u16,
82 deserialize_u32 => visit_u32: u32,
83 deserialize_u64 => visit_u64: u64,
84 deserialize_u128 => visit_u128: u128,
85 deserialize_f32 => visit_f32: f32,
86 deserialize_f64 => visit_f64: f64,
87 deserialize_char => visit_char: char,
88 }
89
90 serde::forward_to_deserialize_any! {
91 str string bytes byte_buf unit unit_struct newtype_struct seq tuple
92 tuple_struct map struct enum identifier ignored_any
93 }
94}
95
96impl<'de> IntoDeserializer<'de, ValueError> for ParamValue {
97 type Deserializer = Self;
98
99 fn into_deserializer(self) -> Self {
100 self
101 }
102}
103
104#[cfg(test)]
105mod tests {
106 use super::*;
107 use crate::router::PathParams;
108 use serde::Deserialize;
109 use std::collections::HashMap;
110
111 #[derive(Debug, Deserialize)]
112 struct Item {
113 id: u64,
114 }
115
116 fn request_with_params(params: HashMap<String, String>) -> Request<()> {
117 Request::builder()
118 .extension(PathParams(params))
119 .body(())
120 .unwrap()
121 }
122
123 #[test]
124 fn extracts_typed_numeric_field() {
125 let mut params = HashMap::new();
126 params.insert("id".to_string(), "42".to_string());
127 let req = request_with_params(params);
128
129 let item: Item = path_params(&req).unwrap();
130 assert_eq!(item.id, 42);
131 }
132
133 #[test]
134 fn unparseable_segment_returns_400() {
135 let mut params = HashMap::new();
136 params.insert("id".to_string(), "not-a-number".to_string());
137 let req = request_with_params(params);
138
139 let err = path_params::<Item, _>(&req).unwrap_err();
140 assert_eq!(err.code, 400);
141 assert_eq!(err.message, "invalid path parameters");
142 }
143
144 #[test]
145 fn missing_extensions_returns_500() {
146 let req = Request::builder().body(()).unwrap();
147 let err = path_params::<Item, _>(&req).unwrap_err();
148 assert_eq!(err.code, 500);
149 }
150}