use std::net::SocketAddr;
use axum::extract::{ConnectInfo, FromRef, FromRequestParts, Request};
use axum::middleware::Next;
use axum::response::Response;
use axum::{Extension, RequestPartsExt};
use clap::ValueEnum;
use http::StatusCode;
use http::request::Parts;
use super::AppState;
#[derive(ValueEnum, Clone, Copy, Debug)]
pub enum ClientIp {
None,
Socket,
#[clap(name = "CF-Connecting-IP")]
CfConnectingIp,
#[clap(name = "Fly-Client-IP")]
FlyClientIp,
#[clap(name = "True-Client-IP")]
TrueClientIp,
#[clap(name = "X-Real-IP")]
XRealIp,
#[clap(name = "X-Forwarded-For")]
XForwardedFor,
#[clap(name = "Forwarded")]
Forwarded,
}
impl std::fmt::Display for ClientIp {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ClientIp::None => write!(f, "None"),
ClientIp::Socket => write!(f, "Socket"),
ClientIp::CfConnectingIp => write!(f, "CF-Connecting-IP"),
ClientIp::FlyClientIp => write!(f, "Fly-Client-IP"),
ClientIp::TrueClientIp => write!(f, "True-Client-IP"),
ClientIp::XRealIp => write!(f, "X-Real-IP"),
ClientIp::XForwardedFor => write!(f, "X-Forwarded-For"),
ClientIp::Forwarded => write!(f, "Forwarded"),
}
}
}
impl ClientIp {
fn is_header(self) -> bool {
match self {
ClientIp::None => false,
ClientIp::Socket => false,
ClientIp::CfConnectingIp => true,
ClientIp::FlyClientIp => true,
ClientIp::TrueClientIp => true,
ClientIp::XRealIp => true,
ClientIp::XForwardedFor => true,
ClientIp::Forwarded => true,
}
}
}
fn parse_forwarded_for(header: &str) -> Option<String> {
let first_element = split_top_level(header, ',').next()?;
for pair in split_top_level(first_element, ';') {
let Some((name, value)) = pair.trim().split_once('=') else {
continue;
};
if name.trim().eq_ignore_ascii_case("for") {
return Some(unquote(value.trim()));
}
}
None
}
fn split_top_level(input: &str, delim: char) -> impl Iterator<Item = &str> {
let bytes = input.as_bytes();
let mut start = 0;
let mut idx = 0;
let mut in_quotes = false;
let mut escape = false;
std::iter::from_fn(move || {
while idx < bytes.len() {
let c = bytes[idx] as char;
if escape {
escape = false;
} else if in_quotes {
match c {
'\\' => escape = true,
'"' => in_quotes = false,
_ => {}
}
} else if c == '"' {
in_quotes = true;
} else if c == delim {
let segment = &input[start..idx];
idx += 1;
start = idx;
return Some(segment);
}
idx += 1;
}
if start <= bytes.len() {
let segment = &input[start..bytes.len()];
start = bytes.len() + 1;
Some(segment)
} else {
None
}
})
}
fn unquote(value: &str) -> String {
let bytes = value.as_bytes();
if bytes.len() < 2 || bytes[0] != b'"' || bytes[bytes.len() - 1] != b'"' {
return value.to_owned();
}
let inner = &value[1..value.len() - 1];
let mut out = String::with_capacity(inner.len());
let mut escape = false;
for c in inner.chars() {
if escape {
out.push(c);
escape = false;
} else if c == '\\' {
escape = true;
} else {
out.push(c);
}
}
out
}
#[derive(Clone)]
pub(super) struct ExtractClientIP(pub Option<String>);
impl<S> FromRequestParts<S> for ExtractClientIP
where
AppState: FromRef<S>,
S: Send + Sync,
{
type Rejection = (StatusCode, &'static str);
async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Self::Rejection> {
let app_state = AppState::from_ref(state);
let res = match app_state.client_ip {
ClientIp::None => ExtractClientIP(None),
ClientIp::Socket => {
match ConnectInfo::<SocketAddr>::from_request_parts(parts, state).await {
Ok(ConnectInfo(addr)) => ExtractClientIP(Some(addr.ip().to_string())),
_ => ExtractClientIP(None),
}
}
var if var.is_header() => {
if let Some(ip) = parts.headers.get(var.to_string()) {
match ip.to_str() {
Ok(s) => {
let parsed = match var {
ClientIp::Forwarded => parse_forwarded_for(s),
_ => Some(s.to_string()),
};
ExtractClientIP(parsed)
}
Err(err) => {
debug!("Invalid header value for {}: {}", var, err);
ExtractClientIP(None)
}
}
} else {
ExtractClientIP(None)
}
}
_ => {
warn!("Unexpected ClientIp variant: {:?}", app_state.client_ip);
ExtractClientIP(None)
}
};
Ok(res)
}
}
pub(super) async fn client_ip_middleware(
request: Request,
next: Next,
) -> Result<Response, StatusCode> {
let (mut parts, body) = request.into_parts();
match parts.extract::<Extension<AppState>>().await {
Ok(Extension(state)) => {
if let Ok(client_ip) =
parts.extract_with_state::<ExtractClientIP, AppState>(&state).await
{
parts.extensions.insert(client_ip);
}
}
_ => {
trace!("No AppState found, skipping client_ip_middleware");
}
}
Ok(next.run(Request::from_parts(parts, body)).await)
}
#[cfg(test)]
mod tests {
use super::parse_forwarded_for;
#[test]
fn forwarded_simple() {
assert_eq!(parse_forwarded_for("for=192.0.2.43"), Some("192.0.2.43".to_owned()));
}
#[test]
fn forwarded_takes_first_element() {
assert_eq!(
parse_forwarded_for("for=192.0.2.43, for=198.51.100.17"),
Some("192.0.2.43".to_owned())
);
}
#[test]
fn forwarded_skips_other_parameters() {
assert_eq!(
parse_forwarded_for("by=203.0.113.43;proto=http;for=192.0.2.60"),
Some("192.0.2.60".to_owned())
);
}
#[test]
fn forwarded_is_case_insensitive() {
assert_eq!(parse_forwarded_for("For=192.0.2.60"), Some("192.0.2.60".to_owned()));
assert_eq!(parse_forwarded_for("FOR=192.0.2.60"), Some("192.0.2.60".to_owned()));
}
#[test]
fn forwarded_ipv6_quoted() {
assert_eq!(
parse_forwarded_for(r#"for="[2001:db8:cafe::17]:4711""#),
Some("[2001:db8:cafe::17]:4711".to_owned())
);
}
#[test]
fn forwarded_quoted_with_escape() {
assert_eq!(parse_forwarded_for(r#"for="\"weird\"""#), Some(r#""weird""#.to_owned()));
}
#[test]
fn forwarded_quoted_value_with_semicolon() {
assert_eq!(
parse_forwarded_for(r#"for="192.0.2.43;not-a-param";by=203.0.113.43"#),
Some("192.0.2.43;not-a-param".to_owned())
);
}
#[test]
fn forwarded_quoted_value_with_comma() {
assert_eq!(
parse_forwarded_for(r#"for="192.0.2.43,still-first", for=198.51.100.17"#),
Some("192.0.2.43,still-first".to_owned())
);
}
#[test]
fn forwarded_obfuscated_identifier() {
assert_eq!(parse_forwarded_for("for=_hidden"), Some("_hidden".to_owned()));
}
#[test]
fn forwarded_unknown_identifier() {
assert_eq!(parse_forwarded_for("for=unknown"), Some("unknown".to_owned()));
}
#[test]
fn forwarded_no_for_parameter() {
assert_eq!(parse_forwarded_for("by=203.0.113.43;proto=http"), None);
}
#[test]
fn forwarded_empty() {
assert_eq!(parse_forwarded_for(""), None);
}
#[test]
fn forwarded_skips_empty_pairs() {
assert_eq!(
parse_forwarded_for(";by=203.0.113.43;;for=192.0.2.60;"),
Some("192.0.2.60".to_owned())
);
}
}