use alloc::{boxed::Box, collections::VecDeque, sync::Arc, vec::Vec};
use core::sync::atomic::{AtomicBool, Ordering};
use ax_io::{IoBuf, Read, Write};
use ax_sync::SpinLock;
use axpoll::{ExclusiveRegistrationSink, IoEvents, Pollable, SharedRegistrationSink};
use axpoll_set::PollSet;
use ringbuf::{
HeapCons, HeapProd, HeapRb,
traits::{Consumer, Observer, Producer, Split},
};
use crate::{
CMsgData, NetError, NetResult, RecvOptions, SendOptions, Shutdown,
general::GeneralOptions,
options::{Configurable, GetSocketOption, SetSocketOption, UnixCredentials},
unix::{Transport, TransportOps, UnixSocketAddr},
};
const BUF_SIZE: usize = 64 * 1024;
struct PendingCmsg {
start_byte: u64,
end_byte: u64,
cmsg: Vec<CMsgData>,
}
type CmsgQueue = Arc<SpinLock<VecDeque<PendingCmsg>>>;
fn new_uni_channel() -> (HeapProd<u8>, HeapCons<u8>) {
let rb = HeapRb::new(BUF_SIZE);
rb.split()
}
fn new_channels(
credentials: UnixCredentials,
first_receive_credentials: Arc<AtomicBool>,
second_receive_credentials: Arc<AtomicBool>,
) -> (Channel, Channel) {
let (client_tx, server_rx) = new_uni_channel();
let (server_tx, client_rx) = new_uni_channel();
let poll_update = Arc::new(PollSet::new());
let c2s_cmsg = CmsgQueue::default();
let s2c_cmsg = CmsgQueue::default();
let client_tx_closed = Arc::new(AtomicBool::new(false));
let server_tx_closed = Arc::new(AtomicBool::new(false));
(
Channel {
tx: client_tx,
rx: client_rx,
tx_cmsg: c2s_cmsg.clone(),
rx_cmsg: s2c_cmsg.clone(),
tx_bytes_total: 0,
rx_bytes_total: 0,
my_tx_closed: client_tx_closed.clone(),
peer_tx_closed: server_tx_closed.clone(),
poll_update: poll_update.clone(),
peer_credentials: credentials.clone(),
peer_receive_credentials: second_receive_credentials,
},
Channel {
tx: server_tx,
rx: server_rx,
tx_cmsg: s2c_cmsg,
rx_cmsg: c2s_cmsg,
tx_bytes_total: 0,
rx_bytes_total: 0,
my_tx_closed: server_tx_closed,
peer_tx_closed: client_tx_closed,
poll_update,
peer_credentials: credentials,
peer_receive_credentials: first_receive_credentials,
},
)
}
struct Channel {
tx: HeapProd<u8>,
rx: HeapCons<u8>,
tx_cmsg: CmsgQueue,
rx_cmsg: CmsgQueue,
tx_bytes_total: u64,
rx_bytes_total: u64,
my_tx_closed: Arc<AtomicBool>,
peer_tx_closed: Arc<AtomicBool>,
poll_update: Arc<PollSet>,
peer_credentials: UnixCredentials,
peer_receive_credentials: Arc<AtomicBool>,
}
pub struct Bind {
conn_tx: async_channel::Sender<ConnRequest>,
poll_new_conn: Arc<PollSet>,
listening: Arc<AtomicBool>,
credentials: UnixCredentials,
receive_credentials: Arc<AtomicBool>,
}
impl Bind {
fn connect(
&self,
local_addr: UnixSocketAddr,
credentials: UnixCredentials,
client_receive_credentials: Arc<AtomicBool>,
) -> NetResult<(Channel, Arc<PollSet>)> {
if !self.listening.load(Ordering::Acquire) {
return Err(NetError::ConnectionRefused);
}
let server_receive_credentials = Arc::new(AtomicBool::new(
self.receive_credentials.load(Ordering::Acquire),
));
let (mut client_chan, mut server_chan) = new_channels(
UnixCredentials::new(0),
client_receive_credentials,
server_receive_credentials.clone(),
);
client_chan.peer_credentials = self.credentials.clone();
server_chan.peer_credentials = credentials.clone();
self.conn_tx
.try_send(ConnRequest {
channel: server_chan,
addr: local_addr,
credentials,
receive_credentials: server_receive_credentials,
})
.map_err(|_| NetError::ConnectionRefused)?;
Ok((client_chan, self.poll_new_conn.clone()))
}
}
struct ConnRequest {
channel: Channel,
addr: UnixSocketAddr,
credentials: UnixCredentials,
receive_credentials: Arc<AtomicBool>,
}
pub struct StreamTransport {
channel: SpinLock<Option<Channel>>,
conn_rx: SpinLock<Option<(async_channel::Receiver<ConnRequest>, Arc<PollSet>)>>,
listening: Arc<AtomicBool>,
poll_state: PollSet,
general: GeneralOptions,
receive_credentials: Arc<AtomicBool>,
credentials: UnixCredentials,
rx_closed: AtomicBool,
tx_closed: AtomicBool,
}
impl StreamTransport {
pub fn new(credentials: impl Into<UnixCredentials>) -> Self {
StreamTransport::new_channel(None, credentials.into(), Arc::new(AtomicBool::new(false)))
}
fn new_channel(
channel: Option<Channel>,
credentials: UnixCredentials,
receive_credentials: Arc<AtomicBool>,
) -> Self {
StreamTransport {
channel: SpinLock::new(channel),
conn_rx: SpinLock::new(None),
listening: Arc::new(AtomicBool::new(false)),
poll_state: PollSet::new(),
general: GeneralOptions::new(1, 1, 0), receive_credentials,
credentials,
rx_closed: AtomicBool::new(false),
tx_closed: AtomicBool::new(false),
}
}
pub fn new_pair(credentials: impl Into<UnixCredentials>) -> (Self, Self) {
let credentials = credentials.into();
let credentials1 = Arc::new(AtomicBool::new(false));
let credentials2 = Arc::new(AtomicBool::new(false));
let (chan1, chan2) = new_channels(
credentials.clone(),
credentials1.clone(),
credentials2.clone(),
);
let transport1 =
StreamTransport::new_channel(Some(chan1), credentials.clone(), credentials1);
let transport2 = StreamTransport::new_channel(Some(chan2), credentials, credentials2);
(transport1, transport2)
}
pub(super) fn wake_connected(&self) {
unsafe { self.poll_state.wake(IoEvents::IN | IoEvents::OUT) };
}
}
impl Configurable for StreamTransport {
fn get_option_inner(&self, opt: &mut GetSocketOption) -> NetResult<bool> {
use GetSocketOption as O;
if self.general.get_option_inner(opt)? {
return Ok(true);
}
match opt {
O::SendBuffer(size) => {
**size = BUF_SIZE;
}
O::PassCredentials(enabled) => {
**enabled = self.receive_credentials.load(Ordering::Acquire);
}
O::PeerCredentials(cred) => {
let peer_credentials = self.channel.lock().as_ref().map_or_else(
|| self.credentials.clone(),
|chan| chan.peer_credentials.clone(),
);
**cred = peer_credentials;
}
_ => return Ok(false),
}
Ok(true)
}
fn set_option_inner(&self, opt: SetSocketOption) -> NetResult<bool> {
use SetSocketOption as O;
if self.general.set_option_inner(opt)? {
return Ok(true);
}
match opt {
O::PassCredentials(enabled) => {
self.receive_credentials.store(*enabled, Ordering::Release);
}
_ => return Ok(false),
}
Ok(true)
}
}
impl TransportOps for StreamTransport {
fn bind(&self, slot: &super::BindSlot, _local_addr: &UnixSocketAddr) -> NetResult<()> {
let mut slot = slot.stream.lock();
if slot.is_some() {
return Err(NetError::AddrInUse);
}
let mut guard = self.conn_rx.lock();
if guard.is_some() {
return Err(NetError::InvalidInput);
}
let (tx, rx) = async_channel::unbounded();
let poll = Arc::new(PollSet::new());
*slot = Some(Bind {
conn_tx: tx,
poll_new_conn: poll.clone(),
listening: self.listening.clone(),
credentials: self.credentials.clone(),
receive_credentials: self.receive_credentials.clone(),
});
*guard = Some((rx, poll));
drop(guard);
drop(slot);
unsafe { self.poll_state.wake(IoEvents::IN | IoEvents::OUT) };
Ok(())
}
fn listen(&self) -> NetResult<()> {
if self.conn_rx.lock().is_none() {
return Err(NetError::InvalidInput);
}
self.listening.store(true, Ordering::Release);
Ok(())
}
fn is_listening(&self) -> bool {
self.listening.load(Ordering::Acquire)
}
fn connect(
&self,
slot: &super::BindSlot,
local_addr: &UnixSocketAddr,
) -> NetResult<Option<Arc<PollSet>>> {
let mut guard = self.channel.lock();
if guard.is_some() {
return Err(NetError::AlreadyConnected);
}
let (channel, accept_poll) = {
let slot = slot.stream.lock();
slot.as_ref().ok_or(NetError::NotConnected)?.connect(
local_addr.clone(),
self.credentials.clone(),
self.receive_credentials.clone(),
)?
};
*guard = Some(channel);
Ok(Some(accept_poll))
}
fn try_accept(&self) -> NetResult<(Transport, UnixSocketAddr)> {
if !self.is_listening() {
return Err(NetError::InvalidInput);
}
let Some((rx, _)) = self.conn_rx.lock().clone() else {
return Err(NetError::InvalidInput);
};
match rx.try_recv() {
Ok(ConnRequest {
channel,
addr: peer_addr,
credentials,
receive_credentials,
}) => Ok((
Transport::Stream(StreamTransport::new_channel(
Some(channel),
credentials,
receive_credentials,
)),
peer_addr,
)),
Err(async_channel::TryRecvError::Empty) => Err(NetError::WouldBlock),
Err(async_channel::TryRecvError::Closed) => Err(NetError::ConnectionReset),
}
}
fn try_send(&self, mut src: impl Read + IoBuf, options: &mut SendOptions) -> NetResult<usize> {
if options.to.is_some() {
return Err(NetError::InvalidInput);
}
let size = src.remaining();
if size == 0 {
return Ok(0);
}
let mut wake_poll = None;
let mut guard = self.channel.lock();
let result = {
let Some(chan) = guard.as_mut() else {
return Err(NetError::NotConnected);
};
if !chan.tx.read_is_held() {
return Err(NetError::BrokenPipe);
}
let count = {
let (left, right) = chan.tx.vacant_slices_mut();
let mut count = src.read(unsafe { left.assume_init_mut() })?;
if count >= left.len() {
count += src.read(unsafe { right.assume_init_mut() })?;
}
unsafe { chan.tx.advance_write_index(count) };
count
};
if count == 0 {
Err(NetError::WouldBlock)
} else {
if chan.peer_receive_credentials.load(Ordering::Acquire)
&& let Some(credentials) = options.sender_credentials.clone()
{
options
.cmsg
.push(Box::new(crate::SocketCmsg::Credentials(credentials)));
}
let cmsg = core::mem::take(&mut options.cmsg);
if !cmsg.is_empty() {
chan.tx_cmsg.lock().push_back(PendingCmsg {
start_byte: chan.tx_bytes_total.saturating_add(1),
end_byte: chan.tx_bytes_total.saturating_add(count as u64),
cmsg,
});
}
chan.tx_bytes_total = chan.tx_bytes_total.saturating_add(count as u64);
wake_poll = Some(chan.poll_update.clone());
Ok(count)
}
};
drop(guard);
if let Some(poll) = wake_poll {
unsafe { poll.wake(IoEvents::IN | IoEvents::OUT) };
}
result
}
fn try_recv(&self, mut dst: impl Write, options: &mut RecvOptions) -> NetResult<usize> {
let peek = options.flags.contains(crate::RecvFlags::PEEK);
let recv_count = {
let mut wake_poll = None;
let mut guard = self.channel.lock();
let result = {
let Some(chan) = guard.as_mut() else {
return Err(NetError::NotConnected);
};
let cap_bytes: Option<usize> = {
let q = chan.rx_cmsg.lock();
q.front().and_then(|front| {
if front.end_byte > chan.rx_bytes_total {
let cap = front.end_byte.saturating_sub(chan.rx_bytes_total);
Some(cap as usize)
} else {
None
}
})
};
let count = {
let (left, right) = chan.rx.as_slices();
let left_cap = cap_bytes.map_or(left.len(), |c| c.min(left.len()));
let mut count = dst.write(&left[..left_cap])?;
let remaining_cap = cap_bytes.map_or(usize::MAX, |c| c.saturating_sub(count));
if count >= left_cap && remaining_cap > 0 {
let right_cap = right.len().min(remaining_cap);
count += dst.write(&right[..right_cap])?;
}
if !peek {
unsafe { chan.rx.advance_read_index(count) };
}
count
};
if count > 0 {
if !peek {
chan.rx_bytes_total = chan.rx_bytes_total.saturating_add(count as u64);
wake_poll = Some(chan.poll_update.clone());
}
Ok(count)
} else if !chan.rx.write_is_held() || chan.peer_tx_closed.load(Ordering::Acquire) {
Ok(0)
} else {
Err(NetError::WouldBlock)
}
};
drop(guard);
if let Some(poll) = wake_poll {
unsafe { poll.wake(IoEvents::OUT) };
}
result
}?;
if peek {
if let Some(dst) = options.cmsg.as_deref_mut() {
let mut guard = self.channel.lock();
if let Some(chan) = guard.as_mut() {
let ready_upto = chan.rx_bytes_total.saturating_add(recv_count as u64);
let q = chan.rx_cmsg.lock();
for entry in q.iter() {
if entry.start_byte > ready_upto {
break;
}
dst.extend(entry.cmsg.iter().map(|c| c.clone_box()));
}
}
}
return Ok(recv_count);
}
let mut dst_cmsg = options.cmsg.as_deref_mut();
let mut guard = self.channel.lock();
if let Some(chan) = guard.as_mut() {
let mut q = chan.rx_cmsg.lock();
while let Some(front) = q.front()
&& front.start_byte <= chan.rx_bytes_total
{
let entry = q.pop_front().unwrap();
if let Some(dst) = dst_cmsg.as_deref_mut() {
dst.extend(entry.cmsg);
}
}
}
Ok(recv_count)
}
fn shutdown(&self, how: Shutdown) -> NetResult<()> {
if how.has_read() {
self.rx_closed.store(true, Ordering::Release);
}
let mut peer_poll = None;
if how.has_write() {
self.tx_closed.store(true, Ordering::Release);
if let Some(chan) = self.channel.lock().as_ref() {
chan.my_tx_closed.store(true, Ordering::Release);
peer_poll = Some(chan.poll_update.clone());
}
}
if self.rx_closed.load(Ordering::Acquire)
&& self.tx_closed.load(Ordering::Acquire)
&& let Some(chan) = self.channel.lock().take()
{
peer_poll.get_or_insert(chan.poll_update);
}
if let Some(poll) = peer_poll {
unsafe { poll.wake(IoEvents::IN | IoEvents::OUT | IoEvents::RDHUP) };
}
if how.has_read() || how.has_write() {
unsafe {
self.poll_state
.wake(IoEvents::IN | IoEvents::OUT | IoEvents::RDHUP)
};
}
Ok(())
}
}
impl Pollable for StreamTransport {
fn poll(&self) -> IoEvents {
let mut events = IoEvents::empty();
let rx_closed = self.rx_closed.load(Ordering::Acquire);
let mut peer_eof = false;
if let Some(chan) = self.channel.lock().as_ref() {
peer_eof = chan.peer_tx_closed.load(Ordering::Acquire);
events.set(
IoEvents::IN,
!rx_closed && (chan.rx.occupied_len() > 0 || peer_eof),
);
events.set(
IoEvents::OUT,
!self.tx_closed.load(Ordering::Acquire) && chan.tx.vacant_len() > 0,
);
} else if let Some((conn_tx, _)) = self.conn_rx.lock().as_ref() {
events.set(IoEvents::IN, !conn_tx.is_empty());
}
events.set(IoEvents::RDHUP, peer_eof || rx_closed);
events
}
unsafe fn register_shared(&self, sink: &mut dyn SharedRegistrationSink, events: IoEvents) {
self.register_poll_sources(events, |poll, interests| unsafe {
sink.register_shared(poll, interests)
});
}
unsafe fn register_exclusive(
&self,
sink: &mut dyn ExclusiveRegistrationSink,
events: IoEvents,
) {
self.register_poll_sources(events, |poll, interests| unsafe {
sink.register_exclusive(poll, interests)
});
}
}
impl StreamTransport {
fn register_poll_sources(
&self,
events: IoEvents,
mut register: impl FnMut(&PollSet, IoEvents),
) {
let chan_poll = if events.intersects(IoEvents::IN | IoEvents::OUT | IoEvents::RDHUP) {
self.channel
.lock()
.as_ref()
.map(|chan| chan.poll_update.clone())
} else {
None
};
if let Some(poll) = chan_poll {
register(&poll, events);
} else if let Some((_, poll_new_conn)) = self.conn_rx.lock().as_ref()
&& events.contains(IoEvents::IN)
{
register(poll_new_conn, IoEvents::IN);
}
register(&self.poll_state, events);
}
}
impl Drop for StreamTransport {
fn drop(&mut self) {
let peer_poll = if let Some(chan) = self.channel.lock().as_ref() {
chan.my_tx_closed.store(true, Ordering::Release);
Some(chan.poll_update.clone())
} else {
None
};
if let Some(poll) = peer_poll {
unsafe { poll.wake(IoEvents::IN | IoEvents::RDHUP) };
}
unsafe {
self.poll_state
.wake(IoEvents::IN | IoEvents::OUT | IoEvents::RDHUP)
};
}
}