use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4};
use std::sync::Arc;
use axum::body::Body;
use axum::extract::{ConnectInfo, Request, State};
use axum::http::header::{CONTENT_TYPE, LOCATION, WWW_AUTHENTICATE};
use axum::http::{HeaderMap, StatusCode};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use axum::routing::{get, on, MethodFilter, MethodRouter};
use axum::{Json, Router};
use crate::hooks::{AuthDecider, Decision, ExtraRoute, OidcBackend, RequestParts, RouteRequest};
use crate::ProxyState;
const MAX_EXTRA_ROUTE_BODY: usize = 16 * 1024 * 1024;
fn unknown_peer() -> SocketAddr {
SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0))
}
fn peer_of(req: &Request) -> SocketAddr {
req.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|ci| ci.0)
.unwrap_or_else(unknown_peer)
}
pub(crate) async fn auth_decider_gate(
State(decider): State<Arc<dyn AuthDecider>>,
mut request: Request,
next: Next,
) -> Response {
let peer = peer_of(&request);
let decision = {
let uri = request.uri();
let parts = RequestParts {
method: request.method(),
path: uri.path(),
query: uri.query(),
headers: request.headers(),
peer,
};
decider.decide(&parts).await
};
match decision {
Decision::Allow { inject_headers } => {
let dst = request.headers_mut();
strip_then_insert(dst, &inject_headers);
next.run(request).await
}
Decision::Deny { status, body } => deny_response(status, body),
Decision::Redirect { location } => redirect_response(StatusCode::FOUND, &location),
}
}
pub(crate) async fn verify_via_decider(
decider: Arc<dyn AuthDecider>,
request: Request,
) -> Response {
let peer = peer_of(&request);
let headers = request.headers().clone();
let method = original_method(&headers).unwrap_or_else(|| request.method().clone());
let (path, query) = original_target(&headers).unwrap_or_else(|| {
let uri = request.uri();
(uri.path().to_string(), uri.query().map(str::to_string))
});
let decision = {
let parts = RequestParts {
method: &method,
path: &path,
query: query.as_deref(),
headers: &headers,
peer,
};
decider.decide(&parts).await
};
match decision {
Decision::Allow { inject_headers } => (StatusCode::OK, inject_headers).into_response(),
Decision::Deny { status, body } => deny_response(status, body),
Decision::Redirect { location } => redirect_response(StatusCode::UNAUTHORIZED, &location),
}
}
pub(crate) fn oidc_backend_routes(backend: Arc<dyn OidcBackend>) -> Router<ProxyState> {
let mut router = Router::new();
for doc in backend.metadata_documents() {
let json = doc.json;
router = router.route(
&doc.path,
get(move || {
let json = json.clone();
async move { Json(json) }
}),
);
}
let jwks = backend.jwks();
let jwks_body =
serde_json::to_string(&jwks.json).unwrap_or_else(|_| "{\"keys\":[]}".to_string());
router = router.route(
&jwks.path,
get(move || {
let body = jwks_body.clone();
async move { ([(CONTENT_TYPE, "application/jwk-set+json")], body) }
}),
);
let userinfo_path = backend.userinfo_path();
let userinfo_backend = backend.clone();
router.route(
&userinfo_path,
get(move |headers: HeaderMap| {
let backend = userinfo_backend.clone();
async move {
let Some(token) = bearer_token(&headers) else {
return unauthorized_with_challenge(
bytes::Bytes::from_static(
br#"{"error":"invalid_request","message":"missing bearer token"}"#,
),
"Bearer",
);
};
match backend.userinfo(&token).await {
Some(claims) => Json(claims).into_response(),
None => unauthorized_with_challenge(
bytes::Bytes::from_static(
br#"{"error":"invalid_token","message":"invalid or expired token"}"#,
),
r#"Bearer error="invalid_token""#,
),
}
}
}),
)
}
pub(crate) fn extra_routes_router(routes: &[ExtraRoute]) -> Router<ProxyState> {
use std::collections::HashMap;
let mut by_path: HashMap<String, MethodRouter<ProxyState>> = HashMap::new();
for route in routes {
let Ok(filter) = MethodFilter::try_from(route.method.clone()) else {
tracing::warn!(
method = %route.method,
path = %route.path,
"skipping extra route: unsupported HTTP method"
);
continue;
};
let handler = route.handler.clone();
let service = on(filter, move |request: Request| {
let handler = handler.clone();
async move {
let peer = peer_of(&request);
let (parts, body) = request.into_parts();
let body = match axum::body::to_bytes(body, MAX_EXTRA_ROUTE_BODY).await {
Ok(bytes) => bytes,
Err(_) => {
return deny_response(
StatusCode::PAYLOAD_TOO_LARGE,
bytes::Bytes::from_static(
br#"{"error":"payload_too_large","message":"request body exceeded limit or could not be read"}"#,
),
)
}
};
let resp = handler
.handle(RouteRequest {
method: parts.method,
uri: parts.uri,
headers: parts.headers,
body,
peer,
})
.await;
let mut response = Response::new(Body::from(resp.body));
*response.status_mut() = resp.status;
*response.headers_mut() = resp.headers;
response
}
});
match by_path.remove(&route.path) {
Some(existing) => {
by_path.insert(route.path.clone(), existing.merge(service));
}
None => {
by_path.insert(route.path.clone(), service);
}
}
}
let mut router = Router::new();
for (path, method_router) in by_path {
router = router.route(&path, method_router);
}
router
}
fn strip_then_insert(dst: &mut HeaderMap, inject: &HeaderMap) {
for name in inject.keys() {
while dst.remove(name).is_some() {}
}
for (name, value) in inject {
dst.append(name.clone(), value.clone());
}
}
fn deny_response(status: StatusCode, body: bytes::Bytes) -> Response {
(status, [(CONTENT_TYPE, "application/json")], body).into_response()
}
fn unauthorized_with_challenge(body: bytes::Bytes, challenge: &'static str) -> Response {
(
StatusCode::UNAUTHORIZED,
[
(CONTENT_TYPE, "application/json"),
(WWW_AUTHENTICATE, challenge),
],
body,
)
.into_response()
}
fn redirect_response(status: StatusCode, location: &str) -> Response {
let mut response = status.into_response();
if let Ok(value) = location.parse() {
response.headers_mut().insert(LOCATION, value);
}
response
}
fn bearer_token(headers: &HeaderMap) -> Option<String> {
let value = headers.get("authorization")?.to_str().ok()?;
let (scheme, rest) = value.split_once(' ')?;
if !scheme.eq_ignore_ascii_case("bearer") {
return None;
}
let token = rest.trim();
(!token.is_empty()).then(|| token.to_string())
}
fn original_method(headers: &HeaderMap) -> Option<axum::http::Method> {
let raw = forwarded(headers, &["x-forwarded-method", "x-original-method"])?;
axum::http::Method::from_bytes(raw.to_ascii_uppercase().as_bytes()).ok()
}
fn original_target(headers: &HeaderMap) -> Option<(String, Option<String>)> {
let raw = forwarded(headers, &["x-forwarded-uri", "x-original-uri"])?;
Some(match raw.split_once('?') {
Some((path, query)) => (path.to_string(), Some(query.to_string())),
None => (raw, None),
})
}
fn forwarded(headers: &HeaderMap, names: &[&str]) -> Option<String> {
names
.iter()
.filter_map(|n| headers.get(*n).and_then(|v| v.to_str().ok()))
.find(|v| !v.is_empty())
.map(str::to_string)
}
#[cfg(test)]
mod tests;