use std::task::{Context, Poll};
use axum::extract::Request;
use axum::response::{IntoResponse, Response};
use futures::{FutureExt, future::BoxFuture};
use tower::{Layer, Service};
use crate::error::ProxyError;
pub const DEFAULT_FORWARDED_USER_HEADER: &str = "x-forwarded-user";
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ForwardedIdentity(pub Option<String>);
impl ForwardedIdentity {
pub fn anonymous() -> Self {
Self(None)
}
pub fn user(name: impl Into<String>) -> Self {
Self(Some(name.into()))
}
}
impl crate::backend::ForwardedUser for ForwardedIdentity {
fn forwarded_user(&self) -> Option<&str> {
self.0.as_deref()
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum AuthMode {
Anonymous,
ReverseProxy {
header: String,
allow_missing: bool,
},
}
impl AuthMode {
fn authenticate(&self, request: &Request) -> Result<ForwardedIdentity, ProxyError> {
match self {
AuthMode::Anonymous => Ok(ForwardedIdentity::anonymous()),
AuthMode::ReverseProxy {
header,
allow_missing,
} => match request.headers().get(header) {
Some(value) => {
let name = value.to_str().map_err(|_| {
ProxyError::Unauthenticated(format!("`{header}` header is not valid UTF-8"))
})?;
let name = name.trim();
if name.is_empty() {
return Err(ProxyError::Unauthenticated(format!(
"`{header}` header is empty"
)));
}
Ok(ForwardedIdentity::user(name))
}
None if *allow_missing => Ok(ForwardedIdentity::anonymous()),
None => Err(ProxyError::Unauthenticated(format!(
"missing `{header}` header"
))),
},
}
}
}
#[derive(Clone)]
pub struct AuthLayer {
mode: AuthMode,
}
impl AuthLayer {
pub fn new(mode: AuthMode) -> Self {
Self { mode }
}
}
impl<S> Layer<S> for AuthLayer {
type Service = AuthMiddleware<S>;
fn layer(&self, inner: S) -> Self::Service {
AuthMiddleware {
inner,
mode: self.mode.clone(),
}
}
}
#[derive(Clone)]
pub struct AuthMiddleware<S> {
inner: S,
mode: AuthMode,
}
impl<S> Service<Request> for AuthMiddleware<S>
where
S: Service<Request, Response = Response> + Send + 'static,
S::Future: Send + 'static,
{
type Response = S::Response;
type Error = S::Error;
type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut req: Request) -> Self::Future {
match self.mode.authenticate(&req) {
Ok(identity) => {
req.extensions_mut().insert(identity);
self.inner.call(req).boxed()
}
Err(e) => async move { Ok(e.into_response()) }.boxed(),
}
}
}
#[cfg(test)]
mod tests {
use axum::body::Body;
use axum::http::StatusCode;
use tower::{ServiceBuilder, ServiceExt};
use super::*;
fn reverse_proxy(header: &str, allow_missing: bool) -> AuthMode {
AuthMode::ReverseProxy {
header: header.to_string(),
allow_missing,
}
}
#[test]
fn anonymous_always_anonymous() {
let req = Request::get("/").body(Body::empty()).unwrap();
assert_eq!(
AuthMode::Anonymous.authenticate(&req).unwrap(),
ForwardedIdentity::anonymous()
);
}
#[test]
fn reverse_proxy_extracts_forwarded_user() {
let mode = reverse_proxy(DEFAULT_FORWARDED_USER_HEADER, false);
let req = Request::get("/")
.header(DEFAULT_FORWARDED_USER_HEADER, "alice")
.body(Body::empty())
.unwrap();
assert_eq!(
mode.authenticate(&req).unwrap(),
ForwardedIdentity::user("alice")
);
}
#[test]
fn reverse_proxy_honors_custom_header_and_trims() {
let mode = reverse_proxy("x-user", false);
let req = Request::get("/")
.header("x-user", " bob ")
.body(Body::empty())
.unwrap();
assert_eq!(
mode.authenticate(&req).unwrap(),
ForwardedIdentity::user("bob")
);
}
#[test]
fn reverse_proxy_rejects_missing_header_by_default() {
let mode = reverse_proxy(DEFAULT_FORWARDED_USER_HEADER, false);
let req = Request::get("/").body(Body::empty()).unwrap();
assert!(mode.authenticate(&req).is_err());
}
#[test]
fn reverse_proxy_rejects_empty_header() {
let mode = reverse_proxy(DEFAULT_FORWARDED_USER_HEADER, false);
let req = Request::get("/")
.header(DEFAULT_FORWARDED_USER_HEADER, " ")
.body(Body::empty())
.unwrap();
assert!(mode.authenticate(&req).is_err());
}
#[test]
fn reverse_proxy_falls_back_to_anonymous_when_configured() {
let mode = reverse_proxy(DEFAULT_FORWARDED_USER_HEADER, true);
let req = Request::get("/").body(Body::empty()).unwrap();
assert_eq!(
mode.authenticate(&req).unwrap(),
ForwardedIdentity::anonymous()
);
}
#[tokio::test]
async fn layer_inserts_identity_extension() {
async fn echo(req: Request) -> Result<Response, std::convert::Infallible> {
assert!(req.extensions().get::<ForwardedIdentity>().is_some());
Ok(Response::new(Body::empty()))
}
let mut service = ServiceBuilder::new()
.layer(AuthLayer::new(AuthMode::Anonymous))
.service_fn(echo);
let req = Request::get("/").body(Body::empty()).unwrap();
let resp = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn layer_missing_identity_maps_to_401() {
async fn ok(_req: Request) -> Result<Response, std::convert::Infallible> {
Ok(Response::new(Body::empty()))
}
let mut service = ServiceBuilder::new()
.layer(AuthLayer::new(reverse_proxy(
DEFAULT_FORWARDED_USER_HEADER,
false,
)))
.service_fn(ok);
let req = Request::get("/").body(Body::empty()).unwrap();
let resp = service.ready().await.unwrap().call(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
}