#![allow(non_camel_case_types)]
use libc::socklen_t;
#[cfg(target_os = "linux")]
use libc::{c_int, c_void};
use pingora_error::{Error, ErrorType::*, OrErr, Result};
use std::io::{self, ErrorKind};
use std::mem;
use std::net::SocketAddr;
use std::os::unix::io::{AsRawFd, RawFd};
use std::time::Duration;
use tokio::net::{TcpSocket, TcpStream, UnixStream};
#[repr(C)]
#[derive(Copy, Clone, Debug)]
pub struct TCP_INFO {
tcpi_state: u8,
tcpi_ca_state: u8,
tcpi_retransmits: u8,
tcpi_probes: u8,
tcpi_backoff: u8,
tcpi_options: u8,
tcpi_snd_wscale_4_rcv_wscale_4: u8,
tcpi_delivery_rate_app_limited: u8,
tcpi_rto: u32,
tcpi_ato: u32,
tcpi_snd_mss: u32,
tcpi_rcv_mss: u32,
tcpi_unacked: u32,
tcpi_sacked: u32,
tcpi_lost: u32,
tcpi_retrans: u32,
tcpi_fackets: u32,
tcpi_last_data_sent: u32,
tcpi_last_ack_sent: u32,
tcpi_last_data_recv: u32,
tcpi_last_ack_recv: u32,
tcpi_pmtu: u32,
tcpi_rcv_ssthresh: u32,
pub tcpi_rtt: u32,
tcpi_rttvar: u32,
}
impl TCP_INFO {
pub unsafe fn new() -> Self {
mem::zeroed()
}
pub fn len() -> socklen_t {
mem::size_of::<Self>() as socklen_t
}
}
#[cfg(target_os = "linux")]
fn set_opt<T: Copy>(sock: c_int, opt: c_int, val: c_int, payload: T) -> io::Result<()> {
unsafe {
let payload = &payload as *const T as *const c_void;
cvt_linux_error(libc::setsockopt(
sock,
opt,
val,
payload as *const _,
mem::size_of::<T>() as socklen_t,
))?;
Ok(())
}
}
#[cfg(target_os = "linux")]
fn get_opt<T>(
sock: c_int,
opt: c_int,
val: c_int,
payload: &mut T,
size: &mut socklen_t,
) -> io::Result<()> {
unsafe {
let payload = payload as *mut T as *mut c_void;
cvt_linux_error(libc::getsockopt(sock, opt, val, payload as *mut _, size))?;
Ok(())
}
}
#[cfg(target_os = "linux")]
fn cvt_linux_error(t: i32) -> io::Result<i32> {
if t == -1 {
Err(io::Error::last_os_error())
} else {
Ok(t)
}
}
#[cfg(target_os = "linux")]
fn ip_bind_addr_no_port(fd: RawFd, val: bool) -> io::Result<()> {
const IP_BIND_ADDRESS_NO_PORT: i32 = 24;
set_opt(fd, libc::IPPROTO_IP, IP_BIND_ADDRESS_NO_PORT, val as c_int)
}
#[cfg(not(target_os = "linux"))]
fn ip_bind_addr_no_port(_fd: RawFd, _val: bool) -> io::Result<()> {
Ok(())
}
#[cfg(target_os = "linux")]
fn set_so_keepalive(fd: RawFd, val: bool) -> io::Result<()> {
set_opt(fd, libc::SOL_SOCKET, libc::SO_KEEPALIVE, val as c_int)
}
#[cfg(target_os = "linux")]
fn set_so_keepalive_idle(fd: RawFd, val: Duration) -> io::Result<()> {
set_opt(
fd,
libc::IPPROTO_TCP,
libc::TCP_KEEPIDLE,
val.as_secs() as c_int, )
}
#[cfg(target_os = "linux")]
fn set_so_keepalive_interval(fd: RawFd, val: Duration) -> io::Result<()> {
set_opt(
fd,
libc::IPPROTO_TCP,
libc::TCP_KEEPINTVL,
val.as_secs() as c_int, )
}
#[cfg(target_os = "linux")]
fn set_so_keepalive_count(fd: RawFd, val: usize) -> io::Result<()> {
set_opt(fd, libc::IPPROTO_TCP, libc::TCP_KEEPCNT, val as c_int)
}
#[cfg(target_os = "linux")]
fn set_keepalive(fd: RawFd, ka: &TcpKeepalive) -> io::Result<()> {
set_so_keepalive(fd, true)?;
set_so_keepalive_idle(fd, ka.idle)?;
set_so_keepalive_interval(fd, ka.interval)?;
set_so_keepalive_count(fd, ka.count)
}
#[cfg(not(target_os = "linux"))]
fn set_keepalive(_fd: RawFd, _ka: &TcpKeepalive) -> io::Result<()> {
Ok(())
}
#[cfg(target_os = "linux")]
pub fn get_tcp_info(fd: RawFd) -> io::Result<TCP_INFO> {
let mut tcp_info = unsafe { TCP_INFO::new() };
let mut data_len: socklen_t = TCP_INFO::len();
get_opt(
fd,
libc::IPPROTO_TCP,
libc::TCP_INFO,
&mut tcp_info,
&mut data_len,
)?;
if data_len != TCP_INFO::len() {
return Err(std::io::Error::new(
std::io::ErrorKind::Other,
"TCP_INFO struct size mismatch",
));
}
Ok(tcp_info)
}
#[cfg(not(target_os = "linux"))]
pub fn get_tcp_info(_fd: RawFd) -> io::Result<TCP_INFO> {
Ok(unsafe { TCP_INFO::new() })
}
#[cfg(target_os = "linux")]
pub fn set_recv_buf(fd: RawFd, val: usize) -> Result<()> {
set_opt(fd, libc::SOL_SOCKET, libc::SO_RCVBUF, val as c_int)
.or_err(ConnectError, "failed to set SO_RCVBUF")
}
#[cfg(not(target_os = "linux"))]
pub fn set_recv_buf(_fd: RawFd, _: usize) -> Result<()> {
Ok(())
}
pub async fn connect(addr: &SocketAddr, bind_to: Option<&SocketAddr>) -> Result<TcpStream> {
let socket = if addr.is_ipv4() {
TcpSocket::new_v4()
} else {
TcpSocket::new_v6()
}
.or_err(SocketError, "failed to create socket")?;
if cfg!(target_os = "linux") {
ip_bind_addr_no_port(socket.as_raw_fd(), true)
.or_err(SocketError, "failed to set socket opts")?;
if let Some(baddr) = bind_to {
socket
.bind(*baddr)
.or_err_with(BindError, || format!("failed to bind to socket {}", *baddr))?;
};
}
socket
.connect(*addr)
.await
.map_err(|e| wrap_os_connect_error(e, format!("Fail to connect to {}", *addr)))
}
pub async fn connect_uds(path: &std::path::Path) -> Result<UnixStream> {
UnixStream::connect(path)
.await
.map_err(|e| wrap_os_connect_error(e, format!("Fail to connect to {}", path.display())))
}
fn wrap_os_connect_error(e: std::io::Error, context: String) -> Box<Error> {
match e.kind() {
ErrorKind::ConnectionRefused => Error::because(ConnectRefused, context, e),
ErrorKind::TimedOut => Error::because(ConnectTimedout, context, e),
ErrorKind::PermissionDenied | ErrorKind::AddrInUse | ErrorKind::AddrNotAvailable => {
Error::because(InternalError, context, e)
}
_ => match e.raw_os_error() {
Some(code) => match code {
libc::ENETUNREACH | libc::EHOSTUNREACH => {
Error::because(ConnectNoRoute, context, e)
}
_ => Error::because(ConnectError, context, e),
},
None => Error::because(ConnectError, context, e),
},
}
}
#[derive(Clone, Debug)]
pub struct TcpKeepalive {
pub idle: Duration,
pub interval: Duration,
pub count: usize,
}
impl std::fmt::Display for TcpKeepalive {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{:?}/{:?}/{}", self.idle, self.interval, self.count)
}
}
pub fn set_tcp_keepalive(stream: &TcpStream, ka: &TcpKeepalive) -> Result<()> {
let fd = stream.as_raw_fd();
set_keepalive(fd, ka).or_err(ConnectError, "failed to set keepalive")
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn test_set_recv_buf() {
use tokio::net::TcpSocket;
let socket = TcpSocket::new_v4().unwrap();
set_recv_buf(socket.as_raw_fd(), 102400).unwrap();
#[cfg(target_os = "linux")]
{
let mut recv_size: c_int = 0;
let mut size = std::mem::size_of::<c_int>() as u32;
get_opt(
socket.as_raw_fd(),
libc::SOL_SOCKET,
libc::SO_RCVBUF,
&mut recv_size,
&mut size,
)
.unwrap();
assert_eq!(recv_size, 102400 * 2);
}
}
}