1use crate::runtime::with_ambient_tokio_runtime;
2use crate::sockets::{
3 ErrorCode, MaybeReady, SocketAddrCheck, SocketAddrUse, SocketAddressFamily, WasiSocketsCtx,
4 get_receive_buffer_size, get_send_buffer_size, get_unicast_hop_limit, is_valid_address_family,
5 is_valid_remote_address, is_valid_unicast_address, set_receive_buffer_size,
6 set_send_buffer_size, set_unicast_hop_limit, unspecified_addr,
7};
8use rustix::fd::AsFd;
9use rustix::io::Errno;
10use rustix::net::sockopt;
11use std::fmt::Debug;
12use std::future::poll_fn;
13use std::mem;
14use std::net::SocketAddr;
15use std::sync::Arc;
16use std::task::{Poll, ready};
17use std::time::Duration;
18
19const DEFAULT_BACKLOG: u32 = 128;
21
22const NANOS_PER_SEC: u64 = 1_000_000_000;
23
24enum TcpState {
29 Default(tokio::net::TcpSocket),
36
37 Listening(Arc<tokio::net::TcpListener>),
41
42 Connecting(MaybeReady<Result<tokio::net::TcpStream, ErrorCode>>),
49
50 Connected {
57 stream: Arc<tokio::net::TcpStream>,
58 peer: Option<SocketAddr>,
61 receive_taken: bool,
62 send_taken: bool,
63 },
64
65 Closed(ErrorCode),
67}
68impl TcpState {
69 fn connected(stream: tokio::net::TcpStream, peer: Option<SocketAddr>) -> Self {
70 TcpState::Connected {
71 stream: Arc::new(stream),
72 peer,
73 receive_taken: false,
74 send_taken: false,
75 }
76 }
77 fn take(&mut self) -> Self {
78 mem::replace(self, TcpState::Closed(ErrorCode::Other))
79 }
80}
81impl Debug for TcpState {
82 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
83 match self {
84 Self::Default(_) => f.debug_tuple("Default").finish(),
85 Self::Listening { .. } => f.debug_tuple("Listening").finish(),
86 Self::Connecting(..) => f.debug_tuple("Connecting").finish(),
87 Self::Connected { .. } => f.debug_tuple("Connected").finish(),
88 Self::Closed(..) => write!(f, "Closed"),
89 }
90 }
91}
92
93pub struct TcpSocket {
95 tcp_state: TcpState,
97
98 listen_backlog_size: u32,
100
101 family: SocketAddressFamily,
102
103 permissions: SocketAddrCheck,
105
106 listener_options: NonInheritedOptions,
109
110 is_bound: bool,
113}
114
115impl TcpSocket {
116 pub(crate) fn new(
118 ctx: &WasiSocketsCtx,
119 family: SocketAddressFamily,
120 ) -> Result<Self, ErrorCode> {
121 ctx.allowed_network_uses.check_allowed_tcp()?;
122
123 let socket = with_ambient_tokio_runtime(|| socket(family))?;
124
125 Ok(Self {
126 tcp_state: TcpState::Default(socket),
127 listen_backlog_size: DEFAULT_BACKLOG,
128 family,
129 is_bound: false,
130 listener_options: Default::default(),
131 permissions: ctx.socket_addr_check.clone(),
132 })
133 }
134
135 fn as_fd(&self) -> Result<rustix::fd::BorrowedFd<'_>, ErrorCode> {
136 match &self.tcp_state {
137 TcpState::Default(socket) => Ok(socket.as_fd()),
138 TcpState::Connected { stream, .. } => Ok(stream.as_fd()),
139 TcpState::Listening(listener) => Ok(listener.as_fd()),
140 TcpState::Connecting(..) => Err(ErrorCode::InvalidState),
141 TcpState::Closed(err) => Err(*err),
142 }
143 }
144
145 pub(crate) fn is_bound(&mut self) -> bool {
146 if !self.is_bound {
150 self.is_bound = match &self.tcp_state {
151 TcpState::Default(socket) => socket
152 .local_addr()
153 .is_ok_and(|addr| addr != unspecified_addr(self.family)),
154 _ => true,
155 };
156 }
157 self.is_bound
158 }
159
160 pub(crate) async fn bind(&mut self, addr: SocketAddr) -> Result<(), ErrorCode> {
161 if self.is_bound() {
162 return Err(ErrorCode::InvalidState);
163 }
164 let TcpState::Default(sock) = &self.tcp_state else {
165 return Err(ErrorCode::InvalidState);
166 };
167
168 if !is_valid_unicast_address(addr.ip()) || !is_valid_address_family(addr.ip(), self.family)
169 {
170 return Err(ErrorCode::InvalidArgument);
171 }
172
173 self.permissions.check(addr, SocketAddrUse::TcpBind).await?;
174 bind(sock, addr)?;
175 Ok(())
176 }
177
178 pub(crate) fn start_connect(&mut self, addr: SocketAddr) -> Result<(), ErrorCode> {
179 let TcpState::Default(_) = &self.tcp_state else {
180 return Err(ErrorCode::InvalidState);
181 };
182
183 let permissions = self.permissions.clone();
184 let family = self.family;
185 let already_bound = self.is_bound();
186
187 if !is_valid_unicast_address(addr.ip())
188 || !is_valid_remote_address(addr)
189 || !is_valid_address_family(addr.ip(), family)
190 {
191 return Err(ErrorCode::InvalidArgument);
192 };
193
194 let TcpState::Default(sock) = self.tcp_state.take() else {
195 unreachable!();
196 };
197
198 self.tcp_state = TcpState::Connecting(MaybeReady::new(Box::pin(async move {
199 {
201 if !already_bound {
202 let implicit = unspecified_addr(family);
212 permissions.check(implicit, SocketAddrUse::TcpBind).await?;
213 }
214
215 permissions.check(addr, SocketAddrUse::TcpConnect).await?;
216 }
217
218 let stream = sock.connect(addr).await?;
219 Ok(stream)
220 })));
221
222 Ok(())
223 }
224
225 pub(crate) fn poll_finish_connect(
226 &mut self,
227 cx: &mut std::task::Context<'_>,
228 ) -> Poll<Result<(), ErrorCode>> {
229 match &mut self.tcp_state {
230 TcpState::Connecting(connect) => {
231 ready!(with_ambient_tokio_runtime(|| connect.poll_ready(cx)));
232 }
233 TcpState::Connected { .. } => return Poll::Ready(Ok(())),
234 TcpState::Closed(e) => return Poll::Ready(Err(*e)),
235 _ => return Poll::Ready(Err(ErrorCode::InvalidState)),
236 }
237 let TcpState::Connecting(connect) = self.tcp_state.take() else {
238 unreachable!();
239 };
240
241 match connect.unwrap_ready() {
242 Ok(stream) => {
243 self.tcp_state = TcpState::connected(stream, None);
244 Poll::Ready(Ok(()))
245 }
246 Err(err) => {
247 self.tcp_state = TcpState::Closed(err);
248 Poll::Ready(Err(err))
249 }
250 }
251 }
252
253 pub(crate) async fn listen(&mut self) -> Result<TcpListenStream, ErrorCode> {
254 let already_bound = self.is_bound();
255 let sock = match self.tcp_state.take() {
256 TcpState::Default(sock) => sock,
257 tcp_state => {
258 self.tcp_state = tcp_state;
259 return Err(ErrorCode::InvalidState);
260 }
261 };
262
263 {
265 if already_bound {
266 self.permissions
267 .check(sock.local_addr()?, SocketAddrUse::TcpListen)
268 .await?;
269 } else {
270 let implicit = unspecified_addr(self.family);
271 self.permissions
272 .check(implicit, SocketAddrUse::TcpBind)
273 .await?;
274 self.permissions
275 .check(implicit, SocketAddrUse::TcpListen)
276 .await?;
277 }
278 }
279
280 if !already_bound {
307 let implicit = unspecified_addr(self.family);
308 bind(&sock, implicit)?;
309 }
310
311 let listener = sock.listen(self.listen_backlog_size).map_err(|err| {
312 match Errno::from_io_error(&err) {
313 #[cfg(windows)]
323 Some(Errno::MFILE) => Errno::NOBUFS.into(),
324
325 _ => err,
326 }
327 })?;
328 let listener = Arc::new(listener);
329 self.tcp_state = TcpState::Listening(listener.clone());
330
331 Ok(TcpListenStream {
332 inner: listener,
333 listener_options: self.listener_options.clone(),
334 family: self.family,
335 permissions: self.permissions.clone(),
336 pending_accept: None,
337 })
338 }
339
340 pub(crate) fn take_send_stream(&mut self) -> Result<TcpSendStream, ErrorCode> {
341 match &mut self.tcp_state {
342 TcpState::Connected {
343 stream, send_taken, ..
344 } if !*send_taken => {
345 *send_taken = true;
346 Ok(TcpSendStream {
347 inner: stream.clone(),
348 })
349 }
350 TcpState::Closed(err) => Err(*err),
351 _ => Err(ErrorCode::InvalidState),
352 }
353 }
354
355 pub(crate) fn take_receive_stream(&mut self) -> Result<TcpReceiveStream, ErrorCode> {
356 match &mut self.tcp_state {
357 TcpState::Connected {
358 stream,
359 receive_taken,
360 ..
361 } if !*receive_taken => {
362 *receive_taken = true;
363 Ok(TcpReceiveStream {
364 inner: stream.clone(),
365 })
366 }
367 TcpState::Closed(err) => Err(*err),
368 _ => Err(ErrorCode::InvalidState),
369 }
370 }
371
372 pub(crate) fn local_address(&mut self) -> Result<SocketAddr, ErrorCode> {
373 if !self.is_bound() {
374 return Err(ErrorCode::InvalidState);
375 }
376
377 match &self.tcp_state {
378 TcpState::Default(socket) => Ok(socket.local_addr()?),
379 TcpState::Connecting(_) => Err(ErrorCode::InvalidState),
380 TcpState::Connected { stream, .. } => Ok(stream.local_addr()?),
381 TcpState::Listening(listener) => Ok(listener.local_addr()?),
382 TcpState::Closed(err) => Err(*err),
383 }
384 }
385
386 pub(crate) fn remote_address(&mut self) -> Result<SocketAddr, ErrorCode> {
387 match &mut self.tcp_state {
388 TcpState::Connected {
389 peer: Some(peer), ..
390 } => Ok(*peer),
391 TcpState::Connected { stream, peer, .. } => {
392 let addr = stream.peer_addr()?;
393 *peer = Some(addr);
394 Ok(addr)
395 }
396 TcpState::Closed(err) => Err(*err),
397 _ => Err(ErrorCode::InvalidState),
398 }
399 }
400
401 pub(crate) fn is_listening(&self) -> bool {
402 matches!(self.tcp_state, TcpState::Listening(_))
403 }
404
405 pub(crate) fn address_family(&self) -> SocketAddressFamily {
406 self.family
407 }
408
409 pub(crate) fn set_listen_backlog_size(&mut self, value: u64) -> Result<(), ErrorCode> {
410 const MIN_BACKLOG: u32 = 1;
411 const MAX_BACKLOG: u32 = i32::MAX as u32; if value == 0 {
414 return Err(ErrorCode::InvalidArgument);
415 }
416 let value = value
418 .try_into()
419 .unwrap_or(MAX_BACKLOG)
420 .clamp(MIN_BACKLOG, MAX_BACKLOG);
421 match &self.tcp_state {
422 TcpState::Default(..) => {
423 self.listen_backlog_size = value;
425 Ok(())
426 }
427 TcpState::Listening(listener) => {
428 if rustix::net::listen(&listener, value.try_into().unwrap_or(i32::MAX)).is_err() {
431 return Err(ErrorCode::NotSupported);
432 }
433 self.listen_backlog_size = value;
434 Ok(())
435 }
436 TcpState::Closed(err) => Err(*err),
437 _ => Err(ErrorCode::InvalidState),
438 }
439 }
440
441 pub(crate) fn keep_alive_enabled(&self) -> Result<bool, ErrorCode> {
442 let fd = self.as_fd()?;
443 let v = sockopt::socket_keepalive(fd)?;
444 Ok(v)
445 }
446
447 pub(crate) fn set_keep_alive_enabled(&self, value: bool) -> Result<(), ErrorCode> {
448 let fd = self.as_fd()?;
449 sockopt::set_socket_keepalive(fd, value)?;
450 Ok(())
451 }
452
453 pub(crate) fn keep_alive_idle_time(&self) -> Result<u64, ErrorCode> {
454 let fd = self.as_fd()?;
455 let v = sockopt::tcp_keepidle(fd)?;
456 Ok(v.as_nanos().try_into().unwrap_or(u64::MAX))
457 }
458
459 pub(crate) fn set_keep_alive_idle_time(&mut self, value: u64) -> Result<(), ErrorCode> {
460 if value == 0 {
461 return Err(ErrorCode::InvalidArgument);
463 }
464 let fd = self.as_fd()?;
465 let value = clamp_keep_alive_time(value);
466 sockopt::set_tcp_keepidle(fd, Duration::from_nanos(value))?;
467 self.listener_options.set_keep_alive_idle_time(value);
468 Ok(())
469 }
470
471 pub(crate) fn keep_alive_interval(&self) -> Result<u64, ErrorCode> {
472 let fd = self.as_fd()?;
473 let v = sockopt::tcp_keepintvl(fd)?;
474 Ok(v.as_nanos().try_into().unwrap_or(u64::MAX))
475 }
476
477 pub(crate) fn set_keep_alive_interval(&self, value: u64) -> Result<(), ErrorCode> {
478 if value == 0 {
479 return Err(ErrorCode::InvalidArgument);
481 }
482 let fd = self.as_fd()?;
483 let value = clamp_keep_alive_time(value);
484 sockopt::set_tcp_keepintvl(fd, Duration::from_nanos(value))?;
485 Ok(())
486 }
487
488 pub(crate) fn keep_alive_count(&self) -> Result<u32, ErrorCode> {
489 let fd = self.as_fd()?;
490 let v = sockopt::tcp_keepcnt(fd)?;
491 Ok(v)
492 }
493
494 pub(crate) fn set_keep_alive_count(&self, value: u32) -> Result<(), ErrorCode> {
495 if value == 0 {
496 return Err(ErrorCode::InvalidArgument);
498 }
499 let value = clamp_keep_alive_count(value);
500 let fd = self.as_fd()?;
501 sockopt::set_tcp_keepcnt(fd, value)?;
502 Ok(())
503 }
504
505 pub(crate) fn hop_limit(&self) -> Result<u8, ErrorCode> {
506 let fd = self.as_fd()?;
507 let n = get_unicast_hop_limit(fd, self.family)?;
508 Ok(n)
509 }
510
511 pub(crate) fn set_hop_limit(&mut self, value: u8) -> Result<(), ErrorCode> {
512 {
513 let fd = self.as_fd()?;
514 set_unicast_hop_limit(fd, self.family, value)?;
515 }
516 self.listener_options.set_hop_limit(value);
517 Ok(())
518 }
519
520 pub(crate) fn receive_buffer_size(&self) -> Result<u64, ErrorCode> {
521 let fd = self.as_fd()?;
522 let n = get_receive_buffer_size(fd)?;
523 Ok(n)
524 }
525
526 pub(crate) fn set_receive_buffer_size(&mut self, value: u64) -> Result<(), ErrorCode> {
527 let res = {
528 let fd = self.as_fd()?;
529 set_receive_buffer_size(fd, value)?
530 };
531 self.listener_options.set_receive_buffer_size(res);
532 Ok(())
533 }
534
535 pub(crate) fn send_buffer_size(&self) -> Result<u64, ErrorCode> {
536 let fd = self.as_fd()?;
537 let n = get_send_buffer_size(fd)?;
538 Ok(n)
539 }
540
541 pub(crate) fn set_send_buffer_size(&mut self, value: u64) -> Result<(), ErrorCode> {
542 let res = {
543 let fd = self.as_fd()?;
544 set_send_buffer_size(fd, value)?
545 };
546 self.listener_options.set_send_buffer_size(res);
547 Ok(())
548 }
549}
550
551pub(crate) struct TcpListenStream {
552 inner: Arc<tokio::net::TcpListener>,
553 family: SocketAddressFamily,
554 listener_options: NonInheritedOptions,
555 permissions: SocketAddrCheck,
556 pending_accept: Option<MaybeReady<Result<(tokio::net::TcpStream, SocketAddr), ErrorCode>>>,
557}
558impl TcpListenStream {
559 pub(crate) fn poll_accept(&mut self, cx: &mut std::task::Context<'_>) -> Poll<TcpSocket> {
560 ready!(self.poll_ready(cx));
561 let result = self.pending_accept.take().unwrap().unwrap_ready();
562 Poll::Ready(TcpSocket {
563 tcp_state: match result {
564 Ok((client, peer)) => {
565 self.listener_options.apply(self.family, &client);
566 TcpState::connected(client, Some(peer))
567 }
568 Err(err) => TcpState::Closed(err),
569 },
570 listen_backlog_size: DEFAULT_BACKLOG,
571 family: self.family,
572 is_bound: true,
573 listener_options: Default::default(),
574 permissions: self.permissions.clone(),
575 })
576 }
577
578 pub(crate) fn poll_ready(&mut self, cx: &mut std::task::Context<'_>) -> Poll<()> {
579 if self.pending_accept.is_none() {
580 let listener = self.inner.clone();
581 let permissions = self.permissions.clone();
582
583 self.pending_accept = Some(MaybeReady::new(Box::pin(async move {
584 loop {
585 match accept(&listener).await {
586 Ok((client, addr)) => {
587 if permissions
588 .check(addr, SocketAddrUse::TcpAccept)
589 .await
590 .is_ok()
591 {
592 return Ok((client, addr));
593 } else {
594 reset(client);
595 continue;
596 }
597 }
598 Err(err) => {
599 return Err(err.into());
600 }
601 }
602 }
603 })));
604 }
605
606 with_ambient_tokio_runtime(|| {
607 self.pending_accept
608 .as_mut()
609 .unwrap()
610 .poll_ready(cx)
611 .map(|_| ())
612 })
613 }
614}
615
616pub(crate) struct TcpSendStream {
617 inner: Arc<tokio::net::TcpStream>,
618}
619impl TcpSendStream {
620 pub(crate) fn poll_ready(&mut self, cx: &mut std::task::Context<'_>) -> Poll<()> {
621 self.inner.poll_write_ready(cx).map(|_| ())
622 }
623
624 pub(crate) fn poll_write(
625 &mut self,
626 cx: &mut std::task::Context<'_>,
627 buf: &[u8],
628 ) -> Poll<Result<usize, ErrorCode>> {
629 loop {
630 return match self.inner.try_write(buf) {
631 Ok(n) => Poll::Ready(Ok(n)),
632 Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
633 match self.inner.poll_write_ready(cx) {
634 Poll::Ready(Ok(())) => continue,
635 Poll::Ready(Err(e)) => Poll::Ready(Err(e.into())),
636 Poll::Pending => Poll::Pending,
637 }
638 }
639 Err(e) => Poll::Ready(Err(match Errno::from_io_error(&e) {
640 #[cfg(windows)]
641 Some(Errno::SHUTDOWN) | Some(Errno::CONNABORTED) => ErrorCode::ConnectionBroken,
642 #[cfg(not(windows))]
643 Some(Errno::PIPE) => ErrorCode::ConnectionBroken,
644
645 _ => e.into(),
646 })),
647 };
648 }
649 }
650 pub(crate) async fn write(&mut self, buf: &[u8]) -> Result<usize, ErrorCode> {
651 poll_fn(|cx| self.poll_write(cx, buf)).await
652 }
653}
654impl Drop for TcpSendStream {
655 fn drop(&mut self) {
656 _ = rustix::net::shutdown(&self.inner, rustix::net::Shutdown::Write);
657 }
658}
659
660pub(crate) struct TcpReceiveStream {
661 inner: Arc<tokio::net::TcpStream>,
662}
663impl TcpReceiveStream {
664 pub(crate) fn poll_ready(&mut self, cx: &mut std::task::Context<'_>) -> Poll<()> {
665 self.inner.poll_read_ready(cx).map(|_| ())
666 }
667
668 pub(crate) fn poll_read(
669 &mut self,
670 cx: &mut std::task::Context<'_>,
671 buf: &mut [u8],
672 ) -> Poll<Result<usize, ErrorCode>> {
673 if buf.is_empty() {
674 return Poll::Ready(Ok(0));
675 }
676 loop {
677 return match self.inner.try_read(buf) {
678 Ok(0) => Poll::Ready(Ok(0)),
679 Ok(n) => Poll::Ready(Ok(n)),
680 Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
681 match self.inner.poll_read_ready(cx) {
682 Poll::Ready(Ok(())) => continue,
683 Poll::Ready(Err(e)) => Poll::Ready(Err(e.into())),
684 Poll::Pending => Poll::Pending,
685 }
686 }
687 Err(e) => Poll::Ready(Err(e.into())),
688 };
689 }
690 }
691}
692impl Drop for TcpReceiveStream {
693 fn drop(&mut self) {
694 _ = rustix::net::shutdown(&self.inner, rustix::net::Shutdown::Read);
695 }
696}
697
698#[cfg(not(target_os = "macos"))]
699pub use inherits_option::*;
700#[cfg(not(target_os = "macos"))]
701mod inherits_option {
702 use crate::sockets::SocketAddressFamily;
703 use tokio::net::TcpStream;
704
705 #[derive(Default, Clone)]
706 pub struct NonInheritedOptions;
707
708 impl NonInheritedOptions {
709 pub fn set_keep_alive_idle_time(&mut self, _value: u64) {}
710
711 pub fn set_hop_limit(&mut self, _value: u8) {}
712
713 pub fn set_receive_buffer_size(&mut self, _value: usize) {}
714
715 pub fn set_send_buffer_size(&mut self, _value: usize) {}
716
717 pub(crate) fn apply(&self, _family: SocketAddressFamily, _stream: &TcpStream) {}
718 }
719}
720
721#[cfg(target_os = "macos")]
722pub use does_not_inherit_options::*;
723#[cfg(target_os = "macos")]
724mod does_not_inherit_options {
725 use crate::sockets::SocketAddressFamily;
726 use rustix::net::sockopt;
727 use std::sync::Arc;
728 use std::sync::atomic::{AtomicU8, AtomicU64, AtomicUsize, Ordering::Relaxed};
729 use std::time::Duration;
730 use tokio::net::TcpStream;
731
732 #[derive(Default, Clone)]
736 pub struct NonInheritedOptions(Arc<Inner>);
737
738 #[derive(Default)]
739 struct Inner {
740 receive_buffer_size: AtomicUsize,
741 send_buffer_size: AtomicUsize,
742 hop_limit: AtomicU8,
743 keep_alive_idle_time: AtomicU64, }
745
746 impl NonInheritedOptions {
747 pub fn set_keep_alive_idle_time(&mut self, value: u64) {
748 self.0.keep_alive_idle_time.store(value, Relaxed);
749 }
750
751 pub fn set_hop_limit(&mut self, value: u8) {
752 self.0.hop_limit.store(value, Relaxed);
753 }
754
755 pub fn set_receive_buffer_size(&mut self, value: usize) {
756 self.0.receive_buffer_size.store(value, Relaxed);
757 }
758
759 pub fn set_send_buffer_size(&mut self, value: usize) {
760 self.0.send_buffer_size.store(value, Relaxed);
761 }
762
763 pub(crate) fn apply(&self, family: SocketAddressFamily, stream: &TcpStream) {
764 let receive_buffer_size = self.0.receive_buffer_size.load(Relaxed);
769 if receive_buffer_size > 0 {
770 _ = sockopt::set_socket_recv_buffer_size(&stream, receive_buffer_size);
772 }
773
774 let send_buffer_size = self.0.send_buffer_size.load(Relaxed);
775 if send_buffer_size > 0 {
776 _ = sockopt::set_socket_send_buffer_size(&stream, send_buffer_size);
778 }
779
780 if family == SocketAddressFamily::Ipv6 {
782 let hop_limit = self.0.hop_limit.load(Relaxed);
783 if hop_limit > 0 {
784 _ = sockopt::set_ipv6_unicast_hops(&stream, Some(hop_limit));
786 }
787 }
788
789 let keep_alive_idle_time = self.0.keep_alive_idle_time.load(Relaxed);
790 if keep_alive_idle_time > 0 {
791 _ = sockopt::set_tcp_keepidle(&stream, Duration::from_nanos(keep_alive_idle_time));
793 }
794 }
795 }
796}
797
798fn socket(family: SocketAddressFamily) -> std::io::Result<tokio::net::TcpSocket> {
799 match family {
800 SocketAddressFamily::Ipv4 => tokio::net::TcpSocket::new_v4(),
801 SocketAddressFamily::Ipv6 => {
802 let socket = tokio::net::TcpSocket::new_v6()?;
803
804 sockopt::set_ipv6_v6only(&socket, true)?;
808 Ok(socket)
809 }
810 }
811}
812
813fn bind(socket: &tokio::net::TcpSocket, local_address: SocketAddr) -> Result<(), ErrorCode> {
814 #[cfg(not(windows))]
821 {
822 _ = sockopt::set_socket_reuseaddr(&socket, true);
823 }
824
825 socket
827 .bind(local_address)
828 .map_err(|err| match Errno::from_io_error(&err) {
829 Some(Errno::AFNOSUPPORT) => ErrorCode::InvalidArgument,
836 #[cfg(windows)]
839 Some(Errno::NOBUFS) => ErrorCode::AddressInUse,
840 _ => err.into(),
841 })
842}
843
844async fn accept(
845 listener: &tokio::net::TcpListener,
846) -> std::io::Result<(tokio::net::TcpStream, SocketAddr)> {
847 listener
848 .accept()
849 .await
850 .map_err(|err| match Errno::from_io_error(&err) {
851 #[cfg(windows)]
859 Some(Errno::INPROGRESS) => Errno::INTR.into(),
860
861 #[cfg(target_os = "linux")]
868 Some(
869 Errno::CONNRESET
870 | Errno::NETRESET
871 | Errno::HOSTUNREACH
872 | Errno::HOSTDOWN
873 | Errno::NETDOWN
874 | Errno::NETUNREACH
875 | Errno::PROTO
876 | Errno::NOPROTOOPT
877 | Errno::NONET
878 | Errno::OPNOTSUPP,
879 ) => Errno::CONNABORTED.into(),
880
881 _ => err,
882 })
883}
884
885fn reset(socket: tokio::net::TcpStream) {
886 _ = socket.set_zero_linger();
887 drop(socket);
888}
889
890fn clamp_keep_alive_time(value: u64) -> u64 {
891 const MIN: u64 = 1 * NANOS_PER_SEC;
893
894 const MAX: u64 = (i16::MAX as u64) * NANOS_PER_SEC;
896
897 value.clamp(MIN, MAX)
898}
899
900fn clamp_keep_alive_count(value: u32) -> u32 {
901 const MIN_CNT: u32 = 1;
902 const MAX_CNT: u32 = i8::MAX as u32;
904
905 value.clamp(MIN_CNT, MAX_CNT)
906}
907
908#[cfg(all(test, any(target_os = "linux", target_os = "macos")))]
909mod tests {
910 use super::*;
911 use crate::WasiCtxBuilder;
912
913 #[tokio::test]
914 async fn accepted_remote_address_survives_reset() {
915 let mut ctx = WasiCtxBuilder::new();
916 ctx.inherit_network().allow_tcp(true);
917 let ctx = ctx.build();
918 let mut socket = TcpSocket::new(&ctx.sockets, SocketAddressFamily::Ipv4).unwrap();
919 socket.bind("127.0.0.1:0".parse().unwrap()).await.unwrap();
920 let mut listener = socket.listen().await.unwrap();
921
922 let client = tokio::net::TcpStream::connect(socket.local_address().unwrap())
923 .await
924 .unwrap();
925 let peer = client.local_addr().unwrap();
926
927 poll_fn(|cx| listener.poll_ready(cx)).await;
929 client.set_zero_linger().unwrap();
930 drop(client);
931
932 let mut accepted = poll_fn(|cx| listener.poll_accept(cx)).await;
933 let mut input = accepted.take_receive_stream().unwrap();
934 let mut byte = [0];
935 let read = tokio::time::timeout(
936 Duration::from_secs(5),
937 poll_fn(|cx| input.poll_read(cx, &mut byte)),
938 )
939 .await
940 .unwrap();
941 assert!(matches!(read, Err(ErrorCode::ConnectionReset)));
942
943 let TcpState::Connected { stream, .. } = &accepted.tcp_state else {
946 panic!("expected an accepted connection");
947 };
948 assert!(stream.peer_addr().is_err());
949 assert_eq!(accepted.remote_address().unwrap(), peer);
950 }
951
952 #[tokio::test]
953 async fn connected_remote_address_survives_reset() {
954 let mut ctx = WasiCtxBuilder::new();
955 ctx.inherit_network().allow_tcp(true);
956 let ctx = ctx.build();
957 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
958 let peer = listener.local_addr().unwrap();
959
960 let mut client = TcpSocket::new(&ctx.sockets, SocketAddressFamily::Ipv4).unwrap();
961 client.start_connect(peer).unwrap();
962 poll_fn(|cx| client.poll_finish_connect(cx)).await.unwrap();
963 let (server, _) = listener.accept().await.unwrap();
964
965 assert_eq!(client.remote_address().unwrap(), peer);
967 server.set_zero_linger().unwrap();
968 drop(server);
969
970 let mut input = client.take_receive_stream().unwrap();
971 let mut byte = [0];
972 let read = tokio::time::timeout(
973 Duration::from_secs(5),
974 poll_fn(|cx| input.poll_read(cx, &mut byte)),
975 )
976 .await
977 .unwrap();
978 assert!(matches!(read, Err(ErrorCode::ConnectionReset)));
979
980 let TcpState::Connected { stream, .. } = &client.tcp_state else {
982 panic!("expected an outgoing connection");
983 };
984 assert!(stream.peer_addr().is_err());
985 assert_eq!(client.remote_address().unwrap(), peer);
986 }
987}