use axum::http::{HeaderMap, HeaderName, header};
pub(super) enum ForwardedHeaderMode {
Normal,
Upload,
}
pub(super) trait HeaderMapProxyExt {
fn forwarded_for_guest(&self, mode: ForwardedHeaderMode) -> HeaderMap;
fn wants_upgrade(&self) -> bool;
}
impl HeaderMapProxyExt for HeaderMap {
fn forwarded_for_guest(&self, mode: ForwardedHeaderMode) -> HeaderMap {
let strip_expect = matches!(mode, ForwardedHeaderMode::Upload);
let connection_tokens = self
.get_all(header::CONNECTION)
.iter()
.filter_map(|value| value.to_str().ok())
.flat_map(|value| value.split(','))
.filter_map(|token| token.trim().parse::<HeaderName>().ok())
.collect::<Vec<_>>();
let mut forwarded = Self::new();
for (name, value) in self {
if name == header::HOST
|| name == header::CONNECTION
|| (strip_expect && name == header::EXPECT)
|| is_hop_by_hop_header(name)
|| connection_tokens.iter().any(|token| token == name)
{
continue;
}
forwarded.append(name.clone(), value.clone());
}
forwarded
}
fn wants_upgrade(&self) -> bool {
self.get(header::UPGRADE).is_some()
|| self
.get_all(header::CONNECTION)
.iter()
.filter_map(|value| value.to_str().ok())
.any(|value| {
value
.split(',')
.any(|token| token.trim().eq_ignore_ascii_case("upgrade"))
})
}
}
fn is_hop_by_hop_header(name: &HeaderName) -> bool {
matches!(
name.as_str(),
"proxy-connection"
| "keep-alive"
| "proxy-authenticate"
| "proxy-authorization"
| "te"
| "trailer"
| "transfer-encoding"
| "upgrade"
)
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::HeaderValue;
#[test]
fn detects_upgrade_across_all_connection_headers() {
let mut headers = HeaderMap::new();
headers.append(header::CONNECTION, HeaderValue::from_static("keep-alive"));
headers.append(header::CONNECTION, HeaderValue::from_static("Upgrade"));
assert!(headers.wants_upgrade());
}
#[test]
fn detects_upgrade_token_in_comma_separated_connection_header() {
let mut headers = HeaderMap::new();
headers.insert(
header::CONNECTION,
HeaderValue::from_static("keep-alive, upgrade"),
);
assert!(headers.wants_upgrade());
}
}