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