use std::convert::Infallible;
use std::future::Future;
use std::pin::Pin;
use crate::application::ProxyFn;
use crate::axum::body::Body;
use crate::axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode, Uri};
use crate::proxy::action::Action;
use crate::proxy::request::Request as ProxyRequest;
type BoxFuture = Pin<Box<dyn Future<Output = Result<Response<Body>, Infallible>> + Send>>;
#[derive(Clone)]
pub struct ProxyLayer {
proxy: Option<ProxyFn>,
}
impl ProxyLayer {
#[must_use]
pub fn new(proxy: Option<ProxyFn>) -> Self {
Self { proxy }
}
}
impl<S> tower::Layer<S> for ProxyLayer {
type Service = ProxyService<S>;
fn layer(&self, inner: S) -> Self::Service {
ProxyService {
inner,
proxy: self.proxy.clone(),
}
}
}
#[derive(Clone)]
pub struct ProxyService<Inner> {
inner: Inner,
proxy: Option<ProxyFn>,
}
impl<Inner> tower::Service<crate::axum::extract::Request<Body>> for ProxyService<Inner>
where
Inner: tower::Service<
crate::axum::extract::Request<Body>,
Response = Response<Body>,
Error = Infallible,
> + Clone
+ Send
+ 'static,
Inner::Future: Send + 'static,
{
type Response = Response<Body>;
type Error = Infallible;
type Future = BoxFuture;
fn poll_ready(
&mut self,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: crate::axum::extract::Request<Body>) -> Self::Future {
let Some(proxy) = &self.proxy else {
let fut = self.inner.call(req);
return Box::pin(fut);
};
let (mut parts, body) = req.into_parts();
let proxy_request = ProxyRequest::new(&parts.method, &parts.uri, &parts.headers);
let action = proxy(proxy_request);
match action {
Action::Continue { set_headers } => {
if let Err(rejection) = merge_headers(&mut parts.headers, set_headers) {
return Box::pin(async move { Ok(rejection) });
}
let req = crate::axum::extract::Request::from_parts(parts, body);
let fut = self.inner.call(req);
Box::pin(fut)
}
Action::Redirect {
location,
permanent,
} => {
match validate_header_value(&location) {
Ok(value) => {
let status = if permanent {
StatusCode::MOVED_PERMANENTLY
} else {
StatusCode::FOUND
};
let mut headers = HeaderMap::new();
headers.insert(HeaderName::from_static("location"), value);
Box::pin(async move { Ok(build_response(status, Some(headers), None)) })
}
Err(rejection) => Box::pin(async move { Ok(rejection) }),
}
}
Action::Rewrite { uri } => {
match validate_rewrite_uri(&uri) {
Ok(new_uri) => {
parts.uri = new_uri;
let req = crate::axum::extract::Request::from_parts(parts, body);
let fut = self.inner.call(req);
Box::pin(fut)
}
Err(rejection) => Box::pin(async move { Ok(rejection) }),
}
}
Action::ShortCircuit { status, response } => {
if let Some(response) = response {
Box::pin(async move { Ok(response) })
} else {
Box::pin(async move { Ok(build_response(status, None, None)) })
}
}
}
}
}
fn build_response(
status: StatusCode,
headers: Option<HeaderMap>,
body: Option<Body>,
) -> Response<Body> {
let mut builder = Response::builder().status(status);
builder = builder.header(
HeaderName::from_static("x-content-type-options"),
HeaderValue::from_static("nosniff"),
);
if let Some(h) = headers {
for (name, value) in h.iter() {
builder = builder.header(name.clone(), value.clone());
}
}
builder
.body(body.unwrap_or_else(Body::empty))
.unwrap_or_else(|_| {
Response::builder()
.status(StatusCode::INTERNAL_SERVER_ERROR)
.body(Body::empty())
.unwrap_or_else(|_| Response::new(Body::empty()))
})
}
#[allow(clippy::result_large_err)]
fn validate_header_value(value: &str) -> Result<HeaderValue, Response<Body>> {
if value
.bytes()
.any(|b| b == b'\r' || b == b'\n' || (b < 0x20 && b != b'\t'))
{
return Err(build_response(
StatusCode::BAD_REQUEST,
None,
Some(Body::from(
"invalid header value: CRLF or control character detected",
)),
));
}
HeaderValue::try_from(value).map_err(|_| {
build_response(
StatusCode::BAD_REQUEST,
None,
Some(Body::from("invalid header value")),
)
})
}
#[allow(clippy::result_large_err)]
fn validate_rewrite_uri(uri: &str) -> Result<Uri, Response<Body>> {
let parsed = uri.parse::<Uri>().map_err(|_| {
build_response(
StatusCode::BAD_REQUEST,
None,
Some(Body::from("invalid rewrite URI")),
)
})?;
if parsed.scheme_str().is_some() {
return Err(build_response(
StatusCode::BAD_REQUEST,
None,
Some(Body::from("rewrite URI must not contain a scheme")),
));
}
Ok(parsed)
}
#[allow(clippy::result_large_err)]
fn merge_headers(destination: &mut HeaderMap, source: HeaderMap) -> Result<(), Response<Body>> {
for (name, value) in source.iter() {
if value.as_bytes().contains(&b'\r') || value.as_bytes().contains(&b'\n') {
return Err(build_response(
StatusCode::BAD_REQUEST,
None,
Some(Body::from(
"invalid header value in SetHeaders: CRLF detected",
)),
));
}
destination.insert(name.clone(), value.clone());
}
Ok(())
}