1use std::sync::atomic::{AtomicU64, Ordering};
2use std::sync::Arc;
3use std::time::Duration;
4
5use tokio::io::{self, AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
6use tokio::task::{AbortHandle, JoinSet};
7
8use crate::BoxStream;
9
10const RELAY_HALF_CLOSE_DRAIN: Duration = Duration::from_secs(1);
14
15const RELAY_ABORT_GRACE: Duration = Duration::from_secs(1);
17
18#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum TerminationReason {
24 ClientClosed,
25 ServerClosed,
26 BothClosed,
27 Error,
28}
29
30#[derive(Debug)]
32pub struct RelayResult {
33 pub bytes_upstream: u64,
34 pub bytes_downstream: u64,
35 pub termination_reason: TerminationReason,
36}
37
38#[derive(Debug, Clone, Copy)]
40enum Direction {
41 Upstream,
43 Downstream,
45}
46
47async fn copy_direction<R, W>(reader: &mut R, writer: &mut W, counter: &AtomicU64) -> io::Result<()>
48where
49 R: AsyncRead + Unpin,
50 W: AsyncWrite + Unpin,
51{
52 let mut buf = [0u8; 65536];
53 loop {
54 let n = reader.read(&mut buf).await?;
55 if n == 0 {
56 if let Err(error) = writer.shutdown().await {
57 if !matches!(
58 error.kind(),
59 io::ErrorKind::BrokenPipe | io::ErrorKind::ConnectionReset
60 ) {
61 return Err(error);
62 }
63 }
64 return Ok(());
65 }
66 writer.write_all(&buf[..n]).await?;
67 counter.fetch_add(n as u64, Ordering::Relaxed);
68 }
69}
70
71fn termination_reason(
72 first_closed: Option<TerminationReason>,
73 had_error: bool,
74 drain_timed_out: bool,
75) -> TerminationReason {
76 if had_error {
77 TerminationReason::Error
78 } else if drain_timed_out {
79 first_closed.unwrap_or(TerminationReason::Error)
80 } else {
81 first_closed.unwrap_or(TerminationReason::BothClosed)
82 }
83}
84
85pub async fn relay(client: BoxStream, server: BoxStream) -> RelayResult {
90 let (mut client_read, mut client_write) = io::split(client);
91 let (mut server_read, mut server_write) = io::split(server);
92
93 let bytes_upstream = Arc::new(AtomicU64::new(0));
94 let bytes_downstream = Arc::new(AtomicU64::new(0));
95 let mut tasks = JoinSet::new();
96
97 let upstream_counter = Arc::clone(&bytes_upstream);
98 let upstream_abort: AbortHandle = tasks.spawn(async move {
99 let result = copy_direction(&mut client_read, &mut server_write, &upstream_counter).await;
100 (Direction::Upstream, result)
101 });
102
103 let downstream_counter = Arc::clone(&bytes_downstream);
104 let downstream_abort: AbortHandle = tasks.spawn(async move {
105 let result = copy_direction(&mut server_read, &mut client_write, &downstream_counter).await;
106 (Direction::Downstream, result)
107 });
108
109 let mut had_error = false;
110 let mut first_closed: Option<TerminationReason> = None;
113 let mut pending_abort: Option<AbortHandle> = None;
117 let mut drain_timed_out = false;
118
119 match tasks.join_next().await {
120 Some(Ok((direction, Ok(())))) => {
121 let reason = match direction {
122 Direction::Upstream => {
123 pending_abort = Some(downstream_abort.clone());
124 TerminationReason::ClientClosed
125 }
126 Direction::Downstream => {
127 pending_abort = Some(upstream_abort.clone());
128 TerminationReason::ServerClosed
129 }
130 };
131 first_closed = Some(reason);
132 }
133 Some(Ok((direction, Err(error)))) => {
134 tracing::debug!(%error, ?direction, "relay direction failed");
135 had_error = true;
136 upstream_abort.abort();
137 downstream_abort.abort();
138 }
139 Some(Err(error)) => {
140 tracing::debug!(%error, "relay direction task failed");
141 had_error = true;
142 upstream_abort.abort();
143 downstream_abort.abort();
144 }
145 None => {}
146 }
147
148 if let Some(abort) = pending_abort.as_ref() {
149 match tokio::time::timeout(RELAY_HALF_CLOSE_DRAIN, tasks.join_next()).await {
150 Ok(Some(Ok((_, Ok(()))))) => {}
151 Ok(Some(Ok((direction, Err(error))))) => {
152 tracing::debug!(%error, ?direction, "relay direction failed during drain");
153 had_error = true;
154 }
155 Ok(Some(Err(error))) => {
156 tracing::debug!(%error, "relay direction task failed during drain");
157 had_error = true;
158 }
159 Ok(None) => {}
160 Err(_) => {
161 drain_timed_out = true;
162 abort.abort();
163 let _ = tokio::time::timeout(RELAY_ABORT_GRACE, tasks.join_next()).await;
164 }
165 }
166 }
167
168 if !drain_timed_out {
169 while let Some(outcome) = tasks.join_next().await {
170 match outcome {
171 Ok((direction, Err(error))) => {
172 tracing::debug!(%error, ?direction, "relay direction failed");
173 had_error = true;
174 }
175 Err(error) => {
176 tracing::debug!(%error, "relay direction task failed");
177 had_error = true;
178 }
179 Ok((_, Ok(()))) => {}
180 }
181 }
182 }
183
184 let termination_reason = termination_reason(first_closed, had_error, drain_timed_out);
185
186 RelayResult {
187 bytes_upstream: bytes_upstream.load(Ordering::Relaxed),
188 bytes_downstream: bytes_downstream.load(Ordering::Relaxed),
189 termination_reason,
190 }
191}
192
193#[cfg(test)]
194mod tests {
195 use super::*;
196 use tokio::io::{AsyncReadExt, AsyncWriteExt};
197
198 #[test]
199 fn relay_error_during_drain_takes_precedence_over_close_reason() {
200 assert_eq!(
201 termination_reason(Some(TerminationReason::ClientClosed), true, true),
202 TerminationReason::Error
203 );
204 }
205
206 #[tokio::test]
207 async fn test_relay_echo() {
208 let echo = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
209 let echo_addr = echo.local_addr().unwrap();
210
211 let jh = tokio::spawn(async move {
212 let (stream, _) = echo.accept().await.unwrap();
213 let (mut reader, mut writer) = stream.into_split();
214 tokio::spawn(async move {
215 let mut buf = [0u8; 1024];
216 loop {
217 let n = reader.read(&mut buf).await.unwrap();
218 if n == 0 {
219 break;
220 }
221 writer.write_all(&buf[..n]).await.unwrap();
222 }
223 });
224 });
225
226 let proxy_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
227 let proxy_addr = proxy_listener.local_addr().unwrap();
228
229 let proxy_jh = tokio::spawn(async move {
230 let (client_stream, _) = proxy_listener.accept().await.unwrap();
231 let server_stream = tokio::net::TcpStream::connect(echo_addr).await.unwrap();
232 relay(Box::new(client_stream), Box::new(server_stream)).await
233 });
234
235 let mut client = tokio::net::TcpStream::connect(proxy_addr).await.unwrap();
236 client.write_all(b"hello relay").await.unwrap();
237 client.shutdown().await.unwrap();
238
239 let mut buf = String::new();
240 client.read_to_string(&mut buf).await.unwrap();
241 assert_eq!(buf, "hello relay");
242
243 let result = proxy_jh.await.unwrap();
244 assert_eq!(result.bytes_upstream, 11);
245 assert_eq!(result.bytes_downstream, 11);
246
247 jh.await.unwrap();
248 }
249
250 #[tokio::test]
251 async fn test_relay_half_close() {
252 let echo = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
253 let echo_addr = echo.local_addr().unwrap();
254
255 let jh = tokio::spawn(async move {
256 let (mut stream, _) = echo.accept().await.unwrap();
257 let mut buf = [0u8; 1024];
258 let n = stream.read(&mut buf).await.unwrap();
259 stream.write_all(&buf[..n]).await.unwrap();
260 stream.shutdown().await.unwrap();
261 });
262
263 let proxy_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
264 let proxy_addr = proxy_listener.local_addr().unwrap();
265
266 let proxy_jh = tokio::spawn(async move {
267 let (client_stream, _) = proxy_listener.accept().await.unwrap();
268 let server_stream = tokio::net::TcpStream::connect(echo_addr).await.unwrap();
269 relay(Box::new(client_stream), Box::new(server_stream)).await
270 });
271
272 let mut client = tokio::net::TcpStream::connect(proxy_addr).await.unwrap();
273 client.write_all(b"data").await.unwrap();
274 client.shutdown().await.unwrap();
275
276 let mut buf = [0u8; 4];
277 client.read_exact(&mut buf).await.unwrap();
278 assert_eq!(&buf, b"data");
279
280 let result = proxy_jh.await.unwrap();
281 assert_eq!(result.bytes_upstream, 4);
282 assert_eq!(result.bytes_downstream, 4);
283 assert_eq!(result.termination_reason, TerminationReason::ClientClosed);
285
286 jh.await.unwrap();
287 }
288
289 #[tokio::test]
290 async fn test_relay_half_close_server_hangs() {
291 let upstream = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
295 let upstream_addr = upstream.local_addr().unwrap();
296
297 let upstream_jh = tokio::spawn(async move {
298 let (mut stream, _) = upstream.accept().await.unwrap();
299 let mut buf = [0u8; 64];
300 let _ = stream.read(&mut buf).await.unwrap();
301 std::future::pending::<()>().await;
302 });
303
304 let proxy_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
305 let proxy_addr = proxy_listener.local_addr().unwrap();
306
307 let proxy_jh = tokio::spawn(async move {
308 let (client_stream, _) = proxy_listener.accept().await.unwrap();
309 let server_stream = tokio::net::TcpStream::connect(upstream_addr).await.unwrap();
310 relay(Box::new(client_stream), Box::new(server_stream)).await
311 });
312
313 let mut client = tokio::net::TcpStream::connect(proxy_addr).await.unwrap();
314 client.write_all(b"data").await.unwrap();
315 client.shutdown().await.unwrap();
316
317 let result = tokio::time::timeout(std::time::Duration::from_secs(5), proxy_jh)
318 .await
319 .expect("relay should not block forever on a hanging upstream")
320 .unwrap();
321 assert_eq!(result.bytes_upstream, 4);
322 assert_eq!(result.bytes_downstream, 0);
323 assert_eq!(result.termination_reason, TerminationReason::ClientClosed);
324
325 upstream_jh.abort();
326 }
327
328 #[tokio::test]
329 async fn test_relay_cancellation() {
330 let echo = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
331 let echo_addr = echo.local_addr().unwrap();
332
333 let jh = tokio::spawn(async move {
334 let (stream, _) = echo.accept().await.unwrap();
335 let (mut reader, mut writer) = stream.into_split();
336 tokio::spawn(async move {
337 let mut buf = [0u8; 1024];
338 loop {
339 let n = reader.read(&mut buf).await.unwrap();
340 if n == 0 {
341 break;
342 }
343 writer.write_all(&buf[..n]).await.unwrap();
344 }
345 });
346 });
347
348 let proxy_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
349 let proxy_addr = proxy_listener.local_addr().unwrap();
350
351 let proxy_jh = tokio::spawn(async move {
352 let (client_stream, _) = proxy_listener.accept().await.unwrap();
353 let server_stream = tokio::net::TcpStream::connect(echo_addr).await.unwrap();
354 relay(Box::new(client_stream), Box::new(server_stream)).await
355 });
356
357 let mut client = tokio::net::TcpStream::connect(proxy_addr).await.unwrap();
358 client.write_all(b"data").await.unwrap();
359 drop(client);
360
361 let result = proxy_jh.await.unwrap();
362 assert!(result.bytes_upstream > 0 || result.bytes_downstream > 0);
363
364 jh.await.unwrap();
365 }
366}