use std::{
io,
mem::{size_of, zeroed},
net::{Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6},
os::fd::{AsRawFd, RawFd},
ptr,
};
use crate::{address::SocketAddress, client::ConnectorTarget};
use rama_core::{
Layer, Service,
error::{BoxError, ErrorContext as _, ErrorExt as _},
extensions::ExtensionsRef,
};
#[derive(Debug, Clone, Default)]
pub struct ConnectorTargetFromGetSocketnameLayer;
impl ConnectorTargetFromGetSocketnameLayer {
#[inline(always)]
pub fn new() -> Self {
Self
}
}
impl<S> Layer<S> for ConnectorTargetFromGetSocketnameLayer {
type Service = ConnectorTargetFromGetSocketname<S>;
fn layer(&self, inner: S) -> Self::Service {
ConnectorTargetFromGetSocketname { inner }
}
}
#[derive(Debug, Clone)]
pub struct ConnectorTargetFromGetSocketname<S> {
inner: S,
}
impl<S, Input> Service<Input> for ConnectorTargetFromGetSocketname<S>
where
S: Service<Input, Error: Into<BoxError>>,
Input: AsRawFd + ExtensionsRef + Send + 'static,
{
type Output = S::Output;
type Error = BoxError;
async fn serve(&self, input: Input) -> Result<Self::Output, Self::Error> {
let proxy_target = connector_target_from_input(input.as_raw_fd())
.context("get (proxy) connector target from input stream")?;
input
.extensions()
.insert(ConnectorTarget(proxy_target.into()));
self.inner.serve(input).await.context("inner serve tcp")
}
}
fn connector_target_from_input(fd: RawFd) -> Result<SocketAddress, BoxError> {
let mut storage: libc::sockaddr_storage = unsafe { zeroed() };
let mut len = size_of::<libc::sockaddr_storage>() as libc::socklen_t;
let rc = unsafe {
libc::getsockname(fd, &mut storage as *mut _ as *mut libc::sockaddr, &mut len)
};
if rc != 0 {
return Err(io::Error::last_os_error().context("getsockname"));
}
sockaddr_storage_to_socket_addr(&storage, len).context("socketaddr storage to SocketAddress")
}
fn sockaddr_storage_to_socket_addr(
storage: &libc::sockaddr_storage,
len: libc::socklen_t,
) -> io::Result<SocketAddress> {
match storage.ss_family as libc::c_int {
libc::AF_INET => parse_sockaddr_in(storage, len),
libc::AF_INET6 => parse_sockaddr_in6(storage, len),
family => Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("unsupported address family: {family}"),
)),
}
}
fn parse_sockaddr_in(
storage: &libc::sockaddr_storage,
len: libc::socklen_t,
) -> io::Result<SocketAddress> {
ensure_sockaddr_len::<libc::sockaddr_in>(len, "sockaddr_in")?;
let addr: libc::sockaddr_in = unsafe {
ptr::read_unaligned((storage as *const libc::sockaddr_storage).cast())
};
let ip = Ipv4Addr::from(u32::from_be(addr.sin_addr.s_addr));
let port = u16::from_be(addr.sin_port);
Ok(SocketAddr::V4(SocketAddrV4::new(ip, port)).into())
}
fn parse_sockaddr_in6(
storage: &libc::sockaddr_storage,
len: libc::socklen_t,
) -> io::Result<SocketAddress> {
ensure_sockaddr_len::<libc::sockaddr_in6>(len, "sockaddr_in6")?;
let addr: libc::sockaddr_in6 = unsafe {
ptr::read_unaligned((storage as *const libc::sockaddr_storage).cast())
};
let ip = Ipv6Addr::from(addr.sin6_addr.s6_addr);
let port = u16::from_be(addr.sin6_port);
Ok(SocketAddr::V6(SocketAddrV6::new(
ip,
port,
addr.sin6_flowinfo,
addr.sin6_scope_id,
))
.into())
}
fn ensure_sockaddr_len<T>(len: libc::socklen_t, kind: &'static str) -> io::Result<()> {
if len < size_of::<T>() as libc::socklen_t {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("short {kind}"),
));
}
Ok(())
}
#[cfg(test)]
mod tests {
use std::{mem::zeroed, net::IpAddr};
use super::*;
#[test]
fn sockaddr_storage_to_socket_addr_ipv4() {
let ip = Ipv4Addr::new(127, 0, 0, 1);
let port = 15001u16;
let raw = libc::sockaddr_in {
sin_family: libc::AF_INET as _,
sin_port: port.to_be(),
sin_addr: libc::in_addr {
s_addr: u32::from(ip).to_be(),
},
sin_zero: [0; 8],
};
let storage = sockaddr_storage_from(raw);
let addr =
sockaddr_storage_to_socket_addr(&storage, size_of::<libc::sockaddr_in>() as _).unwrap();
assert_eq!(addr.ip_addr, IpAddr::V4(ip));
assert_eq!(addr.port, port);
}
#[test]
fn sockaddr_storage_to_socket_addr_ipv6() {
let ip = Ipv6Addr::LOCALHOST;
let port = 15001u16;
let flowinfo = 42;
let scope_id = 7;
let raw = libc::sockaddr_in6 {
sin6_family: libc::AF_INET6 as _,
sin6_port: port.to_be(),
sin6_flowinfo: flowinfo,
sin6_addr: libc::in6_addr {
s6_addr: ip.octets(),
},
sin6_scope_id: scope_id,
};
let storage = sockaddr_storage_from(raw);
let addr = sockaddr_storage_to_socket_addr(&storage, size_of::<libc::sockaddr_in6>() as _)
.unwrap();
assert_eq!(addr.ip_addr, IpAddr::V6(ip));
assert_eq!(addr.port, port);
}
#[test]
fn sockaddr_storage_to_socket_addr_rejects_short_sockaddr() {
let storage = sockaddr_storage_from(libc::sockaddr_in {
sin_family: libc::AF_INET as _,
sin_port: 0,
sin_addr: libc::in_addr { s_addr: 0 },
sin_zero: [0; 8],
});
let err = sockaddr_storage_to_socket_addr(&storage, 1).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
assert_eq!(err.to_string(), "short sockaddr_in");
}
#[test]
fn sockaddr_storage_to_socket_addr_rejects_unsupported_family() {
let mut storage: libc::sockaddr_storage = unsafe { zeroed() };
storage.ss_family = libc::AF_UNIX as _;
let err =
sockaddr_storage_to_socket_addr(&storage, size_of::<libc::sockaddr_storage>() as _)
.unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
assert_eq!(err.to_string(), "unsupported address family: 1");
}
fn sockaddr_storage_from<T>(raw: T) -> libc::sockaddr_storage {
let mut storage: libc::sockaddr_storage = unsafe { zeroed() };
unsafe {
ptr::write(
(&mut storage as *mut libc::sockaddr_storage).cast::<T>(),
raw,
);
}
storage
}
}