diff --git a/crates/axum-utils/src/client_authorization.rs b/crates/axum-utils/src/client_authorization.rs index f7fed6111..3ba745e42 100644 --- a/crates/axum-utils/src/client_authorization.rs +++ b/crates/axum-utils/src/client_authorization.rs @@ -184,7 +184,8 @@ fn jwks_key_store(jwks: &JwksOrJwksUri) -> Either() + .response_body_to_bytes() + .json_response::() .map_request(move |_: ()| { http::Request::builder() .method("GET") diff --git a/crates/cli/src/commands/debug.rs b/crates/cli/src/commands/debug.rs index bbd0e6833..402cff130 100644 --- a/crates/cli/src/commands/debug.rs +++ b/crates/cli/src/commands/debug.rs @@ -89,7 +89,9 @@ impl Options { json: true, url, } => { - let mut client = mas_http::client("cli-debug-http").json(); + let mut client = mas_http::client("cli-debug-http") + .response_body_to_bytes() + .json_response(); let request = hyper::Request::builder() .uri(url) .body(hyper::Body::empty())?; diff --git a/crates/http/src/ext.rs b/crates/http/src/ext.rs index cd155d2b8..7862f53aa 100644 --- a/crates/http/src/ext.rs +++ b/crates/http/src/ext.rs @@ -14,9 +14,14 @@ use http::header::HeaderName; use once_cell::sync::OnceCell; +use tower::{layer::util::Stack, ServiceBuilder}; use tower_http::cors::CorsLayer; -use crate::layers::json::Json; +use crate::layers::{ + body_to_bytes::{BodyToBytes, BodyToBytesLayer}, + json_request::{JsonRequest, JsonRequestLayer}, + json_response::{JsonResponse, JsonResponseLayer}, +}; static PROPAGATOR_HEADERS: OnceCell> = OnceCell::new(); @@ -60,11 +65,37 @@ impl CorsLayerExt for CorsLayer { } pub trait ServiceExt: Sized { - fn json(self) -> Json; -} + fn response_body_to_bytes(self) -> BodyToBytes { + BodyToBytes::new(self) + } -impl ServiceExt for S { - fn json(self) -> Json { - Json::new(self) + fn json_response(self) -> JsonResponse { + JsonResponse::new(self) + } + + fn json_request(self) -> JsonRequest { + JsonRequest::new(self) + } +} + +impl ServiceExt for S {} + +pub trait ServiceBuilderExt: Sized { + fn response_to_bytes(self) -> ServiceBuilder>; + fn json_response(self) -> ServiceBuilder, L>>; + fn json_request(self) -> ServiceBuilder, L>>; +} + +impl ServiceBuilderExt for ServiceBuilder { + fn response_to_bytes(self) -> ServiceBuilder> { + self.layer(BodyToBytesLayer::default()) + } + + fn json_response(self) -> ServiceBuilder, L>> { + self.layer(JsonResponseLayer::default()) + } + + fn json_request(self) -> ServiceBuilder, L>> { + self.layer(JsonRequestLayer::default()) } } diff --git a/crates/http/src/layers/body_to_bytes.rs b/crates/http/src/layers/body_to_bytes.rs new file mode 100644 index 000000000..b2f833d8f --- /dev/null +++ b/crates/http/src/layers/body_to_bytes.rs @@ -0,0 +1,96 @@ +// Copyright 2022 The Matrix.org Foundation C.I.C. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +use bytes::Bytes; +use futures_util::future::BoxFuture; +use http::{Request, Response}; +use http_body::Body; +use thiserror::Error; +use tower::{Layer, Service}; + +#[derive(Debug, Error)] +pub enum Error { + #[error(transparent)] + Service { inner: ServiceError }, + + #[error(transparent)] + Body { inner: BodyError }, +} + +impl Error { + fn service(inner: S) -> Self { + Self::Service { inner } + } + + fn body(inner: B) -> Self { + Self::Body { inner } + } +} + +#[derive(Clone)] +pub struct BodyToBytes { + inner: S, +} + +impl BodyToBytes { + pub const fn new(inner: S) -> Self { + Self { inner } + } +} + +impl Service> for BodyToBytes +where + S: Service, Response = Response>, + S::Future: Send + 'static, + ResBody: Body + Send, + ResBody::Data: Send, +{ + type Error = Error; + type Response = Response; + type Future = BoxFuture<'static, Result>; + + fn poll_ready( + &mut self, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + self.inner.poll_ready(cx).map_err(Error::service) + } + + fn call(&mut self, request: Request) -> Self::Future { + let inner = self.inner.call(request); + + let fut = async { + let response = inner.await.map_err(Error::service)?; + let (parts, body) = response.into_parts(); + + let body = hyper::body::to_bytes(body).await.map_err(Error::body)?; + + let response = Response::from_parts(parts, body); + Ok(response) + }; + + Box::pin(fut) + } +} + +#[derive(Default, Clone, Copy)] +pub struct BodyToBytesLayer; + +impl Layer for BodyToBytesLayer { + type Service = BodyToBytes; + + fn layer(&self, inner: S) -> Self::Service { + BodyToBytes::new(inner) + } +} diff --git a/crates/http/src/layers/json_request.rs b/crates/http/src/layers/json_request.rs new file mode 100644 index 000000000..74f15705e --- /dev/null +++ b/crates/http/src/layers/json_request.rs @@ -0,0 +1,123 @@ +// Copyright 2022 The Matrix.org Foundation C.I.C. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +use std::{future::Ready, marker::PhantomData, task::Poll}; + +use bytes::Bytes; +use futures_util::{ + future::{Either, MapErr}, + FutureExt, TryFutureExt, +}; +use http::{header::CONTENT_TYPE, HeaderValue, Request}; +use http_body::Full; +use serde::Serialize; +use thiserror::Error; +use tower::{Layer, Service}; + +#[derive(Debug, Error)] +pub enum Error { + #[error(transparent)] + Service { inner: Service }, + + #[error("could not serialize JSON payload")] + Json { + #[source] + inner: serde_json::Error, + }, +} + +impl Error { + fn service(source: S) -> Self { + Self::Service { inner: source } + } + + fn json(source: serde_json::Error) -> Self { + Self::Json { inner: source } + } +} + +#[derive(Clone)] +pub struct JsonRequest { + inner: S, + _t: PhantomData, +} + +impl JsonRequest { + pub const fn new(inner: S) -> Self { + Self { + inner, + _t: PhantomData, + } + } +} + +impl Service> for JsonRequest +where + S: Service>>, + S::Future: Send + 'static, + S::Error: 'static, + T: Serialize, +{ + type Error = Error; + type Response = S::Response; + type Future = Either< + Ready>, + MapErr Self::Error>, + >; + + fn poll_ready(&mut self, cx: &mut std::task::Context<'_>) -> Poll> { + self.inner.poll_ready(cx).map_err(Error::service) + } + + fn call(&mut self, request: Request) -> Self::Future { + let (mut parts, body) = request.into_parts(); + + parts + .headers + .insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); + + let body = match serde_json::to_vec(&body) { + Ok(body) => Full::new(Bytes::from(body)), + Err(err) => return std::future::ready(Err(Error::json(err))).left_future(), + }; + + let request = Request::from_parts(parts, body); + + self.inner + .call(request) + .map_err(Error::service as fn(S::Error) -> Self::Error) + .right_future() + } +} + +#[derive(Clone, Copy)] +pub struct JsonRequestLayer { + _t: PhantomData, +} + +impl Default for JsonRequestLayer { + fn default() -> Self { + Self { + _t: PhantomData::default(), + } + } +} + +impl Layer for JsonRequestLayer { + type Service = JsonRequest; + + fn layer(&self, inner: S) -> Self::Service { + JsonRequest::new(inner) + } +} diff --git a/crates/http/src/layers/json.rs b/crates/http/src/layers/json_response.rs similarity index 61% rename from crates/http/src/layers/json.rs rename to crates/http/src/layers/json_response.rs index 30c282043..c0f565c2e 100644 --- a/crates/http/src/layers/json.rs +++ b/crates/http/src/layers/json_response.rs @@ -14,24 +14,18 @@ use std::{marker::PhantomData, task::Poll}; -use futures_util::future::BoxFuture; +use bytes::Buf; +use futures_util::FutureExt; use http::{header::ACCEPT, HeaderValue, Request, Response}; -use http_body::Body; use serde::de::DeserializeOwned; use thiserror::Error; use tower::{Layer, Service}; #[derive(Debug, Error)] -pub enum Error { +pub enum Error { #[error(transparent)] Service { inner: Service }, - #[error("failed to fully read the request body")] - Body { - #[source] - inner: Body, - }, - #[error("could not parse JSON payload")] Json { #[source] @@ -39,27 +33,23 @@ pub enum Error { }, } -impl Error { +impl Error { fn service(source: S) -> Self { Self::Service { inner: source } } - fn body(source: B) -> Self { - Self::Body { inner: source } - } - fn json(source: serde_json::Error) -> Self { Self::Json { inner: source } } } #[derive(Clone)] -pub struct Json { +pub struct JsonResponse { inner: S, _t: PhantomData, } -impl Json { +impl JsonResponse { pub const fn new(inner: S) -> Self { Self { inner, @@ -68,59 +58,64 @@ impl Json { } } -impl Service> for Json +impl Service> for JsonResponse where S: Service, Response = Response>, S::Future: Send + 'static, - C: Body + Send + 'static, - C::Data: Send + 'static, + C: Buf, T: DeserializeOwned, { - type Error = Error; + type Error = Error; type Response = Response; - type Future = BoxFuture<'static, Result>; + type Future = futures_util::future::Map< + S::Future, + fn(Result, S::Error>) -> Result, + >; fn poll_ready(&mut self, cx: &mut std::task::Context<'_>) -> Poll> { self.inner.poll_ready(cx).map_err(Error::service) } fn call(&mut self, mut request: Request) -> Self::Future { + fn mapper(res: Result, E>) -> Result, Error> + where + C: Buf, + T: DeserializeOwned, + { + let response = res.map_err(Error::service)?; + let (parts, body) = response.into_parts(); + + let body = serde_json::from_reader(body.reader()).map_err(Error::json)?; + + let res = Response::from_parts(parts, body); + Ok(res) + } + request .headers_mut() .insert(ACCEPT, HeaderValue::from_static("application/json")); - let fut = self.inner.call(request); - - let fut = async { - let response = fut.await.map_err(Error::service)?; - let (parts, body) = response.into_parts(); - - futures_util::pin_mut!(body); - let bytes = hyper::body::to_bytes(&mut body) - .await - .map_err(Error::body)?; - - let body = serde_json::from_slice(&bytes).map_err(Error::json)?; - - let res = Response::from_parts(parts, body); - Ok(res) - }; - - Box::pin(fut) + self.inner.call(request).map(mapper::) } } -#[derive(Default, Clone, Copy)] -pub struct JsonResponseLayer(PhantomData<(T, ReqBody)>); +#[derive(Clone, Copy)] +pub struct JsonResponseLayer { + _t: PhantomData, +} -impl Layer for JsonResponseLayer -where - S: Service, Response = Response>, - T: serde::de::DeserializeOwned, -{ - type Service = Json; +impl Default for JsonResponseLayer { + fn default() -> Self { + Self { + _t: PhantomData::default(), + } + } +} + +impl Layer for JsonResponseLayer { + type Service = JsonResponse; fn layer(&self, inner: S) -> Self::Service { - Json::new(inner) + JsonResponse::new(inner) } } diff --git a/crates/http/src/layers/mod.rs b/crates/http/src/layers/mod.rs index baefd54ec..8eefbb64d 100644 --- a/crates/http/src/layers/mod.rs +++ b/crates/http/src/layers/mod.rs @@ -12,7 +12,9 @@ // See the License for the specific language governing permissions and // limitations under the License. +pub(crate) mod body_to_bytes; pub(crate) mod client; -pub(crate) mod json; +pub(crate) mod json_request; +pub(crate) mod json_response; pub mod otel; pub(crate) mod server; diff --git a/crates/http/src/lib.rs b/crates/http/src/lib.rs index 74b484726..27b353615 100644 --- a/crates/http/src/lib.rs +++ b/crates/http/src/lib.rs @@ -50,7 +50,10 @@ mod layers; pub use self::{ ext::{set_propagator, CorsLayerExt, ServiceExt as HttpServiceExt}, future_service::FutureService, - layers::{client::ClientLayer, json::JsonResponseLayer, otel, server::ServerLayer}, + layers::{ + body_to_bytes::BodyToBytesLayer, client::ClientLayer, json_request::JsonRequestLayer, + json_response::JsonResponseLayer, otel, server::ServerLayer, + }, }; pub(crate) type BoxError = Box;