Skip to main content

rama_net/socket/linux/
tproxy.rs

1use std::{
2    io,
3    mem::{size_of, zeroed},
4    net::{Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6},
5    os::fd::{AsRawFd, RawFd},
6    ptr,
7};
8
9use crate::{address::SocketAddress, client::ConnectorTarget};
10
11use rama_core::{
12    Layer, Service,
13    error::{BoxError, ErrorContext as _, ErrorExt as _},
14    extensions::ExtensionsRef,
15};
16
17#[derive(Debug, Clone, Default)]
18/// Layer to create [`ConnectorTargetFromGetSocketname`] middleware.
19pub struct ConnectorTargetFromGetSocketnameLayer;
20
21impl ConnectorTargetFromGetSocketnameLayer {
22    #[inline(always)]
23    /// Create a new [`ConnectorTargetFromGetSocketnameLayer`]
24    pub fn new() -> Self {
25        Self
26    }
27}
28
29impl<S> Layer<S> for ConnectorTargetFromGetSocketnameLayer {
30    type Service = ConnectorTargetFromGetSocketname<S>;
31
32    fn layer(&self, inner: S) -> Self::Service {
33        ConnectorTargetFromGetSocketname { inner }
34    }
35}
36
37#[derive(Debug, Clone)]
38/// Middleware that can be used by Linux transparent proxies,
39/// to insert the [`ConnectorTarget`] based on the address inserted
40/// by the OS in the "socketname" of the underlying OS socket.
41///
42/// Created using [`ConnectorTargetFromGetSocketnameLayer`].
43pub struct ConnectorTargetFromGetSocketname<S> {
44    inner: S,
45}
46
47impl<S, Input> Service<Input> for ConnectorTargetFromGetSocketname<S>
48where
49    S: Service<Input, Error: Into<BoxError>>,
50    Input: AsRawFd + ExtensionsRef + Send + 'static,
51{
52    type Output = S::Output;
53    type Error = BoxError;
54
55    async fn serve(&self, input: Input) -> Result<Self::Output, Self::Error> {
56        let proxy_target = connector_target_from_input(input.as_raw_fd())
57            .context("get (proxy) connector target from input stream")?;
58        input
59            .extensions()
60            .insert(ConnectorTarget(proxy_target.into()));
61        self.inner.serve(input).await.context("inner serve tcp")
62    }
63}
64
65fn connector_target_from_input(fd: RawFd) -> Result<SocketAddress, BoxError> {
66    // SAFETY: `sockaddr_storage` is a plain old data buffer used as out-parameter
67    // storage for `getsockname`, so zero-initializing it is valid.
68    let mut storage: libc::sockaddr_storage = unsafe { zeroed() };
69    let mut len = size_of::<libc::sockaddr_storage>() as libc::socklen_t;
70
71    let rc = unsafe {
72        // SAFETY: `fd` comes from `AsRawFd`; `storage` points to a writable
73        // `sockaddr_storage` buffer of `len` bytes; and `len` is a valid mutable
74        // pointer for the kernel to update with the number of bytes written.
75        libc::getsockname(fd, &mut storage as *mut _ as *mut libc::sockaddr, &mut len)
76    };
77
78    if rc != 0 {
79        return Err(io::Error::last_os_error().context("getsockname"));
80    }
81
82    sockaddr_storage_to_socket_addr(&storage, len).context("socketaddr storage to SocketAddress")
83}
84
85fn sockaddr_storage_to_socket_addr(
86    storage: &libc::sockaddr_storage,
87    len: libc::socklen_t,
88) -> io::Result<SocketAddress> {
89    match storage.ss_family as libc::c_int {
90        libc::AF_INET => parse_sockaddr_in(storage, len),
91        libc::AF_INET6 => parse_sockaddr_in6(storage, len),
92        family => Err(io::Error::new(
93            io::ErrorKind::InvalidData,
94            format!("unsupported address family: {family}"),
95        )),
96    }
97}
98
99fn parse_sockaddr_in(
100    storage: &libc::sockaddr_storage,
101    len: libc::socklen_t,
102) -> io::Result<SocketAddress> {
103    ensure_sockaddr_len::<libc::sockaddr_in>(len, "sockaddr_in")?;
104
105    let addr: libc::sockaddr_in = unsafe {
106        // SAFETY: the family is `AF_INET` and `ensure_sockaddr_len` verified that at
107        // least a full `sockaddr_in` was written. We use `read_unaligned` because
108        // `sockaddr_storage` does not guarantee alignment for `sockaddr_in`.
109        ptr::read_unaligned((storage as *const libc::sockaddr_storage).cast())
110    };
111    let ip = Ipv4Addr::from(u32::from_be(addr.sin_addr.s_addr));
112    let port = u16::from_be(addr.sin_port);
113
114    Ok(SocketAddr::V4(SocketAddrV4::new(ip, port)).into())
115}
116
117fn parse_sockaddr_in6(
118    storage: &libc::sockaddr_storage,
119    len: libc::socklen_t,
120) -> io::Result<SocketAddress> {
121    ensure_sockaddr_len::<libc::sockaddr_in6>(len, "sockaddr_in6")?;
122
123    let addr: libc::sockaddr_in6 = unsafe {
124        // SAFETY: the family is `AF_INET6` and `ensure_sockaddr_len` verified that at
125        // least a full `sockaddr_in6` was written. We use `read_unaligned` because
126        // `sockaddr_storage` does not guarantee alignment for `sockaddr_in6`.
127        ptr::read_unaligned((storage as *const libc::sockaddr_storage).cast())
128    };
129    let ip = Ipv6Addr::from(addr.sin6_addr.s6_addr);
130    let port = u16::from_be(addr.sin6_port);
131
132    Ok(SocketAddr::V6(SocketAddrV6::new(
133        ip,
134        port,
135        addr.sin6_flowinfo,
136        addr.sin6_scope_id,
137    ))
138    .into())
139}
140
141fn ensure_sockaddr_len<T>(len: libc::socklen_t, kind: &'static str) -> io::Result<()> {
142    if len < size_of::<T>() as libc::socklen_t {
143        return Err(io::Error::new(
144            io::ErrorKind::InvalidData,
145            format!("short {kind}"),
146        ));
147    }
148
149    Ok(())
150}
151
152#[cfg(test)]
153mod tests {
154    use std::{mem::zeroed, net::IpAddr};
155
156    use super::*;
157
158    #[test]
159    fn sockaddr_storage_to_socket_addr_ipv4() {
160        let ip = Ipv4Addr::new(127, 0, 0, 1);
161        let port = 15001u16;
162        let raw = libc::sockaddr_in {
163            sin_family: libc::AF_INET as _,
164            sin_port: port.to_be(),
165            sin_addr: libc::in_addr {
166                s_addr: u32::from(ip).to_be(),
167            },
168            sin_zero: [0; 8],
169        };
170
171        let storage = sockaddr_storage_from(raw);
172        let addr =
173            sockaddr_storage_to_socket_addr(&storage, size_of::<libc::sockaddr_in>() as _).unwrap();
174
175        assert_eq!(addr.ip_addr, IpAddr::V4(ip));
176        assert_eq!(addr.port, port);
177    }
178
179    #[test]
180    fn sockaddr_storage_to_socket_addr_ipv6() {
181        let ip = Ipv6Addr::LOCALHOST;
182        let port = 15001u16;
183        let flowinfo = 42;
184        let scope_id = 7;
185        let raw = libc::sockaddr_in6 {
186            sin6_family: libc::AF_INET6 as _,
187            sin6_port: port.to_be(),
188            sin6_flowinfo: flowinfo,
189            sin6_addr: libc::in6_addr {
190                s6_addr: ip.octets(),
191            },
192            sin6_scope_id: scope_id,
193        };
194
195        let storage = sockaddr_storage_from(raw);
196        let addr = sockaddr_storage_to_socket_addr(&storage, size_of::<libc::sockaddr_in6>() as _)
197            .unwrap();
198
199        assert_eq!(addr.ip_addr, IpAddr::V6(ip));
200        assert_eq!(addr.port, port);
201    }
202
203    #[test]
204    fn sockaddr_storage_to_socket_addr_rejects_short_sockaddr() {
205        let storage = sockaddr_storage_from(libc::sockaddr_in {
206            sin_family: libc::AF_INET as _,
207            sin_port: 0,
208            sin_addr: libc::in_addr { s_addr: 0 },
209            sin_zero: [0; 8],
210        });
211
212        let err = sockaddr_storage_to_socket_addr(&storage, 1).unwrap_err();
213
214        assert_eq!(err.kind(), io::ErrorKind::InvalidData);
215        assert_eq!(err.to_string(), "short sockaddr_in");
216    }
217
218    #[test]
219    fn sockaddr_storage_to_socket_addr_rejects_unsupported_family() {
220        // SAFETY: `sockaddr_storage` is POD and zero-initialization is valid for a test buffer.
221        let mut storage: libc::sockaddr_storage = unsafe { zeroed() };
222        storage.ss_family = libc::AF_UNIX as _;
223
224        let err =
225            sockaddr_storage_to_socket_addr(&storage, size_of::<libc::sockaddr_storage>() as _)
226                .unwrap_err();
227
228        assert_eq!(err.kind(), io::ErrorKind::InvalidData);
229        assert_eq!(err.to_string(), "unsupported address family: 1");
230    }
231
232    fn sockaddr_storage_from<T>(raw: T) -> libc::sockaddr_storage {
233        // SAFETY: `sockaddr_storage` is POD and zero-initialization is valid for a test buffer.
234        let mut storage: libc::sockaddr_storage = unsafe { zeroed() };
235        unsafe {
236            // SAFETY: the destination points to stack-allocated storage large enough for
237            // the concrete sockaddr value used in the test, and we only read it back as
238            // that same concrete type.
239            ptr::write(
240                (&mut storage as *mut libc::sockaddr_storage).cast::<T>(),
241                raw,
242            );
243        }
244        storage
245    }
246}