use std::io;
use std::net::{SocketAddr, TcpListener, TcpStream, UdpSocket};
use std::time::Duration;
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct SocketOptions {
pub broadcast: bool,
pub reuse_address: bool,
pub reuse_port: bool,
}
impl SocketOptions {
pub const FANOUT: Self = Self {
broadcast: false,
reuse_address: true,
reuse_port: true,
};
pub const REUSE_ADDRESS: Self = Self {
broadcast: false,
reuse_address: true,
reuse_port: false,
};
}
pub fn udp_socket(local: SocketAddr, opts: SocketOptions) -> io::Result<UdpSocket> {
sys::udp_socket(local, opts)
}
pub fn tcp_listener(
local: SocketAddr,
opts: SocketOptions,
backlog: i32,
) -> io::Result<TcpListener> {
sys::tcp_listener(local, opts, backlog)
}
pub fn tcp_connect(
remote: SocketAddr,
local: Option<SocketAddr>,
opts: SocketOptions,
timeout: Duration,
) -> io::Result<TcpStream> {
sys::tcp_connect(remote, local, opts, timeout)
}
#[cfg(not(epics_embedded_target))]
mod sys {
use super::SocketOptions;
use std::io;
use std::net::{SocketAddr, TcpListener, TcpStream, UdpSocket};
use std::time::Duration;
fn new_socket(
addr_is_v6: bool,
ty: socket2::Type,
protocol: socket2::Protocol,
opts: SocketOptions,
) -> io::Result<socket2::Socket> {
let domain = if addr_is_v6 {
socket2::Domain::IPV6
} else {
socket2::Domain::IPV4
};
let socket = socket2::Socket::new(domain, ty, Some(protocol))?;
if opts.broadcast {
socket.set_broadcast(true)?;
}
if opts.reuse_port {
#[cfg(unix)]
socket.set_reuse_port(true)?;
#[cfg(not(unix))]
socket.set_reuse_address(true)?;
}
if opts.reuse_address {
socket.set_reuse_address(true)?;
}
Ok(socket)
}
pub(super) fn udp_socket(local: SocketAddr, opts: SocketOptions) -> io::Result<UdpSocket> {
let socket = new_socket(
local.is_ipv6(),
socket2::Type::DGRAM,
socket2::Protocol::UDP,
opts,
)?;
socket.bind(&local.into())?;
Ok(UdpSocket::from(socket))
}
pub(super) fn tcp_listener(
local: SocketAddr,
opts: SocketOptions,
backlog: i32,
) -> io::Result<TcpListener> {
let socket = new_socket(
local.is_ipv6(),
socket2::Type::STREAM,
socket2::Protocol::TCP,
opts,
)?;
socket.bind(&local.into())?;
socket.listen(backlog)?;
Ok(TcpListener::from(socket))
}
pub(super) fn tcp_connect(
remote: SocketAddr,
local: Option<SocketAddr>,
opts: SocketOptions,
timeout: Duration,
) -> io::Result<TcpStream> {
let socket = new_socket(
remote.is_ipv6(),
socket2::Type::STREAM,
socket2::Protocol::TCP,
opts,
)?;
if let Some(local) = local {
socket.bind(&local.into())?;
}
match socket.connect_timeout(&remote.into(), timeout) {
Ok(()) => Ok(TcpStream::from(socket)),
Err(e) => match socket.peer_addr() {
Ok(_) => Ok(TcpStream::from(socket)),
Err(_) => Err(e),
},
}
}
}
#[cfg(epics_embedded_target)]
mod sys {
use super::SocketOptions;
use std::io;
use std::net::{SocketAddr, TcpListener, TcpStream, UdpSocket};
use std::os::fd::{FromRawFd, RawFd};
use std::time::Duration;
struct OwnedFd(RawFd);
impl Drop for OwnedFd {
fn drop(&mut self) {
unsafe { libc::close(self.0) };
}
}
impl OwnedFd {
fn into_raw(self) -> RawFd {
let fd = self.0;
std::mem::forget(self);
fd
}
}
fn last_error() -> io::Error {
io::Error::last_os_error()
}
fn set_bool_opt(fd: RawFd, level: libc::c_int, opt: libc::c_int) -> io::Result<()> {
let one: libc::c_int = 1;
let rc = unsafe {
libc::setsockopt(
fd,
level,
opt,
&one as *const libc::c_int as *const libc::c_void,
std::mem::size_of::<libc::c_int>() as libc::socklen_t,
)
};
if rc != 0 {
return Err(last_error());
}
Ok(())
}
fn new_socket(
ty: libc::c_int,
protocol: libc::c_int,
opts: SocketOptions,
) -> io::Result<OwnedFd> {
let fd = unsafe { libc::socket(libc::AF_INET, ty, protocol) };
if fd < 0 {
return Err(last_error());
}
let owned = OwnedFd(fd);
if opts.broadcast {
set_bool_opt(fd, libc::SOL_SOCKET, libc::SO_BROADCAST)?;
}
if opts.reuse_port {
set_bool_opt(fd, libc::SOL_SOCKET, libc::SO_REUSEPORT)?;
}
if opts.reuse_address {
set_bool_opt(fd, libc::SOL_SOCKET, libc::SO_REUSEADDR)?;
}
Ok(owned)
}
fn sockaddr_in(addr: SocketAddr) -> io::Result<libc::sockaddr_in> {
let v4 = match addr {
SocketAddr::V4(v4) => v4,
SocketAddr::V6(_) => {
return Err(io::Error::new(
io::ErrorKind::Unsupported,
"IPv6 is not supported on this target",
));
}
};
let mut sin: libc::sockaddr_in = unsafe { std::mem::zeroed() };
sin.sin_family = libc::AF_INET as libc::sa_family_t;
sin.sin_port = v4.port().to_be();
sin.sin_addr = libc::in_addr {
s_addr: u32::from(*v4.ip()).to_be(),
};
Ok(sin)
}
fn bind_fd(fd: RawFd, addr: SocketAddr) -> io::Result<()> {
let sin = sockaddr_in(addr)?;
let rc = unsafe {
libc::bind(
fd,
&sin as *const libc::sockaddr_in as *const libc::sockaddr,
std::mem::size_of::<libc::sockaddr_in>() as libc::socklen_t,
)
};
if rc != 0 {
return Err(last_error());
}
Ok(())
}
pub(super) fn udp_socket(local: SocketAddr, opts: SocketOptions) -> io::Result<UdpSocket> {
let owned = new_socket(libc::SOCK_DGRAM, libc::IPPROTO_UDP, opts)?;
bind_fd(owned.0, local)?;
Ok(unsafe { UdpSocket::from_raw_fd(owned.into_raw()) })
}
pub(super) fn tcp_listener(
local: SocketAddr,
opts: SocketOptions,
backlog: i32,
) -> io::Result<TcpListener> {
let owned = new_socket(libc::SOCK_STREAM, libc::IPPROTO_TCP, opts)?;
bind_fd(owned.0, local)?;
if unsafe { libc::listen(owned.0, backlog) } != 0 {
return Err(last_error());
}
Ok(unsafe { TcpListener::from_raw_fd(owned.into_raw()) })
}
#[cfg(not(any(target_os = "rtems", target_os = "vxworks")))]
compile_error!(
"epics_embedded_target gained a triple beyond rtems/vxworks: choose \
its connect_fd and set_nonblocking arms explicitly"
);
#[cfg(target_os = "vxworks")]
fn set_nonblocking(fd: RawFd, on: bool) -> io::Result<()> {
let mut flags: libc::c_int = i32::from(on);
let rc = unsafe { libc::ioctl(fd, libc::FIONBIO, &mut flags as *mut libc::c_int) };
if rc < 0 {
return Err(last_error());
}
Ok(())
}
pub(super) fn tcp_connect(
remote: SocketAddr,
local: Option<SocketAddr>,
opts: SocketOptions,
timeout: Duration,
) -> io::Result<TcpStream> {
let owned = new_socket(libc::SOCK_STREAM, libc::IPPROTO_TCP, opts)?;
if let Some(local) = local {
bind_fd(owned.0, local)?;
}
connect_fd(&owned, remote, timeout)?;
Ok(unsafe { TcpStream::from_raw_fd(owned.into_raw()) })
}
#[cfg(target_os = "rtems")]
fn connect_fd(owned: &OwnedFd, remote: SocketAddr, _timeout: Duration) -> io::Result<()> {
let sin = sockaddr_in(remote)?;
let rc = unsafe {
libc::connect(
owned.0,
&sin as *const libc::sockaddr_in as *const libc::sockaddr,
std::mem::size_of::<libc::sockaddr_in>() as libc::socklen_t,
)
};
if rc != 0 {
return Err(last_error());
}
Ok(())
}
#[cfg(target_os = "vxworks")]
fn connect_fd(owned: &OwnedFd, remote: SocketAddr, timeout: Duration) -> io::Result<()> {
let sin = sockaddr_in(remote)?;
set_nonblocking(owned.0, true)?;
let rc = unsafe {
libc::connect(
owned.0,
&sin as *const libc::sockaddr_in as *const libc::sockaddr,
std::mem::size_of::<libc::sockaddr_in>() as libc::socklen_t,
)
};
if rc == 0 {
set_nonblocking(owned.0, false)?;
return Ok(());
}
let err = last_error();
let in_progress = matches!(
err.raw_os_error(),
Some(e) if e == libc::EINPROGRESS || e == libc::EWOULDBLOCK
);
if !in_progress {
return Err(err);
}
let ms = i32::try_from(timeout.as_millis()).unwrap_or(i32::MAX);
let mut pfd = libc::pollfd {
fd: owned.0,
events: libc::POLLOUT,
revents: 0,
};
let n = unsafe { libc::poll(&mut pfd as *mut libc::pollfd, 1, ms) };
if n < 0 {
return Err(last_error());
}
if n == 0 {
return Err(io::Error::new(io::ErrorKind::TimedOut, "connect timed out"));
}
let mut so_error: libc::c_int = 0;
let mut len = std::mem::size_of::<libc::c_int>() as libc::socklen_t;
let rc = unsafe {
libc::getsockopt(
owned.0,
libc::SOL_SOCKET,
libc::SO_ERROR,
&mut so_error as *mut libc::c_int as *mut libc::c_void,
&mut len as *mut libc::socklen_t,
)
};
if rc != 0 {
return Err(last_error());
}
if so_error != 0 {
return Err(io::Error::from_raw_os_error(so_error));
}
set_nonblocking(owned.0, false)?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::{Ipv4Addr, SocketAddrV4};
fn localhost(port: u16) -> SocketAddr {
SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, port))
}
#[test]
fn udp_binds_and_reports_its_port() {
let sock = udp_socket(localhost(0), SocketOptions::default()).unwrap();
assert_ne!(sock.local_addr().unwrap().port(), 0);
}
#[test]
fn fanout_options_let_two_sockets_share_a_port() {
let first = udp_socket(localhost(0), SocketOptions::FANOUT).unwrap();
let port = first.local_addr().unwrap().port();
let second = udp_socket(localhost(port), SocketOptions::FANOUT).unwrap();
assert_eq!(second.local_addr().unwrap().port(), port);
}
#[test]
fn without_fanout_options_a_shared_port_is_refused() {
let first = udp_socket(localhost(0), SocketOptions::default()).unwrap();
let port = first.local_addr().unwrap().port();
assert!(udp_socket(localhost(port), SocketOptions::default()).is_err());
}
#[test]
fn tcp_listener_accepts_a_connect() {
let listener = tcp_listener(localhost(0), SocketOptions::REUSE_ADDRESS, 8).unwrap();
let addr = listener.local_addr().unwrap();
let joiner = std::thread::spawn(move || listener.accept().map(|(s, _)| s));
let client =
tcp_connect(addr, None, SocketOptions::default(), Duration::from_secs(5)).unwrap();
let accepted = joiner.join().unwrap().unwrap();
assert_eq!(accepted.local_addr().unwrap().port(), addr.port());
assert_eq!(client.peer_addr().unwrap().port(), addr.port());
}
#[test]
fn tcp_connect_to_a_closed_port_fails() {
let port = {
let probe = tcp_listener(localhost(0), SocketOptions::default(), 1).unwrap();
probe.local_addr().unwrap().port()
};
let r = tcp_connect(
localhost(port),
None,
SocketOptions::default(),
Duration::from_secs(5),
);
assert!(r.is_err());
}
#[test]
fn tcp_connect_honours_a_local_bind() {
let listener = tcp_listener(localhost(0), SocketOptions::REUSE_ADDRESS, 8).unwrap();
let addr = listener.local_addr().unwrap();
let joiner = std::thread::spawn(move || listener.accept().map(|(s, _)| s));
let client = tcp_connect(
addr,
Some(localhost(0)),
SocketOptions::REUSE_ADDRESS,
Duration::from_secs(5),
)
.unwrap();
let accepted = joiner.join().unwrap().unwrap();
assert_eq!(
accepted.peer_addr().unwrap().port(),
client.local_addr().unwrap().port()
);
}
}