use std::net::{Ipv4Addr, SocketAddrV4};
use std::os::fd::{AsRawFd, RawFd};
use std::{io, mem, ptr};
use super::sndrcv::{RecvFlags, SendFlags};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum L4Protocol {
Tcp,
Udp,
Icmp,
Sctp,
Dccp,
Custom(u8),
}
pub struct L4Socket {
fd: i32,
}
impl L4Socket {
#[inline]
pub fn new(protocol: L4Protocol) -> io::Result<L4Socket> {
let protocol = match protocol {
L4Protocol::Dccp => libc::IPPROTO_DCCP,
L4Protocol::Icmp => libc::IPPROTO_ICMP,
L4Protocol::Sctp => libc::IPPROTO_SCTP,
L4Protocol::Tcp => libc::IPPROTO_TCP,
L4Protocol::Udp => libc::IPPROTO_UDP,
L4Protocol::Custom(protocol) => {
if protocol == libc::IPPROTO_RAW as u8 {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"IPPROTO_RAW not supported for L4Socket",
));
}
protocol as i32
}
};
match unsafe { libc::socket(libc::AF_INET, libc::SOCK_RAW, protocol) } {
..=-1 => Err(std::io::Error::last_os_error()),
fd => Ok(L4Socket { fd }),
}
}
pub fn bind(&self, addr: &SocketAddrV4) -> io::Result<()> {
let ip_addr = addr.ip();
let port = addr.port();
let sockaddr = libc::sockaddr_in {
sin_family: libc::AF_INET as u16,
sin_addr: libc::in_addr {
s_addr: u32::from_le_bytes(ip_addr.octets()), },
sin_port: port,
sin_zero: [0u8; 8],
};
match unsafe {
libc::bind(
self.fd,
ptr::addr_of!(sockaddr) as *const libc::sockaddr,
mem::size_of::<libc::sockaddr_in>() as u32,
)
} {
0 => Ok(()),
_ => Err(io::Error::last_os_error()),
}
}
pub fn send(&self, buf: &[u8]) -> io::Result<usize> {
match unsafe { libc::send(self.fd, buf.as_ptr() as *const libc::c_void, buf.len(), 0) } {
..=-1 => Err(io::Error::last_os_error()),
sent => Ok(sent as usize),
}
}
pub fn recv(&self, buf: &mut [u8]) -> io::Result<usize> {
match unsafe { libc::recv(self.fd, buf.as_mut_ptr() as *mut libc::c_void, buf.len(), 0) } {
..=-1 => Err(io::Error::last_os_error()),
recvd => Ok(recvd as usize),
}
}
pub fn send_to(
&self,
buf: &[u8],
rem_addr: SocketAddrV4,
flags: SendFlags,
) -> io::Result<usize> {
let sockaddr = libc::sockaddr_in {
sin_family: libc::AF_INET as u16,
sin_port: rem_addr.port(),
sin_addr: libc::in_addr {
s_addr: u32::from_be_bytes(rem_addr.ip().octets()), },
sin_zero: [0u8; 8],
};
let addrlen = mem::size_of_val(&sockaddr) as u32;
match unsafe {
libc::sendto(
self.fd,
buf.as_ptr() as *mut libc::c_void,
buf.len(),
flags.bits(),
ptr::addr_of!(sockaddr) as *const libc::sockaddr,
addrlen,
)
} {
..=-1 => Err(io::Error::last_os_error()),
recvd => Ok(recvd as usize),
}
}
pub fn recv_from(&self, buf: &[u8], flags: RecvFlags) -> io::Result<(usize, SocketAddrV4)> {
let sockaddr = libc::sockaddr_in {
sin_family: libc::AF_INET as u16,
sin_port: 0,
sin_addr: libc::in_addr { s_addr: 0 },
sin_zero: [0u8; 8],
};
let addrlen = mem::size_of_val(&sockaddr) as u32;
match unsafe {
libc::sendto(
self.fd,
buf.as_ptr() as *mut libc::c_void,
buf.len(),
flags.bits(),
ptr::addr_of!(sockaddr) as *const libc::sockaddr,
addrlen,
)
} {
..=-1 => Err(io::Error::last_os_error()),
recvd => {
let rem_addr = SocketAddrV4::new(
Ipv4Addr::from(sockaddr.sin_addr.s_addr.to_be_bytes()), sockaddr.sin_port,
);
Ok((recvd as usize, rem_addr))
}
}
}
#[inline]
pub fn nonblocking(&self) -> io::Result<bool> {
let flags = unsafe { libc::fcntl(self.fd, libc::F_GETFL) };
if flags < 0 {
return Err(io::Error::last_os_error());
}
Ok(flags & libc::O_NONBLOCK > 0)
}
#[inline]
pub fn set_nonblocking(&self, nonblocking: bool) -> io::Result<()> {
let mut fcntl_flags = match unsafe { libc::fcntl(self.fd, libc::F_GETFL, 0) } {
..=-1 => return Err(io::Error::last_os_error()),
f => f,
};
if nonblocking {
fcntl_flags |= libc::O_NONBLOCK;
} else {
fcntl_flags &= !libc::O_NONBLOCK;
}
match unsafe { libc::fcntl(self.fd, libc::F_SETFL, fcntl_flags) } {
0 => Ok(()),
_ => Err(io::Error::last_os_error()),
}
}
}
impl AsRawFd for L4Socket {
fn as_raw_fd(&self) -> RawFd {
self.fd
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn bind_localhost() {
let sock = L4Socket::new(L4Protocol::Udp).unwrap();
sock.bind(&SocketAddrV4::new(Ipv4Addr::new(127, 0, 0, 1), 777))
.unwrap();
}
}