use alloc::{sync::Arc, vec, vec::Vec};
use core::{
net::{Ipv4Addr, SocketAddr},
sync::atomic::{AtomicBool, AtomicI32, AtomicU32, Ordering},
task::Waker,
};
use ax_io::prelude::*;
use ax_lazyinit::LazyLock;
use ax_sync::SpinLock;
use axpoll::{ExclusiveRegistrationSink, IoEvents, Pollable, SharedRegistrationSink};
use axpoll_set::PollSet;
use hashbrown::HashMap;
use smoltcp::{
iface::SocketHandle,
socket::tcp as smol,
time::Duration,
wire::{IpEndpoint, IpListenEndpoint, IpProtocol},
};
use crate::{
ConnectStatus, LISTEN_TABLE, NetError, NetResult, ReadinessVersion, RecvFlags, RecvOptions,
SOCKET_SET, SendOptions, Shutdown, Socket, SocketAddrEx, SocketDeferPollWake, SocketOps,
addr::{allocate_ephemeral_port, listen_addrs_conflict},
config::{DeviceBinding, InterfaceId},
consts::{TCP_RX_BUF_LEN, TCP_TX_BUF_LEN},
general::GeneralOptions,
get_control, get_service, interface_by_id,
ip_tos::{EgressIpTosKey, clear_egress_ip_tos, set_egress_ip_tos},
options::{
Configurable, GetSocketOption, SetSocketOption, TcpCongestionControl, TcpInfo,
TcpInfoOptions, TcpState,
},
receive_starts_next_edge, request_poll,
state::*,
};
const TCP_KEEPIDLE_DEFAULT_SECS: u32 = 7200;
const TCP_KEEPINTVL_DEFAULT_SECS: u32 = 75;
const TCP_KEEPCNT_DEFAULT: u32 = 9;
const TCP_USER_TIMEOUT_DEFAULT_MS: u32 = 0;
const TCP_KEEPIDLE_MAX_SECS: u32 = 32767;
const TCP_KEEPINTVL_MAX_SECS: u32 = 32767;
const TCP_KEEPCNT_MAX: u32 = 127;
const TCP_INFO_DEFAULT_MSS: u32 = 1460;
const TCP_INFO_DEFAULT_PMTU: u32 = 1500;
const TCP_INFO_INITIAL_RTO_MICROS: u32 = 1_000_000;
const TCP_INFO_DEFAULT_REORDERING: u32 = 3;
pub struct TcpSocket {
state: StateLock,
handle: SocketHandle,
bound_endpoint: SpinLock<IpListenEndpoint>,
peer_endpoint: SpinLock<Option<IpEndpoint>>,
tos_key: SpinLock<Option<EgressIpTosKey>>,
bound_registered: AtomicBool,
general: GeneralOptions,
pending_error: AtomicI32,
keep_idle_secs: AtomicU32,
keep_interval_secs: AtomicU32,
keep_count: AtomicU32,
user_timeout_millis: AtomicU32,
rx_closed: AtomicBool,
poll_rx: Arc<PollSet>,
poll_tx: Arc<PollSet>,
poll_rx_closed: PollSet,
readiness_version: ReadinessVersion,
}
unsafe impl Sync for TcpSocket {}
impl TcpSocket {
pub fn new() -> Self {
Self {
state: StateLock::new(State::Idle),
handle: SOCKET_SET.add(smol::Socket::new(
smol::SocketBuffer::new(vec![0; TCP_RX_BUF_LEN]),
smol::SocketBuffer::new(vec![0; TCP_TX_BUF_LEN]),
)),
bound_endpoint: SpinLock::new(empty_endpoint()),
peer_endpoint: SpinLock::new(None),
tos_key: SpinLock::new(None),
bound_registered: AtomicBool::new(false),
general: GeneralOptions::new(1, 2, 6), pending_error: AtomicI32::new(0),
keep_idle_secs: AtomicU32::new(TCP_KEEPIDLE_DEFAULT_SECS),
keep_interval_secs: AtomicU32::new(TCP_KEEPINTVL_DEFAULT_SECS),
keep_count: AtomicU32::new(TCP_KEEPCNT_DEFAULT),
user_timeout_millis: AtomicU32::new(TCP_USER_TIMEOUT_DEFAULT_MS),
rx_closed: AtomicBool::new(false),
poll_rx: Arc::new(PollSet::new()),
poll_tx: Arc::new(PollSet::new()),
poll_rx_closed: PollSet::new(),
readiness_version: ReadinessVersion::new(),
}
}
pub fn bind_device(&self, interface_id: InterfaceId) -> NetResult {
if interface_by_id(interface_id).is_none() {
return Err(NetError::NoSuchDevice);
}
self.general.set_device_binding(DeviceBinding {
bound_if: Some(interface_id),
});
Ok(())
}
fn new_connected(
handle: SocketHandle,
local_endpoint: IpEndpoint,
remote_endpoint: IpEndpoint,
) -> Self {
let result = Self {
state: StateLock::new(State::Connected),
handle,
bound_endpoint: SpinLock::new(empty_endpoint()),
peer_endpoint: SpinLock::new(Some(remote_endpoint)),
tos_key: SpinLock::new(None),
bound_registered: AtomicBool::new(false),
general: GeneralOptions::new(1, 2, 6), pending_error: AtomicI32::new(0),
keep_idle_secs: AtomicU32::new(TCP_KEEPIDLE_DEFAULT_SECS),
keep_interval_secs: AtomicU32::new(TCP_KEEPINTVL_DEFAULT_SECS),
keep_count: AtomicU32::new(TCP_KEEPCNT_DEFAULT),
user_timeout_millis: AtomicU32::new(TCP_USER_TIMEOUT_DEFAULT_MS),
rx_closed: AtomicBool::new(false),
poll_rx: Arc::new(PollSet::new()),
poll_tx: Arc::new(PollSet::new()),
poll_rx_closed: PollSet::new(),
readiness_version: ReadinessVersion::new(),
};
let endpoint = IpListenEndpoint {
addr: Some(local_endpoint.addr),
port: local_endpoint.port,
};
*result.bound_endpoint.lock() = endpoint;
result.general.set_device_binding(
get_control()
.local_binding_for(&endpoint)
.unwrap_or_default(),
);
result
}
pub fn readiness_version(&self) -> u64 {
self.readiness_version.current()
}
}
impl Default for TcpSocket {
fn default() -> Self {
Self::new()
}
}
impl TcpSocket {
fn state(&self) -> State {
self.state.get()
}
#[inline]
fn is_listening(&self) -> bool {
self.state() == State::Listening
}
fn with_smol_socket<R>(&self, f: impl FnOnce(&mut smol::Socket) -> R) -> R {
SOCKET_SET.with_socket_mut::<smol::Socket, _, _>(self.handle, f)
}
fn egress_ip_tos_key(&self) -> Option<EgressIpTosKey> {
if self.is_listening() {
return EgressIpTosKey::listener(IpProtocol::Tcp, *self.bound_endpoint.lock());
}
let local = self
.with_smol_socket(|socket| socket.local_endpoint())
.or_else(|| {
let endpoint = *self.bound_endpoint.lock();
endpoint.addr.map(|addr| IpEndpoint {
addr,
port: endpoint.port,
})
});
let remote = self
.with_smol_socket(|socket| socket.remote_endpoint())
.or_else(|| *self.peer_endpoint.lock());
EgressIpTosKey::exact(IpProtocol::Tcp, local?, remote?)
}
fn sync_egress_ip_tos(&self) {
let key = self.egress_ip_tos_key();
let tos = self.general.ip_tos();
let mut tracked = self.tos_key.lock();
if *tracked != key {
if let Some(old) = *tracked {
clear_egress_ip_tos(old);
}
*tracked = key;
}
if let Some(key) = key {
set_egress_ip_tos(key, tos);
}
}
fn clear_tracked_egress_ip_tos(&self) {
if let Some(key) = self.tos_key.lock().take() {
clear_egress_ip_tos(key);
}
}
fn tcp_info_snapshot(&self) -> TcpInfo {
self.with_smol_socket(|socket| {
let send_queue = socket.send_queue().min(u32::MAX as usize) as u32;
let snd_mss = TCP_INFO_DEFAULT_MSS;
let mut options = TcpInfoOptions::empty();
if socket.timestamp_enabled() {
options |= TcpInfoOptions::TIMESTAMPS;
}
TcpInfo {
state: tcp_state_info(socket.state()),
options,
rto_micros: socket
.timeout()
.map(duration_micros_u32)
.unwrap_or(TCP_INFO_INITIAL_RTO_MICROS),
ato_micros: socket.ack_delay().map(duration_micros_u32).unwrap_or(0),
snd_mss,
rcv_mss: snd_mss,
notsent_bytes: send_queue,
pmtu: TCP_INFO_DEFAULT_PMTU,
advmss: snd_mss,
reordering: TCP_INFO_DEFAULT_REORDERING,
snd_wnd: 0,
..Default::default()
}
})
}
fn bound_endpoint(&self) -> NetResult<IpListenEndpoint> {
let endpoint = *self.bound_endpoint.lock();
if endpoint.port == 0 {
return Err(NetError::InvalidInput);
}
Ok(endpoint)
}
fn poll_connect(&self) -> IoEvents {
let mut events = IoEvents::empty();
self.with_smol_socket(|socket| match socket.state() {
smol::State::SynSent | smol::State::SynReceived => {
}
smol::State::Established => {
self.pending_error.store(0, Ordering::Release);
self.state.set(State::Connected); *self.peer_endpoint.lock() = socket.remote_endpoint();
debug!(
"TCP socket {}: connected to {}",
self.handle,
socket.remote_endpoint().unwrap(),
);
events.set(IoEvents::OUT, true);
}
state => {
*self.peer_endpoint.lock() = None;
self.pending_error
.store(syscalls::Errno::ECONNREFUSED.into_raw(), Ordering::Release);
self.state.set(State::Closed); debug!(
"TCP socket {}: connect failed in state {:?}",
self.handle, state
);
events.set(IoEvents::OUT, true);
events.set(IoEvents::ERR, true);
events.set(IoEvents::HUP, true);
}
});
events
}
fn poll_stream(&self) -> IoEvents {
let mut events = IoEvents::empty();
self.with_smol_socket(|socket| {
events.set(
IoEvents::IN,
!self.rx_closed.load(Ordering::Acquire)
&& (!socket.may_recv() || socket.can_recv()),
);
events.set(IoEvents::OUT, !socket.may_send() || socket.can_send());
});
events
}
fn poll_listener(&self) -> IoEvents {
let mut events = IoEvents::empty();
let endpoint = self.bound_endpoint().unwrap();
let sockets = SOCKET_SET.inner.lock();
events.set(
IoEvents::IN,
LISTEN_TABLE.can_accept(endpoint, &sockets).unwrap(),
);
events
}
}
impl Configurable for TcpSocket {
fn get_option_inner(&self, option: &mut GetSocketOption) -> NetResult<bool> {
use GetSocketOption as O;
if let O::Error(error) = option {
**error = self.pending_error.swap(0, Ordering::AcqRel);
return Ok(true);
}
if self.general.get_option_inner(option)? {
return Ok(true);
}
match option {
O::NoDelay(no_delay) => {
**no_delay = self.with_smol_socket(|socket| !socket.nagle_enabled());
}
O::KeepAlive(keep_alive) => {
**keep_alive = self.with_smol_socket(|socket| socket.keep_alive().is_some());
}
O::MaxSegment(max_segment) => {
**max_segment = 1460;
}
O::TcpKeepIdle(keep_idle) => {
**keep_idle = self.keep_idle_secs.load(Ordering::Relaxed);
}
O::TcpKeepInterval(keep_interval) => {
**keep_interval = self.keep_interval_secs.load(Ordering::Relaxed);
}
O::TcpKeepCount(keep_count) => {
**keep_count = self.keep_count.load(Ordering::Relaxed);
}
O::TcpUserTimeout(user_timeout) => {
**user_timeout = self.user_timeout_millis.load(Ordering::Relaxed);
}
O::SendBuffer(size) => {
**size = TCP_TX_BUF_LEN;
}
O::ReceiveBuffer(size) => {
**size = TCP_RX_BUF_LEN;
}
O::TcpInfo(info) => {
**info = self.tcp_info_snapshot();
}
O::TcpCongestionControl(congestion_control) => {
**congestion_control =
self.with_smol_socket(|socket| match socket.congestion_control() {
smol::CongestionControl::None => TcpCongestionControl::None,
});
}
_ => return Ok(false),
}
Ok(true)
}
fn set_option_inner(&self, option: SetSocketOption) -> NetResult<bool> {
use SetSocketOption as O;
if let O::IpTos(tos) = option {
self.general.set_ip_tos(*tos);
self.sync_egress_ip_tos();
return Ok(true);
}
if self.general.set_option_inner(option)? {
return Ok(true);
}
match option {
O::NoDelay(no_delay) => {
self.with_smol_socket(|socket| {
socket.set_nagle_enabled(!no_delay);
});
}
O::KeepAlive(keep_alive) => {
let interval =
Duration::from_secs(self.keep_idle_secs.load(Ordering::Relaxed) as u64);
self.with_smol_socket(|socket| {
socket.set_keep_alive(keep_alive.then_some(interval));
});
}
O::TcpKeepIdle(keep_idle) => {
if *keep_idle == 0 || *keep_idle > TCP_KEEPIDLE_MAX_SECS {
return Err(NetError::InvalidInput);
}
self.keep_idle_secs.store(*keep_idle, Ordering::Relaxed);
let interval = Duration::from_secs(*keep_idle as u64);
self.with_smol_socket(|socket| {
if socket.keep_alive().is_some() {
socket.set_keep_alive(Some(interval));
}
});
}
O::TcpKeepInterval(keep_interval) => {
if *keep_interval == 0 || *keep_interval > TCP_KEEPINTVL_MAX_SECS {
return Err(NetError::InvalidInput);
}
self.keep_interval_secs
.store(*keep_interval, Ordering::Relaxed);
}
O::TcpKeepCount(keep_count) => {
if *keep_count == 0 || *keep_count > TCP_KEEPCNT_MAX {
return Err(NetError::InvalidInput);
}
self.keep_count.store(*keep_count, Ordering::Relaxed);
}
O::TcpUserTimeout(user_timeout) => {
self.user_timeout_millis
.store(*user_timeout, Ordering::Relaxed);
}
O::TcpCongestionControl(congestion_control) => {
self.with_smol_socket(|socket| match congestion_control {
TcpCongestionControl::None => {
socket.set_congestion_control(smol::CongestionControl::None);
}
});
}
_ => return Ok(false),
}
Ok(true)
}
}
impl SocketOps for TcpSocket {
fn bind(&self, local_addr: SocketAddrEx) -> NetResult {
let mut local_addr = local_addr.into_ip()?;
self.state
.lock(State::Idle)
.map_err(|_| NetError::InvalidInput)?
.transit(State::Idle, || {
if local_addr.port() == 0 {
local_addr.set_port(get_ephemeral_port()?);
}
if self.bound_endpoint.lock().port != 0 {
return Err(NetError::InvalidInput);
}
let endpoint = IpListenEndpoint {
addr: if local_addr.ip().is_unspecified() {
None
} else {
Some(local_addr.ip().into())
},
port: local_addr.port(),
};
if !self.general.reuse_address()
&& !self.general.reuse_port()
&& !LISTEN_TABLE.can_listen(endpoint)
{
return Err(NetError::AddrInUse);
}
let binding = get_control().local_binding_for(&endpoint)?;
self.register_bound_endpoint(endpoint)?;
*self.bound_endpoint.lock() = endpoint;
if binding.bound_if.is_some() {
self.general.set_device_binding(binding);
}
debug!("TCP socket {}: binding to {}", self.handle, local_addr);
Ok(())
})
}
fn start_connect(&self, remote_addr: SocketAddrEx) -> NetResult<ConnectStatus> {
let remote_addr = remote_addr.into_ip()?;
self.begin_connect(remote_addr)?;
request_poll();
Ok(ConnectStatus::InProgress)
}
fn connect_status(&self) -> NetResult<ConnectStatus> {
match self.state.get() {
State::Connected => return Ok(ConnectStatus::Connected),
State::Connecting => {}
State::Closed => return Err(NetError::ConnectionRefused),
_ => return Err(NetError::InvalidInput),
}
request_poll();
let events = self.poll_connect();
if !events.contains(IoEvents::OUT) {
Ok(ConnectStatus::InProgress)
} else if self.state.get() == State::Connected {
Ok(ConnectStatus::Connected)
} else {
Err(NetError::ConnectionRefused)
}
}
fn listen(&self, backlog: usize) -> NetResult {
if let Ok(guard) = self.state.lock(State::Idle) {
guard.transit(State::Listening, || {
let mut bound_endpoint = *self.bound_endpoint.lock();
if bound_endpoint.port == 0 {
bound_endpoint.port = get_ephemeral_port()?;
}
let binding = get_control().local_binding_for(&bound_endpoint)?;
self.with_bound_endpoint_registered(bound_endpoint, || {
LISTEN_TABLE.listen(bound_endpoint, backlog, self.general.reuse_port())
})?;
*self.bound_endpoint.lock() = bound_endpoint;
self.sync_egress_ip_tos();
if binding.bound_if.is_some() {
self.general.set_device_binding(binding);
}
debug!("listening on {}", bound_endpoint);
Ok(())
})?;
} else {
}
Ok(())
}
fn is_listening(&self) -> bool {
self.state.get() == State::Listening
}
fn try_accept(&self) -> NetResult<Socket> {
if self.state.get() != State::Listening {
return Err(NetError::InvalidInput);
}
let bound_endpoint = self.bound_endpoint()?;
request_poll();
let accepted = {
let mut sockets = SOCKET_SET.inner.lock();
let accepted = LISTEN_TABLE.accept(bound_endpoint, &mut sockets)?;
if matches!(LISTEN_TABLE.can_accept(bound_endpoint, &sockets), Ok(false)) {
self.readiness_version.publish();
}
accepted
};
Ok({
let socket = TcpSocket::new_connected(
accepted.handle,
accepted.local_endpoint,
accepted.remote_endpoint,
);
socket.general.set_ip_tos(self.general.ip_tos());
socket.sync_egress_ip_tos();
debug!(
"accepted connection from {}, {}",
accepted.handle, accepted.remote_endpoint
);
socket.into()
})
}
fn try_send(&self, mut src: impl Read + IoBuf, _options: &mut SendOptions) -> NetResult<usize> {
if src.remaining() == 0 {
return Ok(0);
}
request_poll();
let result = self.with_smol_socket(|socket| {
if !socket.is_active() {
Err(NetError::NotConnected)
} else if !socket.can_send() {
Err(NetError::WouldBlock)
} else {
let len = socket
.send(|buffer| {
let result = src.read(buffer);
let len = result.unwrap_or(0);
(len, result)
})
.map_err(|_| NetError::NotConnected)??;
Ok(len)
}
});
if result.as_ref().is_ok_and(|sent| *sent > 0) {
request_poll();
}
result
}
fn try_recv(
&self,
mut dst: impl Write + IoBufMut,
options: &mut RecvOptions<'_>,
) -> NetResult<usize> {
if self.rx_closed.load(Ordering::Acquire) {
return Err(NetError::NotConnected);
}
if self.state.get() == State::Closed {
return Err(NetError::NotConnected);
}
request_poll();
self.with_smol_socket(|socket| {
if socket.recv_queue() > 0 {
if options.flags.contains(RecvFlags::PEEK) {
dst.write(
socket
.peek(dst.remaining_mut())
.map_err(|_| NetError::NotConnected)?,
)
.map_err(NetError::from)
} else {
let mut total = 0;
while socket.recv_queue() > 0 && dst.remaining_mut() > 0 {
let len = socket
.recv(|buf| {
let result = dst.write(buf).map_err(NetError::from);
let len = result.unwrap_or(0);
(len, result)
})
.map_err(|_| NetError::NotConnected)??;
if len == 0 {
break;
}
total += len;
}
if receive_starts_next_edge(total, socket.recv_queue()) {
self.readiness_version.publish();
}
Ok(total)
}
} else if !socket.may_recv() {
Ok(0)
} else {
Err(NetError::WouldBlock)
}
})
}
fn recv_available(&self) -> NetResult<usize> {
if self.state.get() == State::Listening {
return Err(NetError::InvalidInput);
}
let available = self.with_smol_socket(|socket| socket.recv_queue());
if available > 0 {
return Ok(available);
}
request_poll();
Ok(self.with_smol_socket(|socket| socket.recv_queue()))
}
fn local_addr(&self) -> NetResult<SocketAddrEx> {
let endpoint = self.with_smol_socket(|socket| {
socket
.local_endpoint()
.map(|endpoint| IpListenEndpoint {
addr: Some(endpoint.addr),
port: endpoint.port,
})
.unwrap_or_else(|| *self.bound_endpoint.lock())
});
Ok(SocketAddrEx::Ip(SocketAddr::new(
endpoint
.addr
.map_or_else(|| Ipv4Addr::UNSPECIFIED.into(), Into::into),
endpoint.port,
)))
}
fn peer_addr(&self) -> NetResult<SocketAddrEx> {
self.with_smol_socket(|socket| {
Ok(SocketAddrEx::Ip(
socket
.remote_endpoint()
.or_else(|| *self.peer_endpoint.lock())
.ok_or(NetError::NotConnected)?
.into(),
))
})
}
fn shutdown(&self, how: Shutdown) -> NetResult {
if how.has_read() {
self.rx_closed.store(true, Ordering::Release);
self.readiness_version.publish();
unsafe { self.poll_rx_closed.wake(IoEvents::RDHUP | IoEvents::IN) };
}
if let Ok(guard) = self.state.lock(State::Connected) {
if how.has_read() && how.has_write() {
guard.transit(State::Closed, || {
self.with_smol_socket(|socket| {
debug!("TCP socket {}: shutting down", self.handle);
socket.close();
});
self.clear_tracked_egress_ip_tos();
self.unregister_bound_endpoint();
*self.bound_endpoint.lock() = empty_endpoint();
request_poll();
Ok(())
})?;
} else if how.has_write() {
self.with_smol_socket(|socket| {
debug!("TCP socket {}: shutting down write side", self.handle);
socket.close();
});
request_poll();
}
}
if let Ok(guard) = self.state.lock(State::Listening) {
guard.transit(State::Closed, || {
LISTEN_TABLE.unlisten(self.bound_endpoint()?);
self.clear_tracked_egress_ip_tos();
self.unregister_bound_endpoint();
*self.bound_endpoint.lock() = empty_endpoint();
request_poll();
Ok(())
})?;
}
Ok(())
}
}
impl Pollable for TcpSocket {
fn poll(&self) -> IoEvents {
request_poll();
let mut events = match self.state.get() {
State::Connecting => self.poll_connect(),
State::Connected | State::Idle | State::Closed => self.poll_stream(),
State::Listening => self.poll_listener(),
State::Busy => IoEvents::empty(),
};
events.set(IoEvents::RDHUP, self.rx_closed.load(Ordering::Acquire));
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 TcpSocket {
fn register_poll_sources(
&self,
events: IoEvents,
mut register: impl FnMut(&PollSet, IoEvents),
) {
let mut accept_registration = None;
if self.state.get() == State::Listening && events.intersects(IoEvents::IN | IoEvents::RDHUP)
{
let port = self.bound_endpoint.lock().port;
if port != 0 {
let endpoint = *self.bound_endpoint.lock();
if let Some(accept_poll) = LISTEN_TABLE.accept_poll(endpoint) {
register(&accept_poll, IoEvents::IN);
let accept_waker = LISTEN_TABLE.accept_waker(accept_poll.clone());
accept_registration = Some((endpoint, accept_poll, accept_waker));
}
}
}
let recv_waker = if events.intersects(IoEvents::IN | IoEvents::RDHUP) {
register(&self.poll_rx, IoEvents::IN | IoEvents::RDHUP);
Some(Waker::from(Arc::new(SocketDeferPollWake::new(
self.poll_rx.clone(),
IoEvents::IN | IoEvents::RDHUP,
self.readiness_version.clone(),
))))
} else {
None
};
let send_waker = if events.contains(IoEvents::OUT) {
register(&self.poll_tx, IoEvents::OUT);
Some(Waker::from(Arc::new(SocketDeferPollWake::new(
self.poll_tx.clone(),
IoEvents::OUT,
self.readiness_version.clone(),
))))
} else {
None
};
if let Some((endpoint, accept_poll, accept_waker)) = accept_registration.as_ref() {
let mut sockets = SOCKET_SET.inner.lock();
LISTEN_TABLE.register_pending_accept_wakers(
*endpoint,
&mut sockets,
accept_poll,
accept_waker,
);
}
self.with_smol_socket(|socket| {
if let Some(waker) = recv_waker.as_ref() {
socket.register_recv_waker(waker);
}
if let Some(waker) = send_waker.as_ref() {
socket.register_send_waker(waker);
}
});
if events.intersects(IoEvents::IN | IoEvents::OUT | IoEvents::RDHUP) {
register(&self.poll_rx, events);
self.general
.register_waker(&Waker::from(Arc::new(SocketDeferPollWake::new(
self.poll_rx.clone(),
events,
self.readiness_version.clone(),
))));
}
if events.contains(IoEvents::RDHUP) {
register(&self.poll_rx_closed, IoEvents::RDHUP | IoEvents::IN);
}
}
}
impl Drop for TcpSocket {
fn drop(&mut self) {
let endpoint = *self.bound_endpoint.lock();
if self.state.get() == State::Listening && endpoint.port != 0 {
LISTEN_TABLE.unlisten(endpoint);
}
let should_orphan = self.with_smol_socket(|socket| {
let state = socket.state();
let should_orphan = matches!(
state,
smol::State::Established
| smol::State::CloseWait
| smol::State::FinWait1
| smol::State::FinWait2
| smol::State::Closing
| smol::State::LastAck
| smol::State::TimeWait
) || socket.send_queue() > 0;
if matches!(
state,
smol::State::Established
| smol::State::SynSent
| smol::State::SynReceived
| smol::State::CloseWait
| smol::State::FinWait1
| smol::State::FinWait2
| smol::State::Closing
| smol::State::LastAck
) {
debug!("TCP socket {}: closing on drop", self.handle);
socket.close();
}
should_orphan
});
self.clear_tracked_egress_ip_tos();
self.unregister_bound_endpoint();
if should_orphan {
let timestamp = smoltcp::time::Instant::from_micros_const(
(ax_hal::time::monotonic_time_nanos() / 1_000) as i64,
);
crate::orphan::add_orphan(self.handle, timestamp);
} else {
SOCKET_SET.remove(self.handle);
}
crate::request_poll();
}
}
fn duration_micros_u32(value: Duration) -> u32 {
value.total_micros().min(u32::MAX as u64) as u32
}
fn tcp_state_info(state: smol::State) -> TcpState {
match state {
smol::State::Closed => TcpState::Closed,
smol::State::Listen => TcpState::Listen,
smol::State::SynSent => TcpState::SynSent,
smol::State::SynReceived => TcpState::SynReceived,
smol::State::Established => TcpState::Established,
smol::State::FinWait1 => TcpState::FinWait1,
smol::State::FinWait2 => TcpState::FinWait2,
smol::State::CloseWait => TcpState::CloseWait,
smol::State::Closing => TcpState::Closing,
smol::State::LastAck => TcpState::LastAck,
smol::State::TimeWait => TcpState::TimeWait,
}
}
const fn empty_endpoint() -> IpListenEndpoint {
IpListenEndpoint {
addr: None,
port: 0,
}
}
impl TcpSocket {
fn begin_connect(&self, remote_addr: SocketAddr) -> NetResult {
self.state
.lock(State::Idle)
.map_err(|state| {
if state == State::Connecting {
NetError::InProgress
} else {
NetError::AlreadyConnected
}
})?
.transit(State::Connecting, || {
self.pending_error.store(0, Ordering::Release);
let remote_endpoint = IpEndpoint::from(remote_addr);
let mut bound_endpoint = *self.bound_endpoint.lock();
let was_unbound_or_unspecified =
bound_endpoint.addr.is_none_or(|addr| addr.is_unspecified());
let had_explicit_device_binding = self.general.device_binding().bound_if.is_some();
if bound_endpoint.addr.is_none_or(|addr| addr.is_unspecified()) {
bound_endpoint.addr = Some(
get_control()
.select_route_with_binding(
&remote_endpoint.addr,
self.general.device_binding(),
)?
.source,
);
}
if bound_endpoint.port == 0 {
bound_endpoint.port = get_ephemeral_port()?;
}
info!(
"TCP connection from {} to {}",
bound_endpoint, remote_endpoint
);
self.with_bound_endpoint_registered(bound_endpoint, || {
let mut service = get_service();
let context = service.iface.context();
self.with_smol_socket(|socket| {
socket
.connect(context, remote_endpoint, bound_endpoint)
.map_err(|e| match e {
smol::ConnectError::InvalidState => NetError::AlreadyConnected,
smol::ConnectError::Unaddressable => NetError::ConnectionRefused,
})?;
Ok::<(), NetError>(())
})
})?;
*self.bound_endpoint.lock() = bound_endpoint;
if !had_explicit_device_binding && was_unbound_or_unspecified {
self.general
.set_device_binding(get_control().local_binding_for(&bound_endpoint)?);
}
self.sync_egress_ip_tos();
Ok(())
})
}
fn register_bound_endpoint(&self, endpoint: IpListenEndpoint) -> NetResult {
if !self.bound_registered.load(Ordering::Acquire) {
register_tcp_bound(endpoint, self.general.reuse_port())?;
self.bound_registered.store(true, Ordering::Release);
}
Ok(())
}
fn with_bound_endpoint_registered<R>(
&self,
endpoint: IpListenEndpoint,
f: impl FnOnce() -> NetResult<R>,
) -> NetResult<R> {
let register_bound = !self.bound_registered.load(Ordering::Acquire);
if register_bound {
register_tcp_bound(endpoint, self.general.reuse_port())?;
}
match f() {
Ok(value) => {
if register_bound {
self.bound_registered.store(true, Ordering::Release);
}
Ok(value)
}
Err(err) => {
if register_bound {
unregister_tcp_bound(endpoint);
}
Err(err)
}
}
}
fn unregister_bound_endpoint(&self) {
if self.bound_registered.swap(false, Ordering::AcqRel) {
unregister_tcp_bound(*self.bound_endpoint.lock());
}
}
}
struct TcpBoundEntry {
addr: Option<smoltcp::wire::IpAddress>,
reuse_port: bool,
}
static TCP_BOUND_PORTS: LazyLock<SpinLock<HashMap<u16, Vec<TcpBoundEntry>>>> =
LazyLock::new(|| SpinLock::new(HashMap::new()));
fn register_tcp_bound(endpoint: IpListenEndpoint, reuse_port: bool) -> NetResult {
if endpoint.port == 0 {
return Ok(());
}
let mut bound_ports = TCP_BOUND_PORTS.lock();
let entries = bound_ports.entry(endpoint.port).or_default();
for entry in entries.iter() {
if listen_addrs_conflict(entry.addr, endpoint.addr)
&& !(reuse_port && entry.reuse_port && entry.addr == endpoint.addr)
{
return Err(NetError::AddrInUse);
}
}
entries.push(TcpBoundEntry {
addr: endpoint.addr,
reuse_port,
});
Ok(())
}
fn unregister_tcp_bound(endpoint: IpListenEndpoint) {
if endpoint.port == 0 {
return;
}
let mut bound_ports = TCP_BOUND_PORTS.lock();
let Some(entries) = bound_ports.get_mut(&endpoint.port) else {
return;
};
if let Some(index) = entries.iter().position(|entry| entry.addr == endpoint.addr) {
entries.swap_remove(index);
}
if entries.is_empty() {
bound_ports.remove(&endpoint.port);
}
}
fn tcp_port_available(port: u16) -> bool {
LISTEN_TABLE.can_listen(IpListenEndpoint { addr: None, port })
&& !TCP_BOUND_PORTS.lock().contains_key(&port)
}
fn get_ephemeral_port() -> NetResult<u16> {
allocate_ephemeral_port(tcp_port_available)
}