pub(crate) mod dgram;
pub mod namespace;
pub(crate) mod stream;
use alloc::sync::Arc;
use ax_io::{IoBuf, Read, Write};
use ax_lazyinit::LazyLock;
use ax_sync::SpinLock;
use axpoll::{ExclusiveRegistrationSink, IoEvents, Pollable, SharedRegistrationSink};
use axpoll_set::PollSet;
use enum_dispatch::enum_dispatch;
use hashbrown::HashMap;
pub use self::{
dgram::DgramTransport,
namespace::{UnixNamespace, register_unix_namespace},
stream::StreamTransport,
};
use crate::{
ConnectStatus, NetError, NetResult, RecvOptions, SendOptions, Shutdown, Socket, SocketAddrEx,
SocketOps,
options::{Configurable, GetSocketOption, SetSocketOption},
};
#[derive(Default, Clone, Debug)]
pub enum UnixSocketAddr {
#[default]
Unnamed,
Abstract(Arc<[u8]>),
Path(Arc<str>),
}
#[enum_dispatch]
pub trait TransportOps: Configurable + Pollable + Send + Sync {
fn bind(&self, slot: &BindSlot, local_addr: &UnixSocketAddr) -> NetResult;
fn connect(
&self,
slot: &BindSlot,
local_addr: &UnixSocketAddr,
) -> NetResult<Option<Arc<PollSet>>>;
fn listen(&self) -> NetResult {
Err(NetError::OperationNotSupported)
}
fn is_listening(&self) -> bool {
false
}
fn try_accept(&self) -> NetResult<(Transport, UnixSocketAddr)> {
Err(NetError::WouldBlock)
}
fn try_send(&self, src: impl Read + IoBuf, options: &mut SendOptions) -> NetResult<usize>;
fn try_recv(&self, dst: impl Write, options: &mut RecvOptions<'_>) -> NetResult<usize>;
fn shutdown(&self, _how: Shutdown) -> NetResult {
Ok(())
}
}
#[enum_dispatch(Configurable, TransportOps)]
pub enum Transport {
Stream(StreamTransport),
Dgram(DgramTransport),
}
impl Transport {
fn finish_connect(&self, accept_poll: Option<Arc<PollSet>>) {
if let Some(poll) = accept_poll {
unsafe { poll.wake(IoEvents::IN) };
}
match self {
Transport::Stream(stream) => stream.wake_connected(),
Transport::Dgram(dgram) => dgram.wake_connected(),
}
}
}
impl Pollable for Transport {
fn poll(&self) -> IoEvents {
match self {
Transport::Stream(stream) => stream.poll(),
Transport::Dgram(dgram) => dgram.poll(),
}
}
unsafe fn register_shared(&self, sink: &mut dyn SharedRegistrationSink, events: IoEvents) {
match self {
Transport::Stream(stream) => unsafe { stream.register_shared(sink, events) },
Transport::Dgram(dgram) => unsafe { dgram.register_shared(sink, events) },
}
}
unsafe fn register_exclusive(
&self,
sink: &mut dyn ExclusiveRegistrationSink,
events: IoEvents,
) {
match self {
Transport::Stream(stream) => unsafe { stream.register_exclusive(sink, events) },
Transport::Dgram(dgram) => unsafe { dgram.register_exclusive(sink, events) },
}
}
}
#[derive(Default)]
pub struct BindSlot {
stream: SpinLock<Option<stream::Bind>>,
dgram: SpinLock<Option<dgram::Bind>>,
seqpacket: SpinLock<Option<dgram::SeqBind>>,
}
static ABSTRACT_BINDS: LazyLock<SpinLock<HashMap<Arc<[u8]>, BindSlot>>> =
LazyLock::new(|| SpinLock::new(HashMap::new()));
pub(crate) fn with_slot<R>(
addr: &UnixSocketAddr,
f: impl FnOnce(&BindSlot) -> NetResult<R>,
) -> NetResult<R> {
match addr {
UnixSocketAddr::Unnamed => Err(NetError::InvalidInput),
UnixSocketAddr::Abstract(name) => {
let binds = ABSTRACT_BINDS.lock();
if let Some(slot) = binds.get(name) {
f(slot)
} else {
Err(NetError::NotFound)
}
}
UnixSocketAddr::Path(path) => namespace::with_namespace(|ns| {
let slot = ns.resolve(path.as_ref())?;
f(slot.as_ref())
}),
}
}
fn with_slot_or_insert<R>(
addr: &UnixSocketAddr,
f: impl FnOnce(&BindSlot) -> NetResult<R>,
) -> NetResult<R> {
match addr {
UnixSocketAddr::Unnamed => Err(NetError::InvalidInput),
UnixSocketAddr::Abstract(name) => {
let mut binds = ABSTRACT_BINDS.lock();
f(binds.entry(name.clone()).or_default())
}
UnixSocketAddr::Path(path) => namespace::with_namespace(|ns| {
let slot = ns.bind(path.as_ref())?;
f(slot.as_ref())
}),
}
}
pub struct UnixSocket {
transport: Transport,
local_addr: SpinLock<UnixSocketAddr>,
remote_addr: SpinLock<UnixSocketAddr>,
}
impl UnixSocket {
pub fn new(transport: impl Into<Transport>) -> Self {
Self {
transport: transport.into(),
local_addr: SpinLock::new(UnixSocketAddr::Unnamed),
remote_addr: SpinLock::new(UnixSocketAddr::Unnamed),
}
}
}
impl Configurable for UnixSocket {
fn get_option_inner(&self, opt: &mut GetSocketOption) -> NetResult<bool> {
self.transport.get_option_inner(opt)
}
fn set_option_inner(&self, opt: SetSocketOption) -> NetResult<bool> {
self.transport.set_option_inner(opt)
}
}
impl SocketOps for UnixSocket {
fn bind(&self, local_addr: SocketAddrEx) -> NetResult {
let local_addr = local_addr.into_unix()?;
let mut guard = self.local_addr.lock();
if matches!(&*guard, UnixSocketAddr::Unnamed) {
with_slot_or_insert(&local_addr, |slot| self.transport.bind(slot, &local_addr))?;
*guard = local_addr;
} else {
return Err(NetError::InvalidInput);
}
Ok(())
}
fn start_connect(&self, remote_addr: SocketAddrEx) -> NetResult<ConnectStatus> {
let remote_addr = remote_addr.into_unix()?;
let local_addr = self.local_addr.lock().clone();
let accept_poll = {
let mut guard = self.remote_addr.lock();
if !matches!(&*guard, UnixSocketAddr::Unnamed) {
return Err(NetError::InvalidInput);
}
let accept_poll = with_slot(&remote_addr, |slot| {
self.transport.connect(slot, &local_addr)
})?;
*guard = remote_addr;
accept_poll
};
self.transport.finish_connect(accept_poll);
Ok(ConnectStatus::Connected)
}
fn listen(&self, _backlog: usize) -> NetResult {
self.transport.listen()
}
fn is_listening(&self) -> bool {
self.transport.is_listening()
}
fn try_accept(&self) -> NetResult<Socket> {
let (transport, peer_addr) = self.transport.try_accept()?;
Ok(Self {
transport,
local_addr: SpinLock::new(self.local_addr.lock().clone()),
remote_addr: SpinLock::new(peer_addr),
}
.into())
}
fn try_send(&self, src: impl Read + IoBuf, options: &mut SendOptions) -> NetResult<usize> {
self.transport.try_send(src, options)
}
fn try_recv(&self, dst: impl Write, options: &mut RecvOptions<'_>) -> NetResult<usize> {
self.transport.try_recv(dst, options)
}
fn local_addr(&self) -> NetResult<SocketAddrEx> {
Ok(SocketAddrEx::Unix(self.local_addr.lock().clone()))
}
fn peer_addr(&self) -> NetResult<SocketAddrEx> {
Ok(SocketAddrEx::Unix(self.remote_addr.lock().clone()))
}
fn shutdown(&self, how: Shutdown) -> NetResult {
self.transport.shutdown(how)
}
}
impl Pollable for UnixSocket {
fn poll(&self) -> IoEvents {
self.transport.poll()
}
unsafe fn register_shared(&self, sink: &mut dyn SharedRegistrationSink, events: IoEvents) {
unsafe { self.transport.register_shared(sink, events) };
}
unsafe fn register_exclusive(
&self,
sink: &mut dyn ExclusiveRegistrationSink,
events: IoEvents,
) {
unsafe { self.transport.register_exclusive(sink, events) };
}
}