1use libc::{
18 __errno_location,
19 __u32,
20 AF_INET,
21 AF_INET6,
22 EINPROGRESS,
23 SOCK_STREAM,
24 SOL_TCP,
25 TCP_QUEUE_SEQ,
27 TCP_REPAIR,
28 TCP_REPAIR_QUEUE,
29 c_int,
30 c_void,
31 close as libc_close,
32 connect as libc_connect,
33 in_addr,
34 in6_addr,
35 setsockopt,
36 sockaddr,
37 sockaddr_in,
38 sockaddr_in6,
39 socket,
40};
41
42static TCP_SEND_QUEUE: __u32 = 2;
45
46use std::io::{Error, ErrorKind, Result};
47use std::mem::size_of;
48use std::net::SocketAddr;
49use std::net::ToSocketAddrs;
50use std::os::unix::io::FromRawFd;
51
52use log::debug;
53
54#[cfg(feature = "async")]
55use async_io;
56#[cfg(feature = "async")]
57use async_net;
58
59#[derive(Clone, Copy, PartialEq)]
60pub enum Family {
61 V4,
62 V6,
63}
64
65fn family_matches(socket_addr: &SocketAddr, family: Option<Family>) -> bool {
66 if let Some(f) = family {
67 if f == Family::V4 && !socket_addr.is_ipv4() {
68 return false;
69 }
70 if f == Family::V6 && !socket_addr.is_ipv6() {
71 return false;
72 }
73 }
74 true
75}
76
77pub fn connect<A: ToSocketAddrs>(
78 sequence_no: u32,
79 addr: A,
80 force_family: Option<Family>,
81) -> Result<std::net::TcpStream> {
82 unsafe {
83 let socket_addrs = addr.to_socket_addrs()?;
85
86 let mut maybe_err = None;
88 for socket_addr in socket_addrs {
89 if !family_matches(&socket_addr, force_family) {
90 debug!("skipping {}, not of requested family", socket_addr);
91 continue;
92 }
93 debug!("Trying to connect to {}", socket_addr);
94
95 let sock = create_socket(family_of(&socket_addr), sequence_no)?;
96
97 match connect_socket(sock, socket_addr, true) {
98 Ok(()) => {
99 debug!("Connected to {}", socket_addr);
100 return Ok(std::net::TcpStream::from_raw_fd(sock));
102 }
103 Err(e) => maybe_err = Some(e),
104 }
105
106 libc_close(sock);
107 }
108
109 if let Some(e) = maybe_err {
110 return Err(e);
111 }
112 }
113 Err(Error::new(
114 ErrorKind::AddrNotAvailable,
115 "No address entries for hostname",
116 ))
117}
118
119#[cfg(feature = "async")]
120pub async fn connect_async<A>(
121 sequence_no: u32,
122 socket_addr: A,
123 force_family: Option<Family>,
124) -> Result<async_net::TcpStream>
125where
126 A: async_net::AsyncToSocketAddrs,
127{
128 let socket_addrs = async_net::resolve(socket_addr).await?;
129
130 unsafe {
131 let mut maybe_err = None;
133 for socket_addr in socket_addrs {
134 if !family_matches(&socket_addr, force_family) {
135 debug!("skipping {}, not of requested family", socket_addr);
136 continue;
137 }
138 debug!("Trying to connect to {}", socket_addr);
139
140 let sock = create_socket(family_of(&socket_addr), sequence_no)?;
141
142 match connect_socket(sock, socket_addr, true) {
143 Ok(()) => {
144 let stream = match async_io::Async::new(std::net::TcpStream::from_raw_fd(sock))
145 {
146 Ok(s) => s,
147 Err(e) => {
148 maybe_err = Some(e);
149 continue;
150 }
151 };
152 match stream.writable().await {
153 Ok(_) => match stream.get_ref().take_error()? {
154 None => {
155 debug!("Connected to {}", socket_addr);
156 return Ok(stream.into());
157 }
158 Some(e) => {
159 maybe_err = Some(e);
160 continue;
161 }
162 },
163 Err(e) => maybe_err = Some(e),
164 }
165 }
166 Err(e) => maybe_err = Some(e),
167 }
168
169 libc_close(sock);
170 }
171
172 if let Some(e) = maybe_err {
173 return Err(e);
174 }
175 }
176
177 Err(Error::new(
178 ErrorKind::AddrNotAvailable,
179 "No address entries for hostname",
180 ))
181}
182
183unsafe fn connect_socket(
184 sock: c_int,
185 socket_addr: std::net::SocketAddr,
186 blocking: bool,
187) -> Result<()> {
188 unsafe {
189 match socket_addr {
190 SocketAddr::V4(v4addr) => {
191 let octets = v4addr.ip().octets();
192 let u32_addr: u32 = (octets[0] as u32)
193 | (octets[1] as u32) << 8
194 | (octets[2] as u32) << 16
195 | (octets[3] as u32) << 24;
196 let saddr = sockaddr_in {
197 sin_family: AF_INET as u16,
198 sin_port: v4addr.port().to_be(),
199 sin_addr: in_addr { s_addr: u32_addr },
200 sin_zero: [0; 8],
201 };
202 let result = libc_connect(
203 sock,
204 &saddr as *const sockaddr_in as *const sockaddr,
205 size_of::<sockaddr_in>() as u32,
206 );
207 if result < 0 && (blocking || (*__errno_location()) != EINPROGRESS) {
208 return Err(Error::last_os_error());
209 }
210 }
211 SocketAddr::V6(v6addr) => {
212 let saddr = sockaddr_in6 {
213 sin6_family: AF_INET6 as u16,
214 sin6_port: v6addr.port().to_be(),
215 sin6_flowinfo: 0,
216 sin6_addr: in6_addr {
217 s6_addr: v6addr.ip().octets(),
218 },
219 sin6_scope_id: 0,
220 };
221 let result = libc_connect(
222 sock,
223 &saddr as *const sockaddr_in6 as *const sockaddr,
224 size_of::<sockaddr_in6>() as u32,
225 );
226 if result < 0 && (blocking || (*__errno_location()) != EINPROGRESS) {
227 return Err(Error::last_os_error());
228 }
229 }
230 }
231
232 Ok(())
233 }
234}
235
236unsafe fn sso_tcp_wrapper(sock: c_int, cmd: c_int, data: u32) -> Result<()> {
237 unsafe {
238 let dataptr = &data as *const __u32 as *const c_void;
239 if setsockopt(sock, SOL_TCP, cmd, dataptr, 4) < 0 {
240 return Err(Error::last_os_error());
241 }
242 Ok(())
243 }
244}
245
246unsafe fn create_socket(family: c_int, sequence_no: u32) -> Result<c_int> {
247 unsafe {
248 let sock: c_int = socket(family, SOCK_STREAM, 0);
250 if sock < 0 {
251 return Err(Error::last_os_error());
252 }
253
254 sso_tcp_wrapper(sock, TCP_REPAIR, 1)?;
256 sso_tcp_wrapper(sock, TCP_REPAIR_QUEUE, TCP_SEND_QUEUE)?;
258 sso_tcp_wrapper(sock, TCP_QUEUE_SEQ, sequence_no)?;
260 sso_tcp_wrapper(sock, TCP_REPAIR, 0)?;
262
263 Ok(sock)
264 }
265}
266
267fn family_of(socket_addr: &std::net::SocketAddr) -> c_int {
268 match socket_addr {
269 SocketAddr::V4(_) => AF_INET,
270 SocketAddr::V6(_) => AF_INET6,
271 }
272}