use futures_util::ready;
use http::{header::HeaderName, Request, Response};
use pin_project_lite::pin_project;
use std::{
future::Future,
pin::Pin,
sync::Arc,
task::{Context, Poll},
};
use tower_layer::Layer;
use tower_service::Service;
#[derive(Clone, Debug)]
pub struct SetSensitiveHeadersLayer {
headers: Arc<[HeaderName]>,
}
impl SetSensitiveHeadersLayer {
pub fn new<I>(headers: I) -> Self
where
I: IntoIterator<Item = HeaderName>,
{
let headers = headers.into_iter().collect::<Vec<_>>();
Self::from_shared(headers.into())
}
pub fn from_shared(headers: Arc<[HeaderName]>) -> Self {
Self { headers }
}
}
impl<S> Layer<S> for SetSensitiveHeadersLayer {
type Service = SetSensitiveHeaders<S>;
fn layer(&self, inner: S) -> Self::Service {
SetSensitiveRequestHeaders::from_shared(
SetSensitiveResponseHeaders::from_shared(inner, self.headers.clone()),
self.headers.clone(),
)
}
}
pub type SetSensitiveHeaders<S> = SetSensitiveRequestHeaders<SetSensitiveResponseHeaders<S>>;
#[derive(Clone, Debug)]
pub struct SetSensitiveRequestHeadersLayer {
headers: Arc<[HeaderName]>,
}
impl SetSensitiveRequestHeadersLayer {
pub fn new<I>(headers: I) -> Self
where
I: IntoIterator<Item = HeaderName>,
{
let headers = headers.into_iter().collect::<Vec<_>>();
Self::from_shared(headers.into())
}
pub fn from_shared(headers: Arc<[HeaderName]>) -> Self {
Self { headers }
}
}
impl<S> Layer<S> for SetSensitiveRequestHeadersLayer {
type Service = SetSensitiveRequestHeaders<S>;
fn layer(&self, inner: S) -> Self::Service {
SetSensitiveRequestHeaders {
inner,
headers: self.headers.clone(),
}
}
}
#[derive(Clone, Debug)]
pub struct SetSensitiveRequestHeaders<S> {
inner: S,
headers: Arc<[HeaderName]>,
}
impl<S> SetSensitiveRequestHeaders<S> {
pub fn new<I>(inner: S, headers: I) -> Self
where
I: IntoIterator<Item = HeaderName>,
{
let headers = headers.into_iter().collect::<Vec<_>>();
Self::from_shared(inner, headers.into())
}
pub fn from_shared(inner: S, headers: Arc<[HeaderName]>) -> Self {
Self { inner, headers }
}
define_inner_service_accessors!();
pub fn layer<I>(headers: I) -> SetSensitiveRequestHeadersLayer
where
I: IntoIterator<Item = HeaderName>,
{
SetSensitiveRequestHeadersLayer::new(headers)
}
}
impl<ReqBody, ResBody, S> Service<Request<ReqBody>> for SetSensitiveRequestHeaders<S>
where
S: Service<Request<ReqBody>, Response = Response<ResBody>>,
{
type Response = S::Response;
type Error = S::Error;
type Future = S::Future;
#[inline]
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut req: Request<ReqBody>) -> Self::Future {
for header in &*self.headers {
if let Some(value) = req.headers_mut().get_mut(header) {
value.set_sensitive(true);
}
}
self.inner.call(req)
}
}
#[derive(Clone, Debug)]
pub struct SetSensitiveResponseHeadersLayer {
headers: Arc<[HeaderName]>,
}
impl SetSensitiveResponseHeadersLayer {
pub fn new<I>(headers: I) -> Self
where
I: IntoIterator<Item = HeaderName>,
{
let headers = headers.into_iter().collect::<Vec<_>>();
Self::from_shared(headers.into())
}
pub fn from_shared(headers: Arc<[HeaderName]>) -> Self {
Self { headers }
}
}
impl<S> Layer<S> for SetSensitiveResponseHeadersLayer {
type Service = SetSensitiveResponseHeaders<S>;
fn layer(&self, inner: S) -> Self::Service {
SetSensitiveResponseHeaders {
inner,
headers: self.headers.clone(),
}
}
}
#[derive(Clone, Debug)]
pub struct SetSensitiveResponseHeaders<S> {
inner: S,
headers: Arc<[HeaderName]>,
}
impl<S> SetSensitiveResponseHeaders<S> {
pub fn new<I>(inner: S, headers: I) -> Self
where
I: IntoIterator<Item = HeaderName>,
{
let headers = headers.into_iter().collect::<Vec<_>>();
Self::from_shared(inner, headers.into())
}
pub fn from_shared(inner: S, headers: Arc<[HeaderName]>) -> Self {
Self { inner, headers }
}
define_inner_service_accessors!();
pub fn layer<I>(headers: I) -> SetSensitiveResponseHeadersLayer
where
I: IntoIterator<Item = HeaderName>,
{
SetSensitiveResponseHeadersLayer::new(headers)
}
}
impl<ReqBody, ResBody, S> Service<Request<ReqBody>> for SetSensitiveResponseHeaders<S>
where
S: Service<Request<ReqBody>, Response = Response<ResBody>>,
{
type Response = S::Response;
type Error = S::Error;
type Future = SetSensitiveResponseHeadersResponseFuture<S::Future>;
#[inline]
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: Request<ReqBody>) -> Self::Future {
SetSensitiveResponseHeadersResponseFuture {
future: self.inner.call(req),
headers: self.headers.clone(),
}
}
}
pin_project! {
#[derive(Debug)]
pub struct SetSensitiveResponseHeadersResponseFuture<F> {
#[pin]
future: F,
headers: Arc<[HeaderName]>,
}
}
impl<F, ResBody, E> Future for SetSensitiveResponseHeadersResponseFuture<F>
where
F: Future<Output = Result<Response<ResBody>, E>>,
{
type Output = F::Output;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.project();
let mut res = ready!(this.future.poll(cx)?);
for header in &**this.headers {
if let Some(value) = res.headers_mut().get_mut(header) {
value.set_sensitive(true);
}
}
Poll::Ready(Ok(res))
}
}