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::resolve::resolve;
10use crate::net::types::{ConnectionId, ConnectionPool, ServerStats};
11
12pub struct TcpStream {
14 inner: AsyncTcpStream,
15 stats: Arc<ServerStats>,
16 connection_pool: Arc<ConnectionPool>,
17 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 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 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 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 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 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 pub async fn flush(&mut self) -> io::Result<()> {
105 self.inner.flush().await
106 }
107
108 pub async fn shutdown(&mut self) -> io::Result<()> {
110 self.inner.shutdown_write()
111 }
112
113 pub fn peer_addr(&self) -> io::Result<SocketAddr> {
115 self.inner.peer_addr()
116 }
117
118 pub fn local_addr(&self) -> io::Result<SocketAddr> {
120 self.inner.local_addr()
121 }
122
123 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 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}