Skip to main content

eggress_core/
relay.rs

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
10/// Time to wait for the opposite direction to drain naturally after one side
11/// completes. Without this, a half-closing peer whose FIN is never echoed by
12/// the upstream causes the relay to block forever on the other side's read.
13const RELAY_HALF_CLOSE_DRAIN: Duration = Duration::from_secs(1);
14
15/// Upper bound for a forced abort to take effect after the drain timeout.
16const RELAY_ABORT_GRACE: Duration = Duration::from_secs(1);
17
18/// Reason the relay terminated.
19///
20/// `ClientClosed`/`ServerClosed` report which side hung up first when both
21/// directions completed cleanly; `Error` means at least one direction failed.
22#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum TerminationReason {
24    ClientClosed,
25    ServerClosed,
26    BothClosed,
27    Error,
28}
29
30/// Result of a relay operation.
31#[derive(Debug)]
32pub struct RelayResult {
33    pub bytes_upstream: u64,
34    pub bytes_downstream: u64,
35    pub termination_reason: TerminationReason,
36}
37
38/// Which relay direction a spawned task was copying.
39#[derive(Debug, Clone, Copy)]
40enum Direction {
41    /// Client → server (upstream).
42    Upstream,
43    /// Server → client (downstream).
44    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
85/// Relay data bidirectionally between two streams.
86///
87/// When one side closes its write half, the other side's write half is shut down
88/// (half-close semantics). Both directions must complete before returning.
89pub 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    // Records the first direction to finish cleanly, so diagnostics can tell
111    // which side hung up first.
112    let mut first_closed: Option<TerminationReason> = None;
113    // When half-close leaves one side's reader stuck on a peer that never
114    // answers the FIN we sent, the relay aborts the surviving direction after
115    // a short drain window so the connection does not leak.
116    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        // The client closed its write half first.
284        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        // The upstream reads the client payload and then never writes or closes.
292        // Without the half-close drain + abort, the relay would block forever
293        // on the downstream reader waiting for an EOF that never arrives.
294        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}