1use 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
77pub struct TcpSocket {
79 state: StateLock,
81 handle: SocketHandle,
83 bound_endpoint: SpinLock<IpListenEndpoint>,
85 peer_endpoint: SpinLock<Option<IpEndpoint>>,
87 tos_key: SpinLock<Option<EgressIpTosKey>>,
89 bound_registered: AtomicBool,
91
92 general: GeneralOptions,
94 pending_error: AtomicI32,
96 keep_idle_secs: AtomicU32,
98 keep_interval_secs: AtomicU32,
100 keep_count: AtomicU32,
102 user_timeout_millis: AtomicU32,
104 rx_closed: AtomicBool,
106 poll_rx: Arc<PollSet>,
108 poll_tx: Arc<PollSet>,
110 poll_rx_closed: PollSet,
112 readiness_version: ReadinessVersion,
114}
115
116unsafe impl Sync for TcpSocket {}
117
118impl TcpSocket {
119 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), 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 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 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), 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 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
208impl 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 }
309 smol::State::Established => {
310 self.pending_error.store(0, Ordering::Release);
311 self.state.set(State::Connected); *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); 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 **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 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 }
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 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 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 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 if how.has_read() {
743 self.rx_closed.store(true, Ordering::Release);
744 self.readiness_version.publish();
746 unsafe { self.poll_rx_closed.wake(IoEvents::RDHUP | IoEvents::IN) };
747 }
748
749 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 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 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 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 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 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 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 self.clear_tracked_egress_ip_tos();
935 self.unregister_bound_endpoint();
936
937 if should_orphan {
938 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 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 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 NetError::AlreadyConnected
990 }
991 })?
992 .transit(State::Connecting, || {
993 self.pending_error.store(0, Ordering::Release);
994 let remote_endpoint = IpEndpoint::from(remote_addr);
997 let mut bound_endpoint = *self.bound_endpoint.lock();
998
999 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 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 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 self.sync_egress_ip_tos();
1045
1046 Ok(())
1047 })
1048 }
1049
1050 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 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
1092struct 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
1103fn 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
1129fn 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
1146fn tcp_port_available(port: u16) -> bool {
1148 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}