moirai_async/net/
stream.rs1use 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
11pub struct TcpStream {
13 inner: AsyncTcpStream,
14 stats: Arc<ServerStats>,
15 connection_pool: Arc<ConnectionPool>,
16 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 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 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 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 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 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 pub async fn flush(&mut self) -> io::Result<()> {
90 self.inner.flush().await
91 }
92
93 pub async fn shutdown(&mut self) -> io::Result<()> {
95 self.inner.shutdown_write()
96 }
97
98 pub fn peer_addr(&self) -> io::Result<SocketAddr> {
100 self.inner.peer_addr()
101 }
102
103 pub fn local_addr(&self) -> io::Result<SocketAddr> {
105 self.inner.local_addr()
106 }
107
108 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 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}