use hyper::Request;
use serde::de::value::{Error as ValueError, MapDeserializer};
use serde::de::{self, DeserializeOwned, Deserializer, IntoDeserializer, Visitor};
use crate::error::ServeError;
use crate::router::PathParams;
pub fn path_params<T: DeserializeOwned, B>(req: &Request<B>) -> Result<T, ServeError> {
let params = req
.extensions()
.get::<PathParams>()
.ok_or_else(|| ServeError::new(500, "no path params in request extensions"))?;
let pairs = params.0.iter().map(|(k, v)| (k.clone(), ParamValue(v.clone())));
let deserializer = MapDeserializer::<_, ValueError>::new(pairs);
T::deserialize(deserializer)
.map_err(|_| ServeError::new(400, "invalid path parameters"))
}
struct ParamValue(String);
macro_rules! deserialize_parsed {
($($method:ident => $visit:ident : $ty:ty),* $(,)?) => {
$(
fn $method<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, Self::Error> {
match self.0.parse::<$ty>() {
Ok(v) => visitor.$visit(v),
Err(_) => Err(de::Error::invalid_value(
de::Unexpected::Str(&self.0),
&stringify!($ty),
)),
}
}
)*
};
}
impl<'de> Deserializer<'de> for ParamValue {
type Error = ValueError;
fn deserialize_any<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, Self::Error> {
visitor.visit_string(self.0)
}
fn deserialize_option<V: Visitor<'de>>(self, visitor: V) -> Result<V::Value, Self::Error> {
visitor.visit_some(self)
}
deserialize_parsed! {
deserialize_bool => visit_bool: bool,
deserialize_i8 => visit_i8: i8,
deserialize_i16 => visit_i16: i16,
deserialize_i32 => visit_i32: i32,
deserialize_i64 => visit_i64: i64,
deserialize_i128 => visit_i128: i128,
deserialize_u8 => visit_u8: u8,
deserialize_u16 => visit_u16: u16,
deserialize_u32 => visit_u32: u32,
deserialize_u64 => visit_u64: u64,
deserialize_u128 => visit_u128: u128,
deserialize_f32 => visit_f32: f32,
deserialize_f64 => visit_f64: f64,
deserialize_char => visit_char: char,
}
serde::forward_to_deserialize_any! {
str string bytes byte_buf unit unit_struct newtype_struct seq tuple
tuple_struct map struct enum identifier ignored_any
}
}
impl<'de> IntoDeserializer<'de, ValueError> for ParamValue {
type Deserializer = Self;
fn into_deserializer(self) -> Self {
self
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::router::PathParams;
use serde::Deserialize;
use std::collections::HashMap;
#[derive(Debug, Deserialize)]
struct Item {
id: u64,
}
fn request_with_params(params: HashMap<String, String>) -> Request<()> {
Request::builder()
.extension(PathParams(params))
.body(())
.unwrap()
}
#[test]
fn extracts_typed_numeric_field() {
let mut params = HashMap::new();
params.insert("id".to_string(), "42".to_string());
let req = request_with_params(params);
let item: Item = path_params(&req).unwrap();
assert_eq!(item.id, 42);
}
#[test]
fn unparseable_segment_returns_400() {
let mut params = HashMap::new();
params.insert("id".to_string(), "not-a-number".to_string());
let req = request_with_params(params);
let err = path_params::<Item, _>(&req).unwrap_err();
assert_eq!(err.code, 400);
assert_eq!(err.message, "invalid path parameters");
}
#[test]
fn missing_extensions_returns_500() {
let req = Request::builder().body(()).unwrap();
let err = path_params::<Item, _>(&req).unwrap_err();
assert_eq!(err.code, 500);
}
}