Skip to main content

eggress_core/
relay.rs

1use tokio::io::{self, AsyncWriteExt};
2
3use crate::BoxStream;
4
5/// Reason the relay terminated.
6#[derive(Debug, Clone, Copy, PartialEq, Eq)]
7pub enum TerminationReason {
8    ClientClosed,
9    ServerClosed,
10    BothClosed,
11    Error,
12    Cancelled,
13}
14
15/// Result of a relay operation.
16#[derive(Debug)]
17pub struct RelayResult {
18    pub bytes_upstream: u64,
19    pub bytes_downstream: u64,
20    pub termination_reason: TerminationReason,
21}
22
23/// Relay data bidirectionally between two streams.
24///
25/// When one side closes its write half, the other side's write half is shut down
26/// (half-close semantics). Both directions must complete before returning.
27pub async fn relay(client: BoxStream, server: BoxStream) -> RelayResult {
28    let (mut client_read, mut client_write) = io::split(client);
29    let (mut server_read, mut server_write) = io::split(server);
30
31    let client_to_server = tokio::spawn(async move {
32        let n = io::copy(&mut client_read, &mut server_write).await?;
33        server_write.shutdown().await?;
34        Ok::<u64, std::io::Error>(n)
35    });
36
37    let server_to_client = tokio::spawn(async move {
38        let n = io::copy(&mut server_read, &mut client_write).await?;
39        client_write.shutdown().await?;
40        Ok::<u64, std::io::Error>(n)
41    });
42
43    let a_result = client_to_server.await;
44    let b_result = server_to_client.await;
45
46    let (a_bytes, a_error) = match a_result {
47        Ok(Ok(n)) => (n, false),
48        Ok(Err(_)) => (0, true),
49        Err(_) => (0, true),
50    };
51
52    let (b_bytes, b_error) = match b_result {
53        Ok(Ok(n)) => (n, false),
54        Ok(Err(_)) => (0, true),
55        Err(_) => (0, true),
56    };
57
58    let termination_reason = match (a_error, b_error) {
59        (true, true) => TerminationReason::Error,
60        (true, false) => TerminationReason::ClientClosed,
61        (false, true) => TerminationReason::ServerClosed,
62        (false, false) => TerminationReason::BothClosed,
63    };
64
65    RelayResult {
66        bytes_upstream: a_bytes,
67        bytes_downstream: b_bytes,
68        termination_reason,
69    }
70}
71
72#[cfg(test)]
73mod tests {
74    use super::*;
75    use tokio::io::{AsyncReadExt, AsyncWriteExt};
76
77    #[tokio::test]
78    async fn test_relay_echo() {
79        let echo = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
80        let echo_addr = echo.local_addr().unwrap();
81
82        let jh = tokio::spawn(async move {
83            let (stream, _) = echo.accept().await.unwrap();
84            let (mut reader, mut writer) = stream.into_split();
85            tokio::spawn(async move {
86                let mut buf = [0u8; 1024];
87                loop {
88                    let n = reader.read(&mut buf).await.unwrap();
89                    if n == 0 {
90                        break;
91                    }
92                    writer.write_all(&buf[..n]).await.unwrap();
93                }
94            });
95        });
96
97        let proxy_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
98        let proxy_addr = proxy_listener.local_addr().unwrap();
99
100        let proxy_jh = tokio::spawn(async move {
101            let (client_stream, _) = proxy_listener.accept().await.unwrap();
102            let server_stream = tokio::net::TcpStream::connect(echo_addr).await.unwrap();
103            relay(Box::new(client_stream), Box::new(server_stream)).await
104        });
105
106        let mut client = tokio::net::TcpStream::connect(proxy_addr).await.unwrap();
107        client.write_all(b"hello relay").await.unwrap();
108        client.shutdown().await.unwrap();
109
110        let mut buf = String::new();
111        client.read_to_string(&mut buf).await.unwrap();
112        assert_eq!(buf, "hello relay");
113
114        let result = proxy_jh.await.unwrap();
115        assert_eq!(result.bytes_upstream, 11);
116        assert_eq!(result.bytes_downstream, 11);
117
118        jh.await.unwrap();
119    }
120
121    #[tokio::test]
122    async fn test_relay_half_close() {
123        let echo = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
124        let echo_addr = echo.local_addr().unwrap();
125
126        let jh = tokio::spawn(async move {
127            let (mut stream, _) = echo.accept().await.unwrap();
128            let mut buf = [0u8; 1024];
129            let n = stream.read(&mut buf).await.unwrap();
130            stream.write_all(&buf[..n]).await.unwrap();
131            stream.shutdown().await.unwrap();
132        });
133
134        let proxy_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
135        let proxy_addr = proxy_listener.local_addr().unwrap();
136
137        let proxy_jh = tokio::spawn(async move {
138            let (client_stream, _) = proxy_listener.accept().await.unwrap();
139            let server_stream = tokio::net::TcpStream::connect(echo_addr).await.unwrap();
140            relay(Box::new(client_stream), Box::new(server_stream)).await
141        });
142
143        let mut client = tokio::net::TcpStream::connect(proxy_addr).await.unwrap();
144        client.write_all(b"data").await.unwrap();
145        client.shutdown().await.unwrap();
146
147        let mut buf = [0u8; 4];
148        client.read_exact(&mut buf).await.unwrap();
149        assert_eq!(&buf, b"data");
150
151        let result = proxy_jh.await.unwrap();
152        assert_eq!(result.bytes_upstream, 4);
153        assert_eq!(result.bytes_downstream, 4);
154
155        jh.await.unwrap();
156    }
157
158    #[tokio::test]
159    async fn test_relay_cancellation() {
160        let echo = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
161        let echo_addr = echo.local_addr().unwrap();
162
163        let jh = tokio::spawn(async move {
164            let (stream, _) = echo.accept().await.unwrap();
165            let (mut reader, mut writer) = stream.into_split();
166            tokio::spawn(async move {
167                let mut buf = [0u8; 1024];
168                loop {
169                    let n = reader.read(&mut buf).await.unwrap();
170                    if n == 0 {
171                        break;
172                    }
173                    writer.write_all(&buf[..n]).await.unwrap();
174                }
175            });
176        });
177
178        let proxy_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
179        let proxy_addr = proxy_listener.local_addr().unwrap();
180
181        let proxy_jh = tokio::spawn(async move {
182            let (client_stream, _) = proxy_listener.accept().await.unwrap();
183            let server_stream = tokio::net::TcpStream::connect(echo_addr).await.unwrap();
184            relay(Box::new(client_stream), Box::new(server_stream)).await
185        });
186
187        let mut client = tokio::net::TcpStream::connect(proxy_addr).await.unwrap();
188        client.write_all(b"data").await.unwrap();
189        drop(client);
190
191        let result = proxy_jh.await.unwrap();
192        assert!(result.bytes_upstream > 0 || result.bytes_downstream > 0);
193
194        jh.await.unwrap();
195    }
196}