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]
#[diagnostic::do_not_recommend]
impl<T: FromCallParts> FromCall for T {
async fn from_call(mut call: Call) -> Result<Self> {
T::from_call_parts(&mut call).await
}
}
#[async_trait]
pub trait OptionalFromCallParts: Sized + Send {
async fn from_call_parts_opt(call: &mut Call) -> Result<Option<Self>>;
}
#[async_trait]
impl<T: OptionalFromCallParts> FromCallParts for Option<T> {
async fn from_call_parts(call: &mut Call) -> Result<Self> {
T::from_call_parts_opt(call).await
}
}
#[async_trait]
impl<T> OptionalFromCallParts for Query<T>
where
T: DeserializeOwned + Send,
{
async fn from_call_parts_opt(call: &mut Call) -> Result<Option<Self>> {
if call.query_string().is_empty() {
return Ok(None);
}
Query::<T>::from_call_parts(call).await.map(Some)
}
}
#[async_trait]
impl<T> OptionalFromCallParts for Path<T>
where
T: DeserializeOwned + Send,
{
async fn from_call_parts_opt(call: &mut Call) -> Result<Option<Self>> {
if call.params().is_empty() {
return Ok(None);
}
Path::<T>::from_call_parts(call).await.map(Some)
}
}
#[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: serde::de::DeserializeOwned + Send,
{
async fn from_call_parts(call: &mut Call) -> Result<Self> {
crate::path_de::from_params::<T>(call.params())
.map(Path)
.map_err(|e| Error::bad_request(format!("bad path parameters: {e}")))
}
}
#[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_html_form::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
.split_once(' ')
.filter(|(scheme, _)| scheme.eq_ignore_ascii_case("bearer"))
.map(|(_, token)| token)
.ok_or_else(|| Error::new(http::StatusCode::UNAUTHORIZED, "expected Bearer scheme"))?;
Ok(BearerToken(token.trim().to_string()))
}
}
#[async_trait]
impl OptionalFromCallParts for BearerToken {
async fn from_call_parts_opt(call: &mut Call) -> Result<Option<Self>> {
if call.header("authorization").is_none() {
return Ok(None);
}
BearerToken::from_call_parts(call).await.map(Some)
}
}
pub struct Header<T, N: HeaderName>(
pub T,
pub std::marker::PhantomData<N>,
);
pub trait HeaderName: Send + Sync + 'static {
const NAME: &'static str;
}
#[async_trait]
impl<T, N> FromCallParts for Header<T, N>
where
T: std::str::FromStr + Send,
T::Err: std::fmt::Display,
N: HeaderName,
{
async fn from_call_parts(call: &mut Call) -> Result<Self> {
let raw = call
.header(N::NAME)
.ok_or_else(|| Error::bad_request(format!("missing header `{}`", N::NAME)))?;
raw.parse::<T>()
.map(|v| Header(v, std::marker::PhantomData))
.map_err(|e| Error::bad_request(format!("bad header `{}`: {e}", N::NAME)))
}
}
#[async_trait]
impl<T, N> OptionalFromCallParts for Header<T, N>
where
T: std::str::FromStr + Send,
T::Err: std::fmt::Display,
N: HeaderName,
{
async fn from_call_parts_opt(call: &mut Call) -> Result<Option<Self>> {
if call.header(N::NAME).is_none() {
return Ok(None);
}
Header::<T, N>::from_call_parts(call).await.map(Some)
}
}
#[async_trait]
impl FromCall for String {
async fn from_call(mut call: Call) -> Result<Self> {
let b = call.try_receive_bytes().await?;
check_body_limit(&call, b.len())?;
String::from_utf8(b.to_vec())
.map_err(|_| Error::bad_request("request body is not valid UTF-8"))
}
}
#[async_trait]
impl FromCall for bytes::Bytes {
async fn from_call(mut call: Call) -> Result<Self> {
let b = call.try_receive_bytes().await?;
check_body_limit(&call, b.len())?;
Ok(b)
}
}
pub enum Either<L, R> {
Left(L),
Right(R),
}
#[async_trait]
impl<L, R> FromCall for Either<L, R>
where
L: FromCall,
R: FromCall,
{
async fn from_call(call: Call) -> Result<Self> {
let mut call = call;
let body = call.try_receive_bytes().await?;
let left_call = call.clone_with_body(body.clone());
if let Ok(l) = L::from_call(left_call).await {
return Ok(Either::Left(l));
}
R::from_call(call.clone_with_body(body))
.await
.map(Either::Right)
}
}
pub struct Payload(
pub crate::call::BodyStream,
);
#[async_trait]
impl FromCall for Payload {
async fn from_call(mut call: Call) -> Result<Self> {
let route_limit = call.get::<RouteBodyLimit>().map(|RouteBodyLimit(n)| n);
let stream = call
.body_stream()
.unwrap_or_else(|| Box::pin(futures_util::stream::empty()));
let Some(max) = route_limit else {
return Ok(Payload(stream));
};
use futures_util::StreamExt;
let mut seen = 0usize;
let counted = stream.map(move |chunk| match chunk {
Ok(bytes) => {
seen += bytes.len();
if seen > max {
Err(Error::new(
http::StatusCode::PAYLOAD_TOO_LARGE,
"request body too large",
))
} else {
Ok(bytes)
}
}
Err(e) => Err(e),
});
Ok(Payload(Box::pin(counted)))
}
}
#[derive(Debug, Clone, Copy)]
pub struct RouteBodyLimit(pub usize);
pub fn check_body_limit(call: &Call, len: usize) -> Result<()> {
match call.get::<RouteBodyLimit>() {
Some(RouteBodyLimit(max)) if len > max => Err(Error::new(
http::StatusCode::PAYLOAD_TOO_LARGE,
"request body too large",
)),
_ => Ok(()),
}
}
pub struct Form<T>(
pub T,
);
#[async_trait]
impl<T> FromCall for Form<T>
where
T: DeserializeOwned + Send,
{
async fn from_call(mut call: Call) -> Result<Self> {
let ct = call
.header(http::header::CONTENT_TYPE.as_str())
.unwrap_or("")
.to_string();
let media = ct.split(';').next().unwrap_or("").trim();
if !media.eq_ignore_ascii_case("application/x-www-form-urlencoded") {
return Err(Error::new(
http::StatusCode::UNSUPPORTED_MEDIA_TYPE,
"expected application/x-www-form-urlencoded",
));
}
let body = call.try_receive_bytes().await?;
check_body_limit(&call, body.len())?;
serde_html_form::from_bytes::<T>(&body)
.map(Form)
.map_err(|e| Error::bad_request(format!("invalid form body: {e}")))
}
}
#[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);
}
#[tokio::test]
async fn path_extracts_single_param() {
let mut c = call();
let mut p = crate::call::Params::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 = crate::call::Params::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);
}
}