use crate::call::Call;
use crate::error::{Error, Result};
use async_trait::async_trait;
use serde::de::DeserializeOwned;
use std::sync::Arc;
#[async_trait]
pub trait FromCallParts: Sized + Send {
async fn from_call_parts(call: &mut Call) -> Result<Self>;
}
#[async_trait]
pub trait FromCall: Sized + Send {
async fn from_call(call: Call) -> Result<Self>;
}
#[async_trait]
impl<T: FromCallParts> FromCall for T {
async fn from_call(mut call: Call) -> Result<Self> {
T::from_call_parts(&mut call).await
}
}
#[async_trait]
impl FromCall for Call {
async fn from_call(call: Call) -> Result<Self> {
Ok(call)
}
}
#[derive(Debug, Clone)]
pub struct Path<T>(
pub T,
);
#[async_trait]
impl<T> FromCallParts for Path<T>
where
T: std::str::FromStr + Send,
T::Err: std::fmt::Display,
{
async fn from_call_parts(call: &mut Call) -> Result<Self> {
let mut params = call.params_iter();
let (_name, raw) = params
.next()
.ok_or_else(|| Error::bad_request("no path parameter to extract"))?;
let value = raw
.parse::<T>()
.map_err(|e| Error::bad_request(format!("bad path param: {e}")))?;
Ok(Path(value))
}
}
#[derive(Debug, Clone)]
pub struct Query<T>(
pub T,
);
#[async_trait]
impl<T> FromCallParts for Query<T>
where
T: DeserializeOwned + Send,
{
async fn from_call_parts(call: &mut Call) -> Result<Self> {
let q = call.query_string();
let value = serde_urlencoded::from_str::<T>(q)
.map_err(|e| Error::bad_request(format!("invalid query string: {e}")))?;
Ok(Query(value))
}
}
#[derive(Debug, Clone)]
pub struct State<T>(
pub Arc<T>,
);
#[async_trait]
impl<T> FromCallParts for State<T>
where
T: Send + Sync + 'static,
{
async fn from_call_parts(call: &mut Call) -> Result<Self> {
match call.state::<T>() {
Some(v) => Ok(State(v)),
None => Err(Error::internal(format!(
"missing application state: {}",
std::any::type_name::<T>()
))),
}
}
}
impl<T> std::ops::Deref for State<T> {
type Target = T;
fn deref(&self) -> &T {
&self.0
}
}
#[derive(Debug, Clone)]
pub struct BearerToken(
).
pub String,
);
#[async_trait]
impl FromCallParts for BearerToken {
async fn from_call_parts(call: &mut Call) -> Result<Self> {
let raw = call.header("authorization").ok_or_else(|| {
Error::new(
http::StatusCode::UNAUTHORIZED,
"missing Authorization header",
)
})?;
let token = raw
.strip_prefix("Bearer ")
.or_else(|| raw.strip_prefix("bearer "))
.ok_or_else(|| Error::new(http::StatusCode::UNAUTHORIZED, "expected Bearer scheme"))?;
Ok(BearerToken(token.trim().to_string()))
}
}
#[cfg(test)]
mod tests {
use super::*;
use bytes::Bytes;
use http::{HeaderMap, Method, Uri};
struct MethodName(String);
#[async_trait]
impl FromCallParts for MethodName {
async fn from_call_parts(call: &mut Call) -> Result<Self> {
Ok(MethodName(call.method().as_str().to_string()))
}
}
fn call() -> Call {
Call::new(
Method::GET,
"/".parse::<Uri>().unwrap(),
HeaderMap::new(),
Bytes::new(),
)
}
#[tokio::test]
async fn parts_extractor_runs() {
let mut c = call();
let m = MethodName::from_call_parts(&mut c).await.unwrap();
assert_eq!(m.0, "GET");
}
#[tokio::test]
async fn call_is_from_call() {
let c = call();
let back = Call::from_call(c).await.unwrap();
assert_eq!(back.method(), &Method::GET);
}
use std::collections::HashMap;
#[tokio::test]
async fn path_extracts_single_param() {
let mut c = call();
let mut p = HashMap::new();
p.insert("id".to_string(), "42".to_string());
c.set_params(p);
let Path(id) = Path::<u64>::from_call_parts(&mut c).await.unwrap();
assert_eq!(id, 42);
}
#[tokio::test]
async fn path_bad_value_is_400() {
let mut c = call();
let mut p = HashMap::new();
p.insert("id".to_string(), "notnum".to_string());
c.set_params(p);
let err = Path::<u64>::from_call_parts(&mut c).await.unwrap_err();
assert_eq!(err.status(), http::StatusCode::BAD_REQUEST);
}
use serde::Deserialize;
#[derive(Deserialize, Debug, PartialEq)]
struct Pager {
page: u32,
q: String,
}
fn call_with_query(qs: &str) -> Call {
Call::new(
Method::GET,
format!("/s?{qs}").parse::<Uri>().unwrap(),
HeaderMap::new(),
Bytes::new(),
)
}
#[tokio::test]
async fn query_deserializes() {
let mut c = call_with_query("page=2&q=rust");
let Query(p) = Query::<Pager>::from_call_parts(&mut c).await.unwrap();
assert_eq!(
p,
Pager {
page: 2,
q: "rust".into()
}
);
}
#[tokio::test]
async fn query_missing_field_is_400() {
let mut c = call_with_query("q=rust");
let err = Query::<Pager>::from_call_parts(&mut c).await.unwrap_err();
assert_eq!(err.status(), http::StatusCode::BAD_REQUEST);
}
fn call_with_auth(value: &str) -> Call {
let mut headers = HeaderMap::new();
headers.insert(
http::header::AUTHORIZATION,
http::HeaderValue::from_str(value).unwrap(),
);
Call::new(
Method::GET,
"/".parse::<Uri>().unwrap(),
headers,
Bytes::new(),
)
}
#[tokio::test]
async fn bearer_token_extracted() {
let mut c = call_with_auth("Bearer abc123");
let BearerToken(t) = BearerToken::from_call_parts(&mut c).await.unwrap();
assert_eq!(t, "abc123");
}
#[tokio::test]
async fn missing_bearer_is_401() {
let mut c = call();
let err = BearerToken::from_call_parts(&mut c).await.unwrap_err();
assert_eq!(err.status(), http::StatusCode::UNAUTHORIZED);
}
}