use std::{
future::Future,
net::SocketAddr,
pin::Pin,
task::{Context, Poll},
};
use ::axum::{
body::{to_bytes, Body, HttpBody},
extract::{ConnectInfo, Request},
response::Response,
};
use bytes::Bytes;
use futures_util::stream;
use tower::{Layer, Service};
use crate::middleware::{
AuthorizationOutcome, DeniedResponse, Middleware, RequestContext, ResponseContext,
};
const MAX_CAPTURE_BYTES: usize = 4_004;
#[derive(Clone)]
pub struct ReqKeyLayer {
middleware: Middleware,
}
impl ReqKeyLayer {
pub const fn new(middleware: Middleware) -> Self {
Self { middleware }
}
}
impl<S> Layer<S> for ReqKeyLayer {
type Service = ReqKeyService<S>;
fn layer(&self, inner: S) -> Self::Service {
ReqKeyService {
inner,
middleware: self.middleware.clone(),
}
}
}
#[derive(Clone)]
pub struct ReqKeyService<S> {
inner: S,
middleware: Middleware,
}
impl<S> Service<Request> for ReqKeyService<S>
where
S: Service<Request, Response = Response> + Clone + Send + 'static,
S::Future: Send + 'static,
S::Error: Send + 'static,
{
type Response = Response;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Response, S::Error>> + Send>>;
fn poll_ready(&mut self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(context)
}
fn call(&mut self, mut request: Request) -> Self::Future {
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
let middleware = self.middleware.clone();
let context = request_context(&request);
Box::pin(async move {
match middleware.authorize(context).await {
AuthorizationOutcome::Bypass => inner.call(request).await,
AuthorizationOutcome::Denied(denial) => Ok(denial_response(&denial)),
AuthorizationOutcome::Authorized(authorized) => {
if let Some(decision) = &authorized.decision {
request.extensions_mut().insert(decision.clone());
}
if let Some(failure) = &authorized.failure {
request.extensions_mut().insert(failure.clone());
}
let mut response = inner.call(request).await?;
let original_headers = response.headers().clone();
let response_body = capture_response_body(
&mut response,
middleware.config().captures_response_body(),
)
.await;
response.headers_mut().extend(authorized.response_headers());
let response_context = ResponseContext::new(response.status().as_u16())
.with_headers(original_headers)
.with_latency_ms(authorized.elapsed_ms());
let response_context = response_body.map_or(response_context.clone(), |body| {
response_context.with_body(body)
});
middleware.record(authorized, response_context).await;
Ok(response)
}
}
})
}
}
fn request_context(request: &Request) -> RequestContext {
let mut context = RequestContext::new(request.method().clone(), request.uri().path())
.with_headers(request.headers().clone());
if let Some(query) = request.uri().query() {
context = context.with_query(query);
}
if let Some(ConnectInfo(address)) = request.extensions().get::<ConnectInfo<SocketAddr>>() {
context = context.with_client_ip(address.ip().to_string());
}
context
}
fn denial_response(denial: &DeniedResponse) -> Response {
let mut response = Response::new(Body::from(denial.json_body()));
*response.status_mut() = http::StatusCode::from_u16(denial.status_code)
.unwrap_or(http::StatusCode::INTERNAL_SERVER_ERROR);
response.headers_mut().extend(denial.headers());
response
}
async fn capture_response_body(response: &mut Response, enabled: bool) -> Option<String> {
if !enabled
|| response
.headers()
.get(http::header::CONTENT_ENCODING)
.is_some_and(|value| value != "identity")
{
return None;
}
let exact = usize::try_from(response.body().size_hint().exact()?).ok()?;
if exact > MAX_CAPTURE_BYTES {
return None;
}
let replacement = Response::new(Body::empty());
let original = std::mem::replace(response, replacement);
let (parts, body) = original.into_parts();
match to_bytes(body, MAX_CAPTURE_BYTES).await {
Ok(bytes) => {
let text = String::from_utf8_lossy(&bytes)
.chars()
.take(crate::client::MAX_BODY_CHARACTERS)
.collect();
*response = Response::from_parts(parts, Body::from(bytes));
Some(text)
}
Err(error) => {
let errors = stream::once(async move { Err::<Bytes, _>(error) });
*response = Response::from_parts(parts, Body::from_stream(errors));
None
}
}
}