use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use axum::{
body::Body,
extract::{ConnectInfo, State},
http::Request,
middleware::Next,
response::{IntoResponse, Response},
};
use tracing::{Span, field, warn};
use acme_proxy_core::client::ClientIp;
use acme_proxy_core::error::Problem;
use acme_proxy_policy::filter::ConnectionContext;
use acme_proxy_policy::filter::FilterPolicy;
use acme_proxy_policy::filter::Outcome;
pub async fn add_filter_middleware(
State(policy): State<Arc<FilterPolicy>>,
mut request: Request<Body>,
next: Next,
) -> Response {
let peer = request
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|ConnectInfo(addr)| addr.ip());
let client_ip = policy.proxy().resolve(peer, request.headers());
request.extensions_mut().insert(ClientIp(client_ip));
if let Some(ip) = client_ip {
Span::current().record("client_ip", field::display(ip));
}
let path = request.uri().path().to_string();
let context = ConnectionContext {
client_ip,
method: request.method(),
path: &path,
};
match policy.check_connection(&context).await {
Outcome::Allow => next.run(request).await,
outcome => problem_for(&outcome, client_ip, &path).into_response(),
}
}
fn problem_for(outcome: &Outcome, client_ip: Option<IpAddr>, path: &str) -> Problem {
match outcome {
Outcome::Deny(detail) => {
warn!(event = "filter_request_blocked", outcome = "failure", client_ip = ?client_ip, path, %detail);
Problem::access_denied(detail.clone())
}
Outcome::Undecided(_) | Outcome::Allow => {
Problem::server_internal("Request filtering failed")
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::StatusCode;
use axum::{Router, middleware, routing::get};
use tower::ServiceExt;
use acme_proxy_core::client::ProxyPolicy;
use acme_proxy_policy::filter::Effect;
fn app(trusted: &[String]) -> Router {
let proxy = ProxyPolicy::new(trusted, "x-forwarded-for").expect("policy must build");
let policy = Arc::new(FilterPolicy::new(
Vec::new(),
Vec::new(),
Effect::Allow,
proxy,
));
Router::new()
.route("/x", get(|| async { "ok" }))
.layer(middleware::from_fn_with_state(
policy,
add_filter_middleware,
))
.layer(middleware::from_fn(
crate::middlewares::access::add_access_middleware,
))
}
fn request(peer: [u8; 4], forwarded: Option<&str>) -> Request<Body> {
let mut builder = Request::builder().uri("/x");
if let Some(value) = forwarded {
builder = builder.header("x-forwarded-for", value);
}
let mut request = builder.body(Body::empty()).unwrap();
request
.extensions_mut()
.insert(ConnectInfo(SocketAddr::from((peer, 4711))));
request
}
#[tokio::test]
async fn a_trusted_proxys_forwarded_client_replaces_the_peer_on_the_span() {
let app = app(&["10.0.0.0/8".to_string()]);
let fields = acme_proxy_core::testutil::capture_request_span(
app.oneshot(request([10, 0, 0, 1], Some("198.51.100.9"))),
)
.await;
assert_eq!(fields.get("client_ip").as_deref(), Some("198.51.100.9"));
}
#[tokio::test]
async fn an_untrusted_peers_forwarded_header_is_ignored() {
let app = app(&[]);
let fields = acme_proxy_core::testutil::capture_request_span(
app.oneshot(request([203, 0, 113, 5], Some("198.51.100.9"))),
)
.await;
assert_eq!(fields.get("client_ip").as_deref(), Some("203.0.113.5"));
}
#[test]
fn a_refusal_becomes_a_403_naming_the_reason() {
let response = problem_for(
&Outcome::Deny("address 1.2.3.4 is not allowed".to_string()),
None,
"/newOrder",
)
.into_response();
assert_eq!(response.status(), StatusCode::FORBIDDEN);
}
#[test]
fn an_unknown_becomes_a_500() {
let response = problem_for(&Outcome::Undecided("resolver down".to_string()), None, "/x")
.into_response();
assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR);
}
}