Skip to main content

moirai_async/net/
stream.rs

1use moirai_pal::net::AsyncTcpStream;
2use std::io;
3use std::net::SocketAddr;
4use std::pin::Pin;
5use std::sync::Arc;
6use std::task::{Context, Poll};
7
8use crate::io::{AsyncRead, AsyncWrite};
9use crate::net::resolve::resolve;
10use crate::net::types::{ConnectionId, ConnectionPool, ServerStats};
11
12/// Native async TCP stream with statistics tracking
13pub struct TcpStream {
14    inner: AsyncTcpStream,
15    stats: Arc<ServerStats>,
16    connection_pool: Arc<ConnectionPool>,
17    /// Pool tracking id assigned at accept time, or `None` for client-side
18    /// streams (`connect`/`from_std`) that are not pool-tracked. `Drop` removes
19    /// by this id rather than re-querying the socket, so a connection is
20    /// untracked exactly once even if the peer has already reset.
21    connection_id: Option<ConnectionId>,
22}
23
24impl TcpStream {
25    pub(super) fn new(
26        inner: AsyncTcpStream,
27        stats: Arc<ServerStats>,
28        connection_pool: Arc<ConnectionPool>,
29        connection_id: Option<ConnectionId>,
30    ) -> Self {
31        Self {
32            inner,
33            stats,
34            connection_pool,
35            connection_id,
36        }
37    }
38
39    /// Connect to a remote address asynchronously.
40    ///
41    /// `addr` is a literal socket address or `host:port`; hostnames resolve off
42    /// the polling thread (see `resolve`). Each resolved address is tried in
43    /// order and the first successful connection is returned. Neither
44    /// resolution nor the TCP handshake blocks the polling thread, so wrapping
45    /// this future in [`crate::timeout()`] or dropping it bounds the wait.
46    ///
47    /// # Errors
48    /// Returns the resolution error, or the connect error of the last address
49    /// tried when none accepts.
50    pub async fn connect(addr: &str) -> io::Result<Self> {
51        let (first, fallbacks) = resolve(addr).await?.into_parts();
52        let mut connected = AsyncTcpStream::connect(first).await;
53        for candidate in fallbacks {
54            if connected.is_ok() {
55                break;
56            }
57            connected = AsyncTcpStream::connect(candidate).await;
58        }
59        let inner = connected?;
60        let stats = Arc::new(ServerStats::default());
61        let connection_pool = Arc::new(ConnectionPool::new(None));
62
63        Ok(Self::new(inner, stats, connection_pool, None))
64    }
65
66    /// Wrap an existing TCP stream in the Moirai TCP facade.
67    pub fn from_std(stream: std::net::TcpStream) -> io::Result<Self> {
68        let inner = AsyncTcpStream::from_std(stream)?;
69        let stats = Arc::new(ServerStats::default());
70        let connection_pool = Arc::new(ConnectionPool::new(None));
71        Ok(Self::new(inner, stats, connection_pool, None))
72    }
73
74    /// Update per-connection pool tracking (byte counters, `last_activity`)
75    /// for pool-tracked (accept-side) streams; no-op for client-side streams.
76    fn record_io(&self, bytes_received: u64, bytes_sent: u64) {
77        if let Some(id) = self.connection_id {
78            self.connection_pool
79                .record_io(id, bytes_received, bytes_sent);
80        }
81    }
82
83    /// Read data from the stream
84    pub async fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
85        let bytes_read = self.inner.read(buf).await?;
86        self.stats
87            .bytes_received
88            .fetch_add(bytes_read as u64, std::sync::atomic::Ordering::Relaxed);
89        self.record_io(bytes_read as u64, 0);
90        Ok(bytes_read)
91    }
92
93    /// Write data to the stream
94    pub async fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
95        let bytes_written = self.inner.write(buf).await?;
96        self.stats
97            .bytes_sent
98            .fetch_add(bytes_written as u64, std::sync::atomic::Ordering::Relaxed);
99        self.record_io(0, bytes_written as u64);
100        Ok(bytes_written)
101    }
102
103    /// Flush the stream
104    pub async fn flush(&mut self) -> io::Result<()> {
105        self.inner.flush().await
106    }
107
108    /// Shutdown the write side of the stream.
109    pub async fn shutdown(&mut self) -> io::Result<()> {
110        self.inner.shutdown_write()
111    }
112
113    /// Get the peer address
114    pub fn peer_addr(&self) -> io::Result<SocketAddr> {
115        self.inner.peer_addr()
116    }
117
118    /// Get the local address
119    pub fn local_addr(&self) -> io::Result<SocketAddr> {
120        self.inner.local_addr()
121    }
122
123    /// Configure TCP_NODELAY on the stream.
124    pub fn set_nodelay(&self, on: bool) -> io::Result<()> {
125        self.inner.set_nodelay(on)
126    }
127}
128
129impl Drop for TcpStream {
130    fn drop(&mut self) {
131        // Untrack by the id captured at accept time. Re-querying `peer_addr()`
132        // here would fail on an already-reset socket and leak the pool slot plus
133        // the `active_connections` counter for the process lifetime.
134        let Some(id) = self.connection_id else {
135            return;
136        };
137
138        if self.connection_pool.remove_connection(id) {
139            self.stats
140                .active_connections
141                .fetch_sub(1, std::sync::atomic::Ordering::Relaxed);
142        }
143    }
144}
145
146impl AsyncRead for TcpStream {
147    fn poll_read(
148        mut self: Pin<&mut Self>,
149        cx: &mut Context<'_>,
150        buf: &mut [u8],
151    ) -> Poll<io::Result<usize>> {
152        match self.inner.poll_read(cx, buf) {
153            Poll::Ready(Ok(n)) => {
154                self.stats
155                    .bytes_received
156                    .fetch_add(n as u64, std::sync::atomic::Ordering::Relaxed);
157                self.record_io(n as u64, 0);
158                Poll::Ready(Ok(n))
159            }
160            res => res,
161        }
162    }
163}
164
165impl AsyncWrite for TcpStream {
166    fn poll_write(
167        mut self: Pin<&mut Self>,
168        cx: &mut Context<'_>,
169        buf: &[u8],
170    ) -> Poll<io::Result<usize>> {
171        match self.inner.poll_write(cx, buf) {
172            Poll::Ready(Ok(n)) => {
173                self.stats
174                    .bytes_sent
175                    .fetch_add(n as u64, std::sync::atomic::Ordering::Relaxed);
176                self.record_io(0, n as u64);
177                Poll::Ready(Ok(n))
178            }
179            res => res,
180        }
181    }
182
183    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
184        self.inner.poll_flush(cx)
185    }
186
187    fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
188        Poll::Ready(self.inner.shutdown_write())
189    }
190}