Skip to main content

ax_net/
tcp.rs

1//! TCP socket implementation.
2//!
3//! TCP sockets wrap smoltcp stream sockets with POSIX-like behavior: bind and
4//! listen bookkeeping, accept queues, nonblocking readiness, keepalive and
5//! TCP_INFO options, orphan cleanup, and route-aware device binding.
6//!
7//! # smoltcp Boundary
8//!
9//! The actual TCP state machine, retransmission timers, and stream buffers live
10//! in smoltcp. This module owns the public socket state around that core:
11//! ephemeral port allocation, wildcard/specific bind registration, listener
12//! setup, accepted child socket construction, shutdown semantics, and
13//! Linux-compatible error reporting.
14//!
15//! # Polling Model
16//!
17//! Socket methods never synchronously drive the full interface poll loop.
18//! Instead they mutate the smoltcp socket, call `request_poll()`, register
19//! wakers through `PollSet`, and let the unique protocol executor advance
20//! timers, handshakes, retransmission, and close states.
21//!
22//! # Related Side Tables
23//!
24//! - `TCP_BOUND_PORTS` records public bind ownership.
25//! - `LISTEN_TABLE` owns passive-open child sockets and accept wakeups.
26//! - `orphan` keeps dropped sockets alive long enough for FIN/TIME-WAIT cleanup.
27
28use alloc::{sync::Arc, vec, vec::Vec};
29use core::{
30    net::{Ipv4Addr, SocketAddr},
31    sync::atomic::{AtomicBool, AtomicI32, AtomicU32, Ordering},
32    task::Waker,
33};
34
35use ax_io::prelude::*;
36use ax_lazyinit::LazyLock;
37use ax_sync::SpinLock;
38use axpoll::{ExclusiveRegistrationSink, IoEvents, Pollable, SharedRegistrationSink};
39use axpoll_set::PollSet;
40use hashbrown::HashMap;
41use smoltcp::{
42    iface::SocketHandle,
43    socket::tcp as smol,
44    time::Duration,
45    wire::{IpEndpoint, IpListenEndpoint, IpProtocol},
46};
47
48use crate::{
49    ConnectStatus, LISTEN_TABLE, NetError, NetResult, ReadinessVersion, RecvFlags, RecvOptions,
50    SOCKET_SET, SendOptions, Shutdown, Socket, SocketAddrEx, SocketDeferPollWake, SocketOps,
51    addr::{allocate_ephemeral_port, listen_addrs_conflict},
52    config::{DeviceBinding, InterfaceId},
53    consts::{TCP_RX_BUF_LEN, TCP_TX_BUF_LEN},
54    general::GeneralOptions,
55    get_control, get_service, interface_by_id,
56    ip_tos::{EgressIpTosKey, clear_egress_ip_tos, set_egress_ip_tos},
57    options::{
58        Configurable, GetSocketOption, SetSocketOption, TcpCongestionControl, TcpInfo,
59        TcpInfoOptions, TcpState,
60    },
61    receive_starts_next_edge, request_poll,
62    state::*,
63};
64
65const TCP_KEEPIDLE_DEFAULT_SECS: u32 = 7200;
66const TCP_KEEPINTVL_DEFAULT_SECS: u32 = 75;
67const TCP_KEEPCNT_DEFAULT: u32 = 9;
68const TCP_USER_TIMEOUT_DEFAULT_MS: u32 = 0;
69const TCP_KEEPIDLE_MAX_SECS: u32 = 32767;
70const TCP_KEEPINTVL_MAX_SECS: u32 = 32767;
71const TCP_KEEPCNT_MAX: u32 = 127;
72const TCP_INFO_DEFAULT_MSS: u32 = 1460;
73const TCP_INFO_DEFAULT_PMTU: u32 = 1500;
74const TCP_INFO_INITIAL_RTO_MICROS: u32 = 1_000_000;
75const TCP_INFO_DEFAULT_REORDERING: u32 = 3;
76
77/// A TCP socket that provides POSIX-like APIs.
78pub struct TcpSocket {
79    /// Public high-level socket state gate.
80    state: StateLock,
81    /// Handle into the global smoltcp socket set.
82    handle: SocketHandle,
83    /// Bound listen endpoint, or an empty endpoint before bind/connect.
84    bound_endpoint: SpinLock<IpListenEndpoint>,
85    /// Connected peer endpoint once established.
86    peer_endpoint: SpinLock<Option<IpEndpoint>>,
87    /// Currently registered egress IP_TOS policy for this TCP socket.
88    tos_key: SpinLock<Option<EgressIpTosKey>>,
89    /// Whether `bound_endpoint` is registered in `TCP_BOUND_PORTS`.
90    bound_registered: AtomicBool,
91
92    /// Shared socket options and blocking helpers.
93    general: GeneralOptions,
94    /// Pending Linux errno-style connection error.
95    pending_error: AtomicI32,
96    /// TCP_KEEPIDLE value in seconds.
97    keep_idle_secs: AtomicU32,
98    /// TCP_KEEPINTVL value in seconds.
99    keep_interval_secs: AtomicU32,
100    /// TCP_KEEPCNT value.
101    keep_count: AtomicU32,
102    /// TCP_USER_TIMEOUT value in milliseconds.
103    user_timeout_millis: AtomicU32,
104    /// Whether the read half was shut down from the public API.
105    rx_closed: AtomicBool,
106    /// Shared RX readiness poll set.
107    poll_rx: Arc<PollSet>,
108    /// Shared TX readiness poll set.
109    poll_tx: Arc<PollSet>,
110    /// Wakes waiters when the receive side becomes closed.
111    poll_rx_closed: PollSet,
112    /// Generation published for each socket readiness wake.
113    readiness_version: ReadinessVersion,
114}
115
116unsafe impl Sync for TcpSocket {}
117
118impl TcpSocket {
119    /// Creates a new TCP socket.
120    pub fn new() -> Self {
121        Self {
122            state: StateLock::new(State::Idle),
123            handle: SOCKET_SET.add(smol::Socket::new(
124                smol::SocketBuffer::new(vec![0; TCP_RX_BUF_LEN]),
125                smol::SocketBuffer::new(vec![0; TCP_TX_BUF_LEN]),
126            )),
127            bound_endpoint: SpinLock::new(empty_endpoint()),
128            peer_endpoint: SpinLock::new(None),
129            tos_key: SpinLock::new(None),
130            bound_registered: AtomicBool::new(false),
131
132            general: GeneralOptions::new(1, 2, 6), // SOCK_STREAM
133            pending_error: AtomicI32::new(0),
134            keep_idle_secs: AtomicU32::new(TCP_KEEPIDLE_DEFAULT_SECS),
135            keep_interval_secs: AtomicU32::new(TCP_KEEPINTVL_DEFAULT_SECS),
136            keep_count: AtomicU32::new(TCP_KEEPCNT_DEFAULT),
137            user_timeout_millis: AtomicU32::new(TCP_USER_TIMEOUT_DEFAULT_MS),
138            rx_closed: AtomicBool::new(false),
139            poll_rx: Arc::new(PollSet::new()),
140            poll_tx: Arc::new(PollSet::new()),
141            poll_rx_closed: PollSet::new(),
142            readiness_version: ReadinessVersion::new(),
143        }
144    }
145
146    /// Restricts this socket to one interface for route selection.
147    pub fn bind_device(&self, interface_id: InterfaceId) -> NetResult {
148        if interface_by_id(interface_id).is_none() {
149            return Err(NetError::NoSuchDevice);
150        }
151        self.general.set_device_binding(DeviceBinding {
152            bound_if: Some(interface_id),
153        });
154        Ok(())
155    }
156
157    /// Creates a new TCP socket that is already connected.
158    fn new_connected(
159        handle: SocketHandle,
160        local_endpoint: IpEndpoint,
161        remote_endpoint: IpEndpoint,
162    ) -> Self {
163        let result = Self {
164            state: StateLock::new(State::Connected),
165            handle,
166            bound_endpoint: SpinLock::new(empty_endpoint()),
167            peer_endpoint: SpinLock::new(Some(remote_endpoint)),
168            tos_key: SpinLock::new(None),
169            bound_registered: AtomicBool::new(false),
170
171            general: GeneralOptions::new(1, 2, 6), // SOCK_STREAM
172            pending_error: AtomicI32::new(0),
173            keep_idle_secs: AtomicU32::new(TCP_KEEPIDLE_DEFAULT_SECS),
174            keep_interval_secs: AtomicU32::new(TCP_KEEPINTVL_DEFAULT_SECS),
175            keep_count: AtomicU32::new(TCP_KEEPCNT_DEFAULT),
176            user_timeout_millis: AtomicU32::new(TCP_USER_TIMEOUT_DEFAULT_MS),
177            rx_closed: AtomicBool::new(false),
178            poll_rx: Arc::new(PollSet::new()),
179            poll_tx: Arc::new(PollSet::new()),
180            poll_rx_closed: PollSet::new(),
181            readiness_version: ReadinessVersion::new(),
182        };
183        let endpoint = IpListenEndpoint {
184            addr: Some(local_endpoint.addr),
185            port: local_endpoint.port,
186        };
187        *result.bound_endpoint.lock() = endpoint;
188        result.general.set_device_binding(
189            get_control()
190                .local_binding_for(&endpoint)
191                .unwrap_or_default(),
192        );
193        result
194    }
195
196    /// Returns the latest readiness wake generation for edge-triggered pollers.
197    pub fn readiness_version(&self) -> u64 {
198        self.readiness_version.current()
199    }
200}
201
202impl Default for TcpSocket {
203    fn default() -> Self {
204        Self::new()
205    }
206}
207
208/// Private methods
209impl TcpSocket {
210    fn state(&self) -> State {
211        self.state.get()
212    }
213
214    #[inline]
215    fn is_listening(&self) -> bool {
216        self.state() == State::Listening
217    }
218
219    fn with_smol_socket<R>(&self, f: impl FnOnce(&mut smol::Socket) -> R) -> R {
220        SOCKET_SET.with_socket_mut::<smol::Socket, _, _>(self.handle, f)
221    }
222
223    fn egress_ip_tos_key(&self) -> Option<EgressIpTosKey> {
224        if self.is_listening() {
225            return EgressIpTosKey::listener(IpProtocol::Tcp, *self.bound_endpoint.lock());
226        }
227
228        let local = self
229            .with_smol_socket(|socket| socket.local_endpoint())
230            .or_else(|| {
231                let endpoint = *self.bound_endpoint.lock();
232                endpoint.addr.map(|addr| IpEndpoint {
233                    addr,
234                    port: endpoint.port,
235                })
236            });
237        let remote = self
238            .with_smol_socket(|socket| socket.remote_endpoint())
239            .or_else(|| *self.peer_endpoint.lock());
240
241        EgressIpTosKey::exact(IpProtocol::Tcp, local?, remote?)
242    }
243
244    fn sync_egress_ip_tos(&self) {
245        let key = self.egress_ip_tos_key();
246        let tos = self.general.ip_tos();
247        let mut tracked = self.tos_key.lock();
248        if *tracked != key {
249            if let Some(old) = *tracked {
250                clear_egress_ip_tos(old);
251            }
252            *tracked = key;
253        }
254        if let Some(key) = key {
255            set_egress_ip_tos(key, tos);
256        }
257    }
258
259    fn clear_tracked_egress_ip_tos(&self) {
260        if let Some(key) = self.tos_key.lock().take() {
261            clear_egress_ip_tos(key);
262        }
263    }
264
265    fn tcp_info_snapshot(&self) -> TcpInfo {
266        self.with_smol_socket(|socket| {
267            let send_queue = socket.send_queue().min(u32::MAX as usize) as u32;
268            let snd_mss = TCP_INFO_DEFAULT_MSS;
269
270            let mut options = TcpInfoOptions::empty();
271            if socket.timestamp_enabled() {
272                options |= TcpInfoOptions::TIMESTAMPS;
273            }
274
275            TcpInfo {
276                state: tcp_state_info(socket.state()),
277                options,
278                rto_micros: socket
279                    .timeout()
280                    .map(duration_micros_u32)
281                    .unwrap_or(TCP_INFO_INITIAL_RTO_MICROS),
282                ato_micros: socket.ack_delay().map(duration_micros_u32).unwrap_or(0),
283                snd_mss,
284                rcv_mss: snd_mss,
285                notsent_bytes: send_queue,
286                pmtu: TCP_INFO_DEFAULT_PMTU,
287                advmss: snd_mss,
288                reordering: TCP_INFO_DEFAULT_REORDERING,
289                snd_wnd: 0,
290                ..Default::default()
291            }
292        })
293    }
294
295    fn bound_endpoint(&self) -> NetResult<IpListenEndpoint> {
296        let endpoint = *self.bound_endpoint.lock();
297        if endpoint.port == 0 {
298            return Err(NetError::InvalidInput);
299        }
300        Ok(endpoint)
301    }
302
303    fn poll_connect(&self) -> IoEvents {
304        let mut events = IoEvents::empty();
305        self.with_smol_socket(|socket| match socket.state() {
306            smol::State::SynSent | smol::State::SynReceived => {
307                // wait for connection
308            }
309            smol::State::Established => {
310                self.pending_error.store(0, Ordering::Release);
311                self.state.set(State::Connected); // connected
312                *self.peer_endpoint.lock() = socket.remote_endpoint();
313                debug!(
314                    "TCP socket {}: connected to {}",
315                    self.handle,
316                    socket.remote_endpoint().unwrap(),
317                );
318                events.set(IoEvents::OUT, true);
319            }
320            state => {
321                *self.peer_endpoint.lock() = None;
322                self.pending_error
323                    .store(syscalls::Errno::ECONNREFUSED.into_raw(), Ordering::Release);
324                self.state.set(State::Closed); // connection failed
325                debug!(
326                    "TCP socket {}: connect failed in state {:?}",
327                    self.handle, state
328                );
329                events.set(IoEvents::OUT, true);
330                events.set(IoEvents::ERR, true);
331                events.set(IoEvents::HUP, true);
332            }
333        });
334        events
335    }
336
337    fn poll_stream(&self) -> IoEvents {
338        let mut events = IoEvents::empty();
339        self.with_smol_socket(|socket| {
340            events.set(
341                IoEvents::IN,
342                !self.rx_closed.load(Ordering::Acquire)
343                    && (!socket.may_recv() || socket.can_recv()),
344            );
345            events.set(IoEvents::OUT, !socket.may_send() || socket.can_send());
346        });
347        events
348    }
349
350    fn poll_listener(&self) -> IoEvents {
351        let mut events = IoEvents::empty();
352        let endpoint = self.bound_endpoint().unwrap();
353        let sockets = SOCKET_SET.inner.lock();
354        events.set(
355            IoEvents::IN,
356            LISTEN_TABLE.can_accept(endpoint, &sockets).unwrap(),
357        );
358        events
359    }
360}
361
362impl Configurable for TcpSocket {
363    fn get_option_inner(&self, option: &mut GetSocketOption) -> NetResult<bool> {
364        use GetSocketOption as O;
365
366        if let O::Error(error) = option {
367            **error = self.pending_error.swap(0, Ordering::AcqRel);
368            return Ok(true);
369        }
370
371        if self.general.get_option_inner(option)? {
372            return Ok(true);
373        }
374
375        match option {
376            O::NoDelay(no_delay) => {
377                **no_delay = self.with_smol_socket(|socket| !socket.nagle_enabled());
378            }
379            O::KeepAlive(keep_alive) => {
380                **keep_alive = self.with_smol_socket(|socket| socket.keep_alive().is_some());
381            }
382            O::MaxSegment(max_segment) => {
383                // TODO(mivik): get actual MSS
384                **max_segment = 1460;
385            }
386            O::TcpKeepIdle(keep_idle) => {
387                **keep_idle = self.keep_idle_secs.load(Ordering::Relaxed);
388            }
389            O::TcpKeepInterval(keep_interval) => {
390                **keep_interval = self.keep_interval_secs.load(Ordering::Relaxed);
391            }
392            O::TcpKeepCount(keep_count) => {
393                **keep_count = self.keep_count.load(Ordering::Relaxed);
394            }
395            O::TcpUserTimeout(user_timeout) => {
396                **user_timeout = self.user_timeout_millis.load(Ordering::Relaxed);
397            }
398            O::SendBuffer(size) => {
399                **size = TCP_TX_BUF_LEN;
400            }
401            O::ReceiveBuffer(size) => {
402                **size = TCP_RX_BUF_LEN;
403            }
404            O::TcpInfo(info) => {
405                **info = self.tcp_info_snapshot();
406            }
407            O::TcpCongestionControl(congestion_control) => {
408                **congestion_control =
409                    self.with_smol_socket(|socket| match socket.congestion_control() {
410                        smol::CongestionControl::None => TcpCongestionControl::None,
411                    });
412            }
413            _ => return Ok(false),
414        }
415        Ok(true)
416    }
417
418    fn set_option_inner(&self, option: SetSocketOption) -> NetResult<bool> {
419        use SetSocketOption as O;
420
421        if let O::IpTos(tos) = option {
422            self.general.set_ip_tos(*tos);
423            self.sync_egress_ip_tos();
424            return Ok(true);
425        }
426
427        if self.general.set_option_inner(option)? {
428            return Ok(true);
429        }
430
431        match option {
432            O::NoDelay(no_delay) => {
433                self.with_smol_socket(|socket| {
434                    socket.set_nagle_enabled(!no_delay);
435                });
436            }
437            O::KeepAlive(keep_alive) => {
438                let interval =
439                    Duration::from_secs(self.keep_idle_secs.load(Ordering::Relaxed) as u64);
440                self.with_smol_socket(|socket| {
441                    socket.set_keep_alive(keep_alive.then_some(interval));
442                });
443            }
444            O::TcpKeepIdle(keep_idle) => {
445                if *keep_idle == 0 || *keep_idle > TCP_KEEPIDLE_MAX_SECS {
446                    return Err(NetError::InvalidInput);
447                }
448                self.keep_idle_secs.store(*keep_idle, Ordering::Relaxed);
449                let interval = Duration::from_secs(*keep_idle as u64);
450                self.with_smol_socket(|socket| {
451                    if socket.keep_alive().is_some() {
452                        socket.set_keep_alive(Some(interval));
453                    }
454                });
455            }
456            O::TcpKeepInterval(keep_interval) => {
457                if *keep_interval == 0 || *keep_interval > TCP_KEEPINTVL_MAX_SECS {
458                    return Err(NetError::InvalidInput);
459                }
460                self.keep_interval_secs
461                    .store(*keep_interval, Ordering::Relaxed);
462            }
463            O::TcpKeepCount(keep_count) => {
464                if *keep_count == 0 || *keep_count > TCP_KEEPCNT_MAX {
465                    return Err(NetError::InvalidInput);
466                }
467                self.keep_count.store(*keep_count, Ordering::Relaxed);
468            }
469            O::TcpUserTimeout(user_timeout) => {
470                self.user_timeout_millis
471                    .store(*user_timeout, Ordering::Relaxed);
472            }
473            O::TcpCongestionControl(congestion_control) => {
474                self.with_smol_socket(|socket| match congestion_control {
475                    TcpCongestionControl::None => {
476                        socket.set_congestion_control(smol::CongestionControl::None);
477                    }
478                });
479            }
480            _ => return Ok(false),
481        }
482        Ok(true)
483    }
484}
485impl SocketOps for TcpSocket {
486    fn bind(&self, local_addr: SocketAddrEx) -> NetResult {
487        let mut local_addr = local_addr.into_ip()?;
488        self.state
489            .lock(State::Idle)
490            .map_err(|_| NetError::InvalidInput)?
491            .transit(State::Idle, || {
492                // TODO: check addr is available
493                if local_addr.port() == 0 {
494                    local_addr.set_port(get_ephemeral_port()?);
495                }
496                if self.bound_endpoint.lock().port != 0 {
497                    return Err(NetError::InvalidInput);
498                }
499                let endpoint = IpListenEndpoint {
500                    addr: if local_addr.ip().is_unspecified() {
501                        None
502                    } else {
503                        Some(local_addr.ip().into())
504                    },
505                    port: local_addr.port(),
506                };
507                if !self.general.reuse_address()
508                    && !self.general.reuse_port()
509                    && !LISTEN_TABLE.can_listen(endpoint)
510                {
511                    return Err(NetError::AddrInUse);
512                }
513                let binding = get_control().local_binding_for(&endpoint)?;
514                self.register_bound_endpoint(endpoint)?;
515                *self.bound_endpoint.lock() = endpoint;
516                if binding.bound_if.is_some() {
517                    self.general.set_device_binding(binding);
518                }
519                debug!("TCP socket {}: binding to {}", self.handle, local_addr);
520                Ok(())
521            })
522    }
523
524    fn start_connect(&self, remote_addr: SocketAddrEx) -> NetResult<ConnectStatus> {
525        let remote_addr = remote_addr.into_ip()?;
526        self.begin_connect(remote_addr)?;
527        request_poll();
528        Ok(ConnectStatus::InProgress)
529    }
530
531    fn connect_status(&self) -> NetResult<ConnectStatus> {
532        match self.state.get() {
533            State::Connected => return Ok(ConnectStatus::Connected),
534            State::Connecting => {}
535            State::Closed => return Err(NetError::ConnectionRefused),
536            _ => return Err(NetError::InvalidInput),
537        }
538        request_poll();
539        let events = self.poll_connect();
540        if !events.contains(IoEvents::OUT) {
541            Ok(ConnectStatus::InProgress)
542        } else if self.state.get() == State::Connected {
543            Ok(ConnectStatus::Connected)
544        } else {
545            Err(NetError::ConnectionRefused)
546        }
547    }
548
549    fn listen(&self, backlog: usize) -> NetResult {
550        if let Ok(guard) = self.state.lock(State::Idle) {
551            guard.transit(State::Listening, || {
552                let mut bound_endpoint = *self.bound_endpoint.lock();
553                if bound_endpoint.port == 0 {
554                    bound_endpoint.port = get_ephemeral_port()?;
555                }
556                let binding = get_control().local_binding_for(&bound_endpoint)?;
557                self.with_bound_endpoint_registered(bound_endpoint, || {
558                    LISTEN_TABLE.listen(bound_endpoint, backlog, self.general.reuse_port())
559                })?;
560                *self.bound_endpoint.lock() = bound_endpoint;
561                self.sync_egress_ip_tos();
562                if binding.bound_if.is_some() {
563                    self.general.set_device_binding(binding);
564                }
565                debug!("listening on {}", bound_endpoint);
566                Ok(())
567            })?;
568        } else {
569            // ignore simultaneous `listen`s.
570        }
571        Ok(())
572    }
573
574    fn is_listening(&self) -> bool {
575        self.state.get() == State::Listening
576    }
577
578    fn try_accept(&self) -> NetResult<Socket> {
579        if self.state.get() != State::Listening {
580            return Err(NetError::InvalidInput);
581        }
582
583        let bound_endpoint = self.bound_endpoint()?;
584        request_poll();
585        let accepted = {
586            let mut sockets = SOCKET_SET.inner.lock();
587            let accepted = LISTEN_TABLE.accept(bound_endpoint, &mut sockets)?;
588            if matches!(LISTEN_TABLE.can_accept(bound_endpoint, &sockets), Ok(false)) {
589                // Preserve the empty interval for EPOLLET even when another
590                // connection arrives before the next poll. Holding SOCKET_SET
591                // prevents protocol progress between draining and publication.
592                // A concurrent unlisten must not turn an already accepted
593                // child's ownership into an error after it left the queue.
594                self.readiness_version.publish();
595            }
596            accepted
597        };
598        Ok({
599            let socket = TcpSocket::new_connected(
600                accepted.handle,
601                accepted.local_endpoint,
602                accepted.remote_endpoint,
603            );
604            socket.general.set_ip_tos(self.general.ip_tos());
605            socket.sync_egress_ip_tos();
606            debug!(
607                "accepted connection from {}, {}",
608                accepted.handle, accepted.remote_endpoint
609            );
610            socket.into()
611        })
612    }
613
614    fn try_send(&self, mut src: impl Read + IoBuf, _options: &mut SendOptions) -> NetResult<usize> {
615        if src.remaining() == 0 {
616            return Ok(0);
617        }
618        request_poll();
619        let result = self.with_smol_socket(|socket| {
620            if !socket.is_active() {
621                Err(NetError::NotConnected)
622            } else if !socket.can_send() {
623                Err(NetError::WouldBlock)
624            } else {
625                let len = socket
626                    .send(|buffer| {
627                        let result = src.read(buffer);
628                        let len = result.unwrap_or(0);
629                        (len, result)
630                    })
631                    .map_err(|_| NetError::NotConnected)??;
632                Ok(len)
633            }
634        });
635        if result.as_ref().is_ok_and(|sent| *sent > 0) {
636            request_poll();
637        }
638        result
639    }
640
641    fn try_recv(
642        &self,
643        mut dst: impl Write + IoBufMut,
644        options: &mut RecvOptions<'_>,
645    ) -> NetResult<usize> {
646        if self.rx_closed.load(Ordering::Acquire) {
647            return Err(NetError::NotConnected);
648        }
649        if self.state.get() == State::Closed {
650            return Err(NetError::NotConnected);
651        }
652        request_poll();
653        self.with_smol_socket(|socket| {
654            if socket.recv_queue() > 0 {
655                if options.flags.contains(RecvFlags::PEEK) {
656                    dst.write(
657                        socket
658                            .peek(dst.remaining_mut())
659                            .map_err(|_| NetError::NotConnected)?,
660                    )
661                    .map_err(NetError::from)
662                } else {
663                    // Drain currently available bytes from RX queue without waiting.
664                    // This loop copies across smoltcp's internal buffer segments to fill
665                    // the user buffer with as many bytes as are ready, but does not block
666                    // waiting for more data to arrive.
667                    let mut total = 0;
668                    while socket.recv_queue() > 0 && dst.remaining_mut() > 0 {
669                        let len = socket
670                            .recv(|buf| {
671                                let result = dst.write(buf).map_err(NetError::from);
672                                let len = result.unwrap_or(0);
673                                (len, result)
674                            })
675                            .map_err(|_| NetError::NotConnected)??;
676                        if len == 0 {
677                            break;
678                        }
679                        total += len;
680                    }
681                    if receive_starts_next_edge(total, socket.recv_queue()) {
682                        // Linux EPOLLET treats data arriving after the receive
683                        // queue was drained as a new edge even if epoll did not
684                        // sample the empty interval. Preserve that epoch for
685                        // the polling POSIX epoll adapter.
686                        self.readiness_version.publish();
687                    }
688                    Ok(total)
689                }
690            } else if !socket.may_recv() {
691                Ok(0)
692            } else {
693                Err(NetError::WouldBlock)
694            }
695        })
696    }
697
698    fn recv_available(&self) -> NetResult<usize> {
699        if self.state.get() == State::Listening {
700            return Err(NetError::InvalidInput);
701        }
702        let available = self.with_smol_socket(|socket| socket.recv_queue());
703        if available > 0 {
704            return Ok(available);
705        }
706        request_poll();
707        Ok(self.with_smol_socket(|socket| socket.recv_queue()))
708    }
709
710    fn local_addr(&self) -> NetResult<SocketAddrEx> {
711        let endpoint = self.with_smol_socket(|socket| {
712            socket
713                .local_endpoint()
714                .map(|endpoint| IpListenEndpoint {
715                    addr: Some(endpoint.addr),
716                    port: endpoint.port,
717                })
718                .unwrap_or_else(|| *self.bound_endpoint.lock())
719        });
720        Ok(SocketAddrEx::Ip(SocketAddr::new(
721            endpoint
722                .addr
723                .map_or_else(|| Ipv4Addr::UNSPECIFIED.into(), Into::into),
724            endpoint.port,
725        )))
726    }
727
728    fn peer_addr(&self) -> NetResult<SocketAddrEx> {
729        self.with_smol_socket(|socket| {
730            Ok(SocketAddrEx::Ip(
731                socket
732                    .remote_endpoint()
733                    .or_else(|| *self.peer_endpoint.lock())
734                    .ok_or(NetError::NotConnected)?
735                    .into(),
736            ))
737        })
738    }
739
740    fn shutdown(&self, how: Shutdown) -> NetResult {
741        // TODO(mivik): shutdown
742        if how.has_read() {
743            self.rx_closed.store(true, Ordering::Release);
744            // rx_closed is visible before waking RDHUP/EOF waiters.
745            self.readiness_version.publish();
746            unsafe { self.poll_rx_closed.wake(IoEvents::RDHUP | IoEvents::IN) };
747        }
748
749        // stream
750        if let Ok(guard) = self.state.lock(State::Connected) {
751            if how.has_read() && how.has_write() {
752                guard.transit(State::Closed, || {
753                    self.with_smol_socket(|socket| {
754                        debug!("TCP socket {}: shutting down", self.handle);
755                        socket.close();
756                    });
757                    self.clear_tracked_egress_ip_tos();
758                    self.unregister_bound_endpoint();
759                    *self.bound_endpoint.lock() = empty_endpoint();
760                    request_poll();
761                    Ok(())
762                })?;
763            } else if how.has_write() {
764                self.with_smol_socket(|socket| {
765                    debug!("TCP socket {}: shutting down write side", self.handle);
766                    socket.close();
767                });
768                request_poll();
769            }
770        }
771
772        // listener
773        if let Ok(guard) = self.state.lock(State::Listening) {
774            guard.transit(State::Closed, || {
775                LISTEN_TABLE.unlisten(self.bound_endpoint()?);
776                self.clear_tracked_egress_ip_tos();
777                self.unregister_bound_endpoint();
778                *self.bound_endpoint.lock() = empty_endpoint();
779                request_poll();
780                Ok(())
781            })?;
782        }
783
784        // ignore for other states
785        Ok(())
786    }
787}
788
789impl Pollable for TcpSocket {
790    fn poll(&self) -> IoEvents {
791        request_poll();
792        let mut events = match self.state.get() {
793            State::Connecting => self.poll_connect(),
794            State::Connected | State::Idle | State::Closed => self.poll_stream(),
795            State::Listening => self.poll_listener(),
796            State::Busy => IoEvents::empty(),
797        };
798        events.set(IoEvents::RDHUP, self.rx_closed.load(Ordering::Acquire));
799        events
800    }
801
802    unsafe fn register_shared(&self, sink: &mut dyn SharedRegistrationSink, events: IoEvents) {
803        self.register_poll_sources(events, |poll, interests| unsafe {
804            sink.register_shared(poll, interests)
805        });
806    }
807
808    unsafe fn register_exclusive(
809        &self,
810        sink: &mut dyn ExclusiveRegistrationSink,
811        events: IoEvents,
812    ) {
813        self.register_poll_sources(events, |poll, interests| unsafe {
814            sink.register_exclusive(poll, interests)
815        });
816    }
817}
818
819impl TcpSocket {
820    fn register_poll_sources(
821        &self,
822        events: IoEvents,
823        mut register: impl FnMut(&PollSet, IoEvents),
824    ) {
825        let mut accept_registration = None;
826        if self.state.get() == State::Listening && events.intersects(IoEvents::IN | IoEvents::RDHUP)
827        {
828            let port = self.bound_endpoint.lock().port;
829            if port != 0 {
830                let endpoint = *self.bound_endpoint.lock();
831                if let Some(accept_poll) = LISTEN_TABLE.accept_poll(endpoint) {
832                    // accept registration runs from task poll context after
833                    // releasing the listen-table lock.
834                    register(&accept_poll, IoEvents::IN);
835                    let accept_waker = LISTEN_TABLE.accept_waker(accept_poll.clone());
836                    accept_registration = Some((endpoint, accept_poll, accept_waker));
837                }
838            }
839        }
840        let recv_waker = if events.intersects(IoEvents::IN | IoEvents::RDHUP) {
841            // Socket registration runs from task poll context before taking the
842            // socket-set lock.
843            register(&self.poll_rx, IoEvents::IN | IoEvents::RDHUP);
844            Some(Waker::from(Arc::new(SocketDeferPollWake::new(
845                self.poll_rx.clone(),
846                IoEvents::IN | IoEvents::RDHUP,
847                self.readiness_version.clone(),
848            ))))
849        } else {
850            None
851        };
852        let send_waker = if events.contains(IoEvents::OUT) {
853            // Socket registration runs from task poll context before taking the
854            // socket-set lock.
855            register(&self.poll_tx, IoEvents::OUT);
856            Some(Waker::from(Arc::new(SocketDeferPollWake::new(
857                self.poll_tx.clone(),
858                IoEvents::OUT,
859                self.readiness_version.clone(),
860            ))))
861        } else {
862            None
863        };
864        if let Some((endpoint, accept_poll, accept_waker)) = accept_registration.as_ref() {
865            let mut sockets = SOCKET_SET.inner.lock();
866            LISTEN_TABLE.register_pending_accept_wakers(
867                *endpoint,
868                &mut sockets,
869                accept_poll,
870                accept_waker,
871            );
872        }
873        self.with_smol_socket(|socket| {
874            if let Some(waker) = recv_waker.as_ref() {
875                socket.register_recv_waker(waker);
876            }
877            if let Some(waker) = send_waker.as_ref() {
878                socket.register_send_waker(waker);
879            }
880        });
881        if events.intersects(IoEvents::IN | IoEvents::OUT | IoEvents::RDHUP) {
882            register(&self.poll_rx, events);
883            self.general
884                .register_waker(&Waker::from(Arc::new(SocketDeferPollWake::new(
885                    self.poll_rx.clone(),
886                    events,
887                    self.readiness_version.clone(),
888                ))));
889        }
890        if events.contains(IoEvents::RDHUP) {
891            // Registration happens from the OS-owned socket wait context.
892            register(&self.poll_rx_closed, IoEvents::RDHUP | IoEvents::IN);
893        }
894    }
895}
896
897impl Drop for TcpSocket {
898    fn drop(&mut self) {
899        let endpoint = *self.bound_endpoint.lock();
900        if self.state.get() == State::Listening && endpoint.port != 0 {
901            LISTEN_TABLE.unlisten(endpoint);
902        }
903
904        let should_orphan = self.with_smol_socket(|socket| {
905            let state = socket.state();
906            let should_orphan = matches!(
907                state,
908                smol::State::Established
909                    | smol::State::CloseWait
910                    | smol::State::FinWait1
911                    | smol::State::FinWait2
912                    | smol::State::Closing
913                    | smol::State::LastAck
914                    | smol::State::TimeWait
915            ) || socket.send_queue() > 0;
916            if matches!(
917                state,
918                smol::State::Established
919                    | smol::State::SynSent
920                    | smol::State::SynReceived
921                    | smol::State::CloseWait
922                    | smol::State::FinWait1
923                    | smol::State::FinWait2
924                    | smol::State::Closing
925                    | smol::State::LastAck
926            ) {
927                debug!("TCP socket {}: closing on drop", self.handle);
928                socket.close();
929            }
930            should_orphan
931        });
932
933        // Unbind from API layer (port registry, etc.)
934        self.clear_tracked_egress_ip_tos();
935        self.unregister_bound_endpoint();
936
937        if should_orphan {
938            // Keep the smoltcp socket alive after the user-facing handle is gone.
939            let timestamp = smoltcp::time::Instant::from_micros_const(
940                (ax_hal::time::monotonic_time_nanos() / 1_000) as i64,
941            );
942            crate::orphan::add_orphan(self.handle, timestamp);
943        } else {
944            SOCKET_SET.remove(self.handle);
945        }
946
947        // Ask the unique protocol executor to process teardown.
948        crate::request_poll();
949    }
950}
951
952fn duration_micros_u32(value: Duration) -> u32 {
953    value.total_micros().min(u32::MAX as u64) as u32
954}
955
956fn tcp_state_info(state: smol::State) -> TcpState {
957    match state {
958        smol::State::Closed => TcpState::Closed,
959        smol::State::Listen => TcpState::Listen,
960        smol::State::SynSent => TcpState::SynSent,
961        smol::State::SynReceived => TcpState::SynReceived,
962        smol::State::Established => TcpState::Established,
963        smol::State::FinWait1 => TcpState::FinWait1,
964        smol::State::FinWait2 => TcpState::FinWait2,
965        smol::State::CloseWait => TcpState::CloseWait,
966        smol::State::Closing => TcpState::Closing,
967        smol::State::LastAck => TcpState::LastAck,
968        smol::State::TimeWait => TcpState::TimeWait,
969    }
970}
971
972const fn empty_endpoint() -> IpListenEndpoint {
973    IpListenEndpoint {
974        addr: None,
975        port: 0,
976    }
977}
978
979impl TcpSocket {
980    /// Starts an active open and leaves completion to the protocol executor.
981    fn begin_connect(&self, remote_addr: SocketAddr) -> NetResult {
982        self.state
983            .lock(State::Idle)
984            .map_err(|state| {
985                if state == State::Connecting {
986                    NetError::InProgress
987                } else {
988                    // TODO(mivik): error code
989                    NetError::AlreadyConnected
990                }
991            })?
992            .transit(State::Connecting, || {
993                self.pending_error.store(0, Ordering::Release);
994                // TODO: check remote addr unreachable
995                // let (bound_endpoint, remote_endpoint) = self.get_endpoint_pair(remote_addr)?;
996                let remote_endpoint = IpEndpoint::from(remote_addr);
997                let mut bound_endpoint = *self.bound_endpoint.lock();
998
999                // Record original bind state before modifying
1000                let was_unbound_or_unspecified =
1001                    bound_endpoint.addr.is_none_or(|addr| addr.is_unspecified());
1002                let had_explicit_device_binding = self.general.device_binding().bound_if.is_some();
1003
1004                // Fill source address if unbound or bound to 0.0.0.0
1005                if bound_endpoint.addr.is_none_or(|addr| addr.is_unspecified()) {
1006                    bound_endpoint.addr = Some(
1007                        get_control()
1008                            .select_route_with_binding(
1009                                &remote_endpoint.addr,
1010                                self.general.device_binding(),
1011                            )?
1012                            .source,
1013                    );
1014                }
1015                if bound_endpoint.port == 0 {
1016                    bound_endpoint.port = get_ephemeral_port()?;
1017                }
1018                info!(
1019                    "TCP connection from {} to {}",
1020                    bound_endpoint, remote_endpoint
1021                );
1022                self.with_bound_endpoint_registered(bound_endpoint, || {
1023                    let mut service = get_service();
1024                    let context = service.iface.context();
1025                    self.with_smol_socket(|socket| {
1026                        socket
1027                            .connect(context, remote_endpoint, bound_endpoint)
1028                            .map_err(|e| match e {
1029                                smol::ConnectError::InvalidState => NetError::AlreadyConnected,
1030                                smol::ConnectError::Unaddressable => NetError::ConnectionRefused,
1031                            })?;
1032                        Ok::<(), NetError>(())
1033                    })
1034                })?;
1035                *self.bound_endpoint.lock() = bound_endpoint;
1036
1037                // Only set device binding if was originally unbound or bound to 0.0.0.0
1038                // Binding to a specific IP should lock the interface
1039                if !had_explicit_device_binding && was_unbound_or_unspecified {
1040                    self.general
1041                        .set_device_binding(get_control().local_binding_for(&bound_endpoint)?);
1042                }
1043                // else: bound to specific IP, keep existing interface binding
1044                self.sync_egress_ip_tos();
1045
1046                Ok(())
1047            })
1048    }
1049
1050    /// Registers the public TCP bind side table if not already registered.
1051    fn register_bound_endpoint(&self, endpoint: IpListenEndpoint) -> NetResult {
1052        if !self.bound_registered.load(Ordering::Acquire) {
1053            register_tcp_bound(endpoint, self.general.reuse_port())?;
1054            self.bound_registered.store(true, Ordering::Release);
1055        }
1056        Ok(())
1057    }
1058
1059    fn with_bound_endpoint_registered<R>(
1060        &self,
1061        endpoint: IpListenEndpoint,
1062        f: impl FnOnce() -> NetResult<R>,
1063    ) -> NetResult<R> {
1064        let register_bound = !self.bound_registered.load(Ordering::Acquire);
1065        if register_bound {
1066            register_tcp_bound(endpoint, self.general.reuse_port())?;
1067        }
1068        match f() {
1069            Ok(value) => {
1070                if register_bound {
1071                    self.bound_registered.store(true, Ordering::Release);
1072                }
1073                Ok(value)
1074            }
1075            Err(err) => {
1076                if register_bound {
1077                    unregister_tcp_bound(endpoint);
1078                }
1079                Err(err)
1080            }
1081        }
1082    }
1083
1084    /// Removes the public TCP bind side-table entry, if present.
1085    fn unregister_bound_endpoint(&self) {
1086        if self.bound_registered.swap(false, Ordering::AcqRel) {
1087            unregister_tcp_bound(*self.bound_endpoint.lock());
1088        }
1089    }
1090}
1091
1092/// One TCP bind ownership record. Several records may share a port only when
1093/// every binder requested SO_REUSEPORT on the identical local address, mirroring
1094/// Linux's reuseport group semantics.
1095struct TcpBoundEntry {
1096    addr: Option<smoltcp::wire::IpAddress>,
1097    reuse_port: bool,
1098}
1099
1100static TCP_BOUND_PORTS: LazyLock<SpinLock<HashMap<u16, Vec<TcpBoundEntry>>>> =
1101    LazyLock::new(|| SpinLock::new(HashMap::new()));
1102
1103/// Registers TCP bind ownership with wildcard/specific address conflicts.
1104///
1105/// A binder joins an existing reuseport group only when it and every colliding
1106/// owner requested SO_REUSEPORT on the exact same local address; any other
1107/// address overlap on the port is rejected with `EADDRINUSE`.
1108fn register_tcp_bound(endpoint: IpListenEndpoint, reuse_port: bool) -> NetResult {
1109    if endpoint.port == 0 {
1110        return Ok(());
1111    }
1112
1113    let mut bound_ports = TCP_BOUND_PORTS.lock();
1114    let entries = bound_ports.entry(endpoint.port).or_default();
1115    for entry in entries.iter() {
1116        if listen_addrs_conflict(entry.addr, endpoint.addr)
1117            && !(reuse_port && entry.reuse_port && entry.addr == endpoint.addr)
1118        {
1119            return Err(NetError::AddrInUse);
1120        }
1121    }
1122    entries.push(TcpBoundEntry {
1123        addr: endpoint.addr,
1124        reuse_port,
1125    });
1126    Ok(())
1127}
1128
1129/// Removes one TCP bind registration.
1130fn unregister_tcp_bound(endpoint: IpListenEndpoint) {
1131    if endpoint.port == 0 {
1132        return;
1133    }
1134    let mut bound_ports = TCP_BOUND_PORTS.lock();
1135    let Some(entries) = bound_ports.get_mut(&endpoint.port) else {
1136        return;
1137    };
1138    if let Some(index) = entries.iter().position(|entry| entry.addr == endpoint.addr) {
1139        entries.swap_remove(index);
1140    }
1141    if entries.is_empty() {
1142        bound_ports.remove(&endpoint.port);
1143    }
1144}
1145
1146/// Returns whether a port is safe for ephemeral TCP allocation.
1147fn tcp_port_available(port: u16) -> bool {
1148    // Ephemeral ports are selected conservatively: avoid any port that has a
1149    // listener or bound socket on any local address.
1150    LISTEN_TABLE.can_listen(IpListenEndpoint { addr: None, port })
1151        && !TCP_BOUND_PORTS.lock().contains_key(&port)
1152}
1153
1154fn get_ephemeral_port() -> NetResult<u16> {
1155    allocate_ephemeral_port(tcp_port_available)
1156}