1use tokio::io::{self, AsyncWriteExt};
2
3use crate::BoxStream;
4
5#[derive(Debug, Clone, Copy, PartialEq, Eq)]
7pub enum TerminationReason {
8 ClientClosed,
9 ServerClosed,
10 BothClosed,
11 Error,
12 Cancelled,
13}
14
15#[derive(Debug)]
17pub struct RelayResult {
18 pub bytes_upstream: u64,
19 pub bytes_downstream: u64,
20 pub termination_reason: TerminationReason,
21}
22
23pub 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}