Skip to main content

o_sfu/runtime/
request_origin.rs

1use std::{
2    convert::Infallible,
3    net::{IpAddr, SocketAddr},
4};
5
6use axum::{
7    extract::{ConnectInfo, FromRequestParts},
8    http::{HeaderMap, header, request::Parts, uri::Authority},
9};
10use ipnet::IpNet;
11
12use crate::runtime::RuntimeState;
13
14/// Proxy-aware request origin derived by the HTTP edge.
15#[derive(Debug, Clone, PartialEq, Eq)]
16pub struct RequestOrigin {
17    pub base_url: String,
18    pub remote_address: Option<IpAddr>,
19}
20
21impl FromRequestParts<RuntimeState> for RequestOrigin {
22    type Rejection = Infallible;
23
24    async fn from_request_parts(
25        parts: &mut Parts,
26        state: &RuntimeState,
27    ) -> Result<Self, Self::Rejection> {
28        let connect_info = parts
29            .extensions
30            .get::<ConnectInfo<SocketAddr>>()
31            .map(|ConnectInfo(addr)| *addr);
32        Ok(resolve_request_origin(
33            &parts.headers,
34            state.config.http.trust_proxy_headers,
35            &state.config.http.trusted_proxies,
36            state.config.http.bind_address,
37            connect_info,
38        ))
39    }
40}
41
42/// Resolves forwarded metadata only for TCP peers in `trusted_proxies` with proxy mode enabled.
43///
44/// Forwarded addresses must all parse as IP addresses. The rightmost untrusted
45/// address identifies the client, after skipping trusted proxy hops. Missing,
46/// malformed or entirely trusted chains fall back to the TCP peer. IPv4-mapped
47/// IPv6 addresses use their IPv4 identity and require an IPv4 trusted CIDR.
48/// Without a peer, forwarding is disabled.
49///
50/// The trusted edge must overwrite forwarded host and protocol with single
51/// values. Invalid or repeated values fall back to Host and HTTP respectively.
52/// Address traversal follows [NGINX's recursive trust rule](https://nginx.org/en/docs/http/ngx_http_realip_module.html#real_ip_recursive).
53#[must_use]
54pub fn resolve_request_origin(
55    headers: &HeaderMap,
56    trust_proxy_headers: bool,
57    trusted_proxies: &[IpNet],
58    fallback_bind_address: SocketAddr,
59    connect_info: Option<SocketAddr>,
60) -> RequestOrigin {
61    let peer = connect_info.map(|addr| addr.ip().to_canonical());
62    let trust_headers = trust_proxy_headers
63        && peer.is_some_and(|address| is_trusted_proxy(address, trusted_proxies));
64    let remote_address = trust_headers
65        .then(|| forwarded_client_address(headers, trusted_proxies))
66        .flatten()
67        .or(peer);
68    let scheme = trust_headers
69        .then(|| single_header(headers, "x-forwarded-proto"))
70        .flatten()
71        .filter(|scheme| matches!(*scheme, "http" | "https"))
72        .unwrap_or("http");
73    let host = trust_headers
74        .then(|| single_header(headers, "x-forwarded-host"))
75        .flatten()
76        .and_then(valid_authority)
77        .or_else(|| single_header(headers, header::HOST.as_str()).and_then(valid_authority))
78        .map_or_else(
79            || fallback_bind_address.to_string(),
80            |host| host.to_string(),
81        );
82    RequestOrigin {
83        base_url: format!("{scheme}://{host}"),
84        remote_address,
85    }
86}
87
88fn is_trusted_proxy(address: IpAddr, trusted_proxies: &[IpNet]) -> bool {
89    // Canonical IPv4 identity must not inherit trust from an IPv6-wide CIDR.
90    trusted_proxies
91        .iter()
92        .any(|network| network.contains(&address))
93}
94
95fn forwarded_client_address(headers: &HeaderMap, trusted_proxies: &[IpNet]) -> Option<IpAddr> {
96    let mut client = None;
97    // Validate the entire chain even after finding an untrusted hop. Accepting
98    // a valid suffix of malformed input would give it a different trust meaning.
99    for value in headers.get_all("x-forwarded-for") {
100        for item in value.to_str().ok()?.split(',') {
101            let address = item.trim().parse::<IpAddr>().ok()?.to_canonical();
102            if !is_trusted_proxy(address, trusted_proxies) {
103                client = Some(address);
104            }
105        }
106    }
107    client
108}
109
110fn valid_authority(host: &str) -> Option<Authority> {
111    host.parse::<Authority>()
112        .ok()
113        .filter(|host| !host.as_str().contains('@'))
114}
115
116fn single_header<'headers>(headers: &'headers HeaderMap, name: &str) -> Option<&'headers str> {
117    let mut values = headers.get_all(name).iter();
118    let value = values.next()?.to_str().ok()?.trim();
119    (values.next().is_none() && !value.is_empty() && !value.contains(',')).then_some(value)
120}
121
122#[cfg(test)]
123#[path = "TESTS/request_origin.rs"]
124mod tests;