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)]
18pub struct ConnectorTargetFromGetSocketnameLayer;
20
21impl ConnectorTargetFromGetSocketnameLayer {
22 #[inline(always)]
23 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)]
38pub 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 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 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 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 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 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 let mut storage: libc::sockaddr_storage = unsafe { zeroed() };
235 unsafe {
236 ptr::write(
240 (&mut storage as *mut libc::sockaddr_storage).cast::<T>(),
241 raw,
242 );
243 }
244 storage
245 }
246}