use serde::de::DeserializeOwned;
use serde::{Deserialize, Deserializer, Serialize};
use crate::error::{codes, ApiError};
use crate::routes::Route;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum PayloadKind {
Empty,
Json,
Query,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash, Serialize)]
#[non_exhaustive]
pub struct NoPayload {}
impl NoPayload {
pub const fn new() -> Self {
Self {}
}
}
impl<'de> Deserialize<'de> for NoPayload {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
serde::de::IgnoredAny::deserialize(deserializer)?;
Ok(NoPayload {})
}
}
pub static NO_PAYLOAD: NoPayload = NoPayload {};
pub trait HttpCall: Sized {
type Payload: Serialize + DeserializeOwned;
type Response: Serialize + DeserializeOwned;
const ROUTE: Route;
const PAYLOAD: PayloadKind;
fn payload(&self) -> &Self::Payload;
fn path_params(&self) -> PathParams {
PathParams::new()
}
fn from_parts(params: &PathParams, payload: Self::Payload) -> Result<Self, ApiError>;
fn path(&self) -> Option<String> {
self.path_params().fill(Self::ROUTE.path)
}
}
pub fn is_path_safe(value: &str) -> bool {
(1..=256).contains(&value.len())
&& value != "."
&& value != ".."
&& value.bytes().all(|b| {
b.is_ascii_alphanumeric()
|| matches!(b, b'-' | b'.' | b'_' | b'~' | b'!' | b'$' | b'&' | b'\'' | b'(' | b')' | b'*' | b'+' | b',' | b';' | b'=' | b':' | b'@')
})
}
pub fn query_pairs<T: Serialize + ?Sized>(payload: &T) -> Result<Vec<(String, String)>, ApiError> {
let bad = |message: &str| ApiError::new(codes::BAD_REQUEST, message);
let value = serde_json::to_value(payload).map_err(|_| bad("the query payload cannot be encoded"))?;
let serde_json::Value::Object(fields) = value else { return Err(bad("a query payload must be an object")) };
let mut pairs = Vec::with_capacity(fields.len());
for (name, value) in fields {
let text = match value {
serde_json::Value::Null => continue,
serde_json::Value::String(text) => text,
serde_json::Value::Bool(flag) => flag.to_string(),
serde_json::Value::Number(number) => number.to_string(),
_ => return Err(bad("a query payload's fields must be strings, numbers or booleans")),
};
pairs.push((name, text));
}
Ok(pairs)
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct PathParams(Vec<(String, String)>);
impl PathParams {
pub fn new() -> Self {
Self::default()
}
pub fn with(mut self, name: &str, value: impl ToString) -> Self {
self.insert(name, value.to_string());
self
}
pub fn insert(&mut self, name: impl Into<String>, value: impl Into<String>) {
let (name, value) = (name.into(), value.into());
match self.0.iter_mut().find(|(n, _)| *n == name) {
Some(slot) => slot.1 = value,
None => self.0.push((name, value)),
}
}
pub fn get(&self, name: &str) -> Option<&str> {
self.0.iter().find(|(n, _)| n == name).map(|(_, v)| v.as_str())
}
pub fn iter(&self) -> impl Iterator<Item = (&str, &str)> {
self.0.iter().map(|(n, v)| (n.as_str(), v.as_str()))
}
pub fn require(&self, name: &str) -> Result<&str, ApiError> {
self.get(name).ok_or_else(|| ApiError::new(codes::BAD_REQUEST, format!("the path parameter `{name}` is missing")))
}
pub fn id<T: From<i64>>(&self, name: &str) -> Result<T, ApiError> {
self.require(name)?.parse::<i64>().map(T::from).map_err(|_| ApiError::new(codes::BAD_REQUEST, format!("the path parameter `{name}` is not a number")))
}
pub fn checked(&self, name: &str, rule: impl Fn(&str) -> bool, problem: &str) -> Result<String, ApiError> {
let value = self.require(name)?;
if rule(value) {
Ok(value.to_string())
} else {
Err(ApiError::new(codes::BAD_REQUEST, format!("the path parameter `{name}` {problem}")))
}
}
pub fn fill(&self, template: &str) -> Option<String> {
let mut out = String::with_capacity(template.len() + 16);
let mut rest = template;
while let Some(open) = rest.find('{') {
out.push_str(&rest[..open]);
let after = &rest[open + 1..];
let close = after.find('}')?;
let value = self.get(&after[..close])?;
if !is_path_safe(value) {
return None;
}
out.push_str(value);
rest = &after[close + 1..];
}
out.push_str(rest);
Some(out)
}
}
pub fn placeholders(template: &str) -> Vec<&str> {
let mut names = Vec::new();
let mut rest = template;
while let Some(open) = rest.find('{') {
let after = &rest[open + 1..];
let Some(close) = after.find('}') else { break };
names.push(&after[..close]);
rest = &after[close + 1..];
}
names
}
macro_rules! payload_call {
($ty:ty, $method:ident, $path:expr, $auth:expr, $kind:ident, $response:ty) => {
impl $crate::http_call::HttpCall for $ty {
type Payload = Self;
type Response = $response;
const ROUTE: $crate::routes::Route = $crate::routes::Route::new($crate::routes::HttpMethod::$method, $path, $auth);
const PAYLOAD: $crate::http_call::PayloadKind = $crate::http_call::PayloadKind::$kind;
fn payload(&self) -> &Self {
self
}
fn from_parts(_params: &$crate::http_call::PathParams, payload: Self) -> Result<Self, $crate::error::ApiError> {
Ok(payload)
}
}
};
}
pub(crate) use payload_call;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn fill_and_safety() {
let params = PathParams::new().with("collection", "saves").with("key", "slot-1");
assert_eq!(params.fill("/v1/storage/{collection}/{key}").as_deref(), Some("/v1/storage/saves/slot-1"));
assert_eq!(params.fill("/v1/info").as_deref(), Some("/v1/info"));
assert_eq!(PathParams::new().fill("/v1/storage/{collection}"), None);
assert_eq!(PathParams::new().with("x", "a/b").fill("/{x}"), None);
for bad in ["", ".", "..", "a b", "a%2F", "a?b", "a#b", "a\\b", "ä", &"x".repeat(257)] {
assert!(!is_path_safe(bad), "{bad:?}");
}
assert!(is_path_safe("slot-1.json") && is_path_safe("42") && is_path_safe("-7"));
assert_eq!(placeholders("/v1/admin/users/{user}/roles/{role}"), ["user", "role"]);
let mut replaced = PathParams::new().with("a", 1);
replaced.insert("a", "2");
assert_eq!(replaced.get("a"), Some("2"));
assert_eq!(replaced.iter().count(), 1);
}
#[test]
fn parameter_errors() {
let params = PathParams::new().with("user", "x").with("name", "../a");
assert_eq!(params.id::<crate::UserId>("user").err().map(|e| e.code), Some(codes::BAD_REQUEST.to_string()));
assert_eq!(params.id::<crate::UserId>("missing").err().map(|e| e.code), Some(codes::BAD_REQUEST.to_string()));
assert_eq!(PathParams::new().with("user", "-3").id::<crate::UserId>("user").ok(), Some(crate::UserId(-3)));
let error = params.checked("name", crate::storage::is_valid_name, "is not a valid storage name").err();
assert_eq!(error.as_ref().map(|e| e.code.as_str()), Some(codes::BAD_REQUEST));
assert_eq!(error.map(|e| e.message).as_deref(), Some("the path parameter `name` is not a valid storage name"));
assert_eq!(serde_json::to_string(&NoPayload::new()).ok().as_deref(), Some("{}"));
assert_eq!(serde_json::from_str::<NoPayload>("null").ok(), Some(NoPayload::new()));
}
}