Skip to main content

mini_serve/
extract.rs

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
8/// Extract path parameters from the request and deserialize into type `T`.
9///
10/// Path parameters are decoded and matched by the router, then deserialized
11/// via serde's `MapDeserializer`. Returns `400 Bad Request` if deserialization
12/// fails (e.g., an unparseable segment for a numeric type) or if the request
13/// lacks extracted path parameters.
14///
15/// # Example
16///
17/// ```ignore
18/// use serde::Deserialize;
19/// use mini_serve::path_params;
20///
21/// #[derive(Deserialize)]
22/// struct ItemId {
23///     id: u64,
24/// }
25///
26/// let item = path_params::<ItemId, _>(req)?;
27/// println!("Item ID: {}", item.id);
28/// ```
29pub 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
41/// Deserializes a single path-param string into whatever scalar type the
42/// target struct field asks for, parsing on demand rather than going through
43/// an intermediate query-string representation.
44struct 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}