use core::net::IpAddr;
use crate::forwarded::Forwarded;
use rama_core::extensions::ExtensionsRef;
#[cfg(feature = "std")]
use crate::stream::SocketInfo;
pub fn client_ip(ext: &impl ExtensionsRef) -> Option<IpAddr> {
let extensions = ext.extensions();
let forwarded = extensions
.get_ref::<Forwarded>()
.and_then(Forwarded::client_ip);
#[cfg(feature = "std")]
{
forwarded.or_else(|| {
extensions
.get_ref::<SocketInfo>()
.map(|info| info.peer_addr().ip_addr)
})
}
#[cfg(not(feature = "std"))]
{
forwarded
}
}
pub trait ClientIp {
fn client_ip(&self) -> Option<IpAddr>;
}
#[cfg(test)]
mod tests {
use super::*;
use rama_core::extensions::Extensions;
#[cfg(feature = "std")]
use crate::address::SocketAddress;
#[cfg(feature = "std")]
use crate::forwarded::{ForwardedElement, NodeId};
#[cfg(feature = "std")]
use crate::stream::SocketInfo;
#[cfg(feature = "std")]
fn socket_info(ip: &str) -> SocketInfo {
SocketInfo::new(None, SocketAddress::new(ip.parse().unwrap(), 0))
}
#[test]
fn none_when_no_extensions() {
let ext = Extensions::new();
assert_eq!(client_ip(&ext), None);
}
#[cfg(feature = "std")]
fn forwarded_for(node: &str) -> Forwarded {
Forwarded::new(ForwardedElement::new_forwarded_for(
NodeId::try_from(node).unwrap(),
))
}
#[cfg(feature = "std")]
#[test]
fn falls_back_to_socket_info_peer() {
let ext = Extensions::new();
ext.insert(socket_info("203.0.113.5"));
assert_eq!(client_ip(&ext), Some("203.0.113.5".parse().unwrap()));
}
#[cfg(feature = "std")]
#[test]
fn prefers_forwarded_over_socket_info() {
let ext = Extensions::new();
ext.insert(socket_info("203.0.113.5"));
ext.insert(forwarded_for("192.0.2.43"));
assert_eq!(client_ip(&ext), Some("192.0.2.43".parse().unwrap()));
}
#[cfg(feature = "std")]
#[test]
fn forwarded_without_client_ip_falls_back_to_socket_info() {
let ext = Extensions::new();
ext.insert(socket_info("203.0.113.5"));
ext.insert(forwarded_for("_hidden"));
assert_eq!(client_ip(&ext), Some("203.0.113.5".parse().unwrap()));
}
}