1use bytes::Bytes;
2use snafu::ResultExt;
3use std::future::Future;
4use std::pin::Pin;
5use std::time::Duration;
6use tokio::io::{AsyncReadExt, AsyncWriteExt};
7use tokio::time::Instant;
8
9use super::{
10 CodecMessageReader, CodecMessageWriter, MessageReader, MessageWriter, NormalMessageReader,
11 NormalMessageWriter,
12};
13use crate::buffer::{BufferReader, BufferedReader};
14use pb_mapper_core::checksum::AesKeyType;
15use pb_mapper_core::codec::{Decryptor, Encryptor};
16use pb_mapper_core::config::duration_from_env;
17use pb_mapper_core::error::{FwdNetworkWriteWithNormalSnafu, Result};
18use pb_mapper_core::snafu_error_get_or_return_ok;
19use uni_stream::stream::{StreamSplit, TcpStreamImpl, UdpStreamImpl};
20use uni_stream::udp::{UdpStreamReadHalf, UdpStreamWriteHalf};
21
22pub trait ForwardReader {
23 async fn read(&mut self) -> Result<&'_ [u8]>;
24}
25
26pub trait ForwardWriter {
27 async fn write(&mut self, src: &[u8]) -> Result<()>;
28
29 async fn shutdown(&mut self);
31}
32
33pub trait DatagramReader {
34 async fn recv(&mut self) -> Result<Bytes>;
35}
36
37pub trait DatagramWriter {
38 async fn send(&mut self, src: &[u8]) -> Result<()>;
39}
40
41const DEFAULT_TUNNEL_IDLE_TIMEOUT: Duration = Duration::from_secs(60 * 60);
42const DEFAULT_HALF_CLOSE_IDLE_TIMEOUT: Duration = Duration::from_secs(60);
43const PB_MAPPER_TUNNEL_IDLE_TIMEOUT: &str = "PB_MAPPER_TUNNEL_IDLE_TIMEOUT";
44const PB_MAPPER_HALF_CLOSE_IDLE_TIMEOUT: &str = "PB_MAPPER_HALF_CLOSE_IDLE_TIMEOUT";
45
46#[derive(Debug, Clone, Copy)]
47struct ForwardTimeoutConfig {
48 tunnel_idle_timeout: Duration,
49 half_close_idle_timeout: Duration,
50}
51
52impl ForwardTimeoutConfig {
53 fn from_env() -> Self {
54 Self {
55 tunnel_idle_timeout: duration_from_env(
56 PB_MAPPER_TUNNEL_IDLE_TIMEOUT,
57 DEFAULT_TUNNEL_IDLE_TIMEOUT,
58 ),
59 half_close_idle_timeout: duration_from_env(
60 PB_MAPPER_HALF_CLOSE_IDLE_TIMEOUT,
61 DEFAULT_HALF_CLOSE_IDLE_TIMEOUT,
62 ),
63 }
64 }
65}
66
67pub struct NormalForwardReader<'a, T> {
68 buffered_reader: BufferReader<'a, T>,
69}
70
71impl<'a, T: AsyncReadExt + Unpin + Send> NormalForwardReader<'a, T> {
72 pub fn new(reader: &'a mut T) -> Self {
73 Self {
74 buffered_reader: BufferReader::new(reader),
75 }
76 }
77}
78
79impl<'a, T: AsyncReadExt + Unpin + Send> ForwardReader for NormalForwardReader<'a, T> {
80 async fn read(&mut self) -> Result<&'_ [u8]> {
81 self.buffered_reader.read().await
82 }
83}
84
85pub struct NormalDatagramReader<'a, T: AsyncReadExt + Unpin> {
86 reader: NormalMessageReader<'a, T>,
87}
88
89impl<'a, T: AsyncReadExt + Unpin + Send> NormalDatagramReader<'a, T> {
90 pub fn new(reader: &'a mut T) -> Self {
91 Self {
92 reader: NormalMessageReader::new(reader),
93 }
94 }
95
96 pub fn with_checksum_key(self, key: AesKeyType) -> Self {
97 Self {
98 reader: self.reader.with_checksum_key(key),
99 }
100 }
101}
102
103impl<'a, T: AsyncReadExt + Unpin + Send> DatagramReader for NormalDatagramReader<'a, T> {
104 async fn recv(&mut self) -> Result<Bytes> {
105 let msg = self.reader.read_msg().await?;
106 Ok(Bytes::copy_from_slice(msg))
107 }
108}
109
110pub struct NormalForwardWriter<'a, T> {
111 writer: &'a mut T,
112}
113
114impl<'a, T: AsyncWriteExt + Unpin + Send> NormalForwardWriter<'a, T> {
115 pub fn new(writer: &'a mut T) -> Self {
116 Self { writer }
117 }
118
119 async fn write_inner(&mut self, src: &[u8]) -> Result<()> {
120 self.writer
121 .write_all(src)
122 .await
123 .context(FwdNetworkWriteWithNormalSnafu)
124 }
125}
126
127impl<'a, T: AsyncWriteExt + Unpin + Send> ForwardWriter for NormalForwardWriter<'a, T> {
128 async fn write(&mut self, src: &[u8]) -> Result<()> {
129 self.write_inner(src).await
130 }
131
132 async fn shutdown(&mut self) {
133 let _ = self.writer.shutdown().await;
134 }
135}
136
137pub struct NormalDatagramWriter<'a, T: AsyncWriteExt + Unpin> {
138 writer: NormalMessageWriter<'a, T>,
139}
140
141impl<'a, T: AsyncWriteExt + Unpin + Send> NormalDatagramWriter<'a, T> {
142 pub fn new(writer: &'a mut T) -> Self {
143 Self {
144 writer: NormalMessageWriter::new(writer),
145 }
146 }
147
148 pub fn with_checksum_key(self, key: AesKeyType) -> Self {
149 Self {
150 writer: self.writer.with_checksum_key(key),
151 }
152 }
153}
154
155impl<'a, T: AsyncWriteExt + Unpin + Send> DatagramWriter for NormalDatagramWriter<'a, T> {
156 async fn send(&mut self, src: &[u8]) -> Result<()> {
157 self.writer.write_msg(src).await
158 }
159}
160
161pub struct CodecForwardReader<'a, T: AsyncReadExt + Unpin + Send, D: Decryptor>(
162 CodecMessageReader<'a, T, D>,
163);
164
165impl<'a, T: AsyncReadExt + Send + Unpin, D: Decryptor> CodecForwardReader<'a, T, D> {
166 pub fn new(reader: &'a mut T, decryptor: D) -> Self {
167 Self(CodecMessageReader::new(reader, decryptor))
168 }
169
170 pub fn with_checksum_key(self, key: AesKeyType) -> Self {
171 Self(self.0.with_checksum_key(key))
172 }
173}
174
175impl<'a, T: AsyncReadExt + Send + Unpin, D: Decryptor> ForwardReader
176 for CodecForwardReader<'a, T, D>
177{
178 async fn read(&mut self) -> Result<&'_ [u8]> {
179 self.0.read_msg().await
180 }
181}
182
183pub struct CodecDatagramReader<'a, T: AsyncReadExt + Unpin + Send, D: Decryptor>(
184 CodecMessageReader<'a, T, D>,
185);
186
187impl<'a, T: AsyncReadExt + Send + Unpin, D: Decryptor> CodecDatagramReader<'a, T, D> {
188 pub fn new(reader: &'a mut T, decryptor: D) -> Self {
189 Self(CodecMessageReader::new(reader, decryptor))
190 }
191
192 pub fn with_checksum_key(self, key: AesKeyType) -> Self {
193 Self(self.0.with_checksum_key(key))
194 }
195}
196
197impl<'a, T: AsyncReadExt + Send + Unpin, D: Decryptor> DatagramReader
198 for CodecDatagramReader<'a, T, D>
199{
200 async fn recv(&mut self) -> Result<Bytes> {
201 let msg = self.0.read_msg().await?;
202 Ok(Bytes::copy_from_slice(msg))
203 }
204}
205
206pub struct CodecForwardWriter<'a, T: AsyncWriteExt + Send + Unpin, E: Encryptor>(
207 CodecMessageWriter<'a, T, E>,
208);
209
210impl<'a, T: AsyncWriteExt + Send + Unpin, E: Encryptor> CodecForwardWriter<'a, T, E> {
211 pub fn new(writer: &'a mut T, encryptor: E) -> Self {
212 Self(CodecMessageWriter::new(writer, encryptor))
213 }
214
215 pub fn with_checksum_key(self, key: AesKeyType) -> Self {
216 Self(self.0.with_checksum_key(key))
217 }
218}
219
220impl<'a, T: AsyncWriteExt + Send + Unpin, E: Encryptor> ForwardWriter
221 for CodecForwardWriter<'a, T, E>
222{
223 async fn write(&mut self, src: &[u8]) -> Result<()> {
225 self.0.write_msg(src).await
226 }
227
228 async fn shutdown(&mut self) {
229 let _ = self.0.shutdown().await;
230 }
231}
232
233pub struct CodecDatagramWriter<'a, T: AsyncWriteExt + Send + Unpin, E: Encryptor>(
234 CodecMessageWriter<'a, T, E>,
235);
236
237impl<'a, T: AsyncWriteExt + Send + Unpin, E: Encryptor> CodecDatagramWriter<'a, T, E> {
238 pub fn new(writer: &'a mut T, encryptor: E) -> Self {
239 Self(CodecMessageWriter::new(writer, encryptor))
240 }
241
242 pub fn with_checksum_key(self, key: AesKeyType) -> Self {
243 Self(self.0.with_checksum_key(key))
244 }
245}
246
247impl<'a, T: AsyncWriteExt + Send + Unpin, E: Encryptor> DatagramWriter
248 for CodecDatagramWriter<'a, T, E>
249{
250 async fn send(&mut self, src: &[u8]) -> Result<()> {
251 let buf = src.to_vec();
252 self.0.write_msg(&buf).await
253 }
254}
255
256pub async fn copy<R: ForwardReader, W: ForwardWriter>(
257 mut reader: R,
258 mut writer: W,
259) -> Result<usize> {
260 let mut length: usize = 0;
261 loop {
262 let src = reader.read().await?;
263 let n = src.len();
264 if n == 0 {
265 break;
266 }
267 writer.write(src).await?;
268 length += n;
269 }
270 writer.shutdown().await;
271 Ok(length)
272}
273
274pub async fn transfer_datagrams<R: DatagramReader, W: DatagramWriter>(
275 label: &'static str,
276 mut reader: R,
277 mut writer: W,
278) -> Result<usize> {
279 let mut _length: usize = 0;
280 loop {
281 let src = reader.recv().await?;
282 let n = src.len();
283 tracing::debug!("datagram forward {label} {n} bytes");
284 writer.send(&src).await?;
285 _length += n;
286 }
287}
288
289pub async fn start_forward<
290 ClientReader: ForwardReader,
291 ClientWriter: ForwardWriter,
292 ServerReader: ForwardReader,
293 ServerWriter: ForwardWriter,
294>(
295 client_reader: ClientReader,
296 client_writer: ClientWriter,
297 server_reader: ServerReader,
298 server_writer: ServerWriter,
299) {
300 start_forward_with_config(
301 client_reader,
302 client_writer,
303 server_reader,
304 server_writer,
305 ForwardTimeoutConfig::from_env(),
306 )
307 .await
308}
309
310#[derive(Default)]
311struct ForwardDirectionState {
312 len: usize,
313 result: Option<Result<usize>>,
314}
315
316impl ForwardDirectionState {
317 fn is_done(&self) -> bool {
318 self.result.is_some()
319 }
320}
321
322#[derive(Clone, Copy)]
323struct ForwardActivity {
324 bytes: usize,
325 at: Instant,
326}
327
328async fn copy_with_activity<R: ForwardReader, W: ForwardWriter>(
329 mut reader: R,
330 mut writer: W,
331 activity_tx: tokio::sync::mpsc::UnboundedSender<ForwardActivity>,
332) -> Result<usize> {
333 let mut length = 0;
334 loop {
335 let src = reader.read().await?;
336 let n = src.len();
337 if n == 0 {
338 break;
339 }
340 writer.write(src).await?;
341 length += n;
342 let _ = activity_tx.send(ForwardActivity {
343 bytes: n,
344 at: Instant::now(),
345 });
346 }
347 writer.shutdown().await;
348 Ok(length)
349}
350
351async fn start_forward_with_config<
352 ClientReader: ForwardReader,
353 ClientWriter: ForwardWriter,
354 ServerReader: ForwardReader,
355 ServerWriter: ForwardWriter,
356>(
357 client_reader: ClientReader,
358 client_writer: ClientWriter,
359 server_reader: ServerReader,
360 server_writer: ServerWriter,
361 timeout_config: ForwardTimeoutConfig,
362) {
363 let tunnel_idle_enabled = !timeout_config.tunnel_idle_timeout.is_zero();
364 let half_close_idle_enabled = !timeout_config.half_close_idle_timeout.is_zero();
365 let tunnel_idle_sleep = tokio::time::sleep(timeout_config.tunnel_idle_timeout);
366 let half_close_idle_sleep = tokio::time::sleep(timeout_config.half_close_idle_timeout);
367 tokio::pin!(tunnel_idle_sleep);
368 tokio::pin!(half_close_idle_sleep);
369
370 let (client_activity_tx, mut client_activity_rx) = tokio::sync::mpsc::unbounded_channel();
371 let (server_activity_tx, mut server_activity_rx) = tokio::sync::mpsc::unbounded_channel();
372 let client_to_server = copy_with_activity(client_reader, server_writer, client_activity_tx);
376 let server_to_client = copy_with_activity(server_reader, client_writer, server_activity_tx);
377 tokio::pin!(client_to_server);
378 tokio::pin!(server_to_client);
379
380 let mut client_state = ForwardDirectionState::default();
381 let mut server_state = ForwardDirectionState::default();
382
383 loop {
384 let client_done = client_state.is_done();
385 let server_done = server_state.is_done();
386 let half_closed = client_done ^ server_done;
387 if client_done && server_done {
388 break;
389 }
390
391 tokio::select! {
392 biased;
393
394 result = &mut client_to_server, if !client_done => {
395 let failed = result.is_err();
396 client_state.result = Some(result);
397 reset_sleep(&mut half_close_idle_sleep, timeout_config.half_close_idle_timeout);
398 if failed {
399 break;
400 }
401 }
402 result = &mut server_to_client, if !server_done => {
403 let failed = result.is_err();
404 server_state.result = Some(result);
405 reset_sleep(&mut half_close_idle_sleep, timeout_config.half_close_idle_timeout);
406 if failed {
407 break;
408 }
409 }
410 Some(activity) = client_activity_rx.recv(), if !client_done => {
411 record_forward_activity(
412 activity,
413 &mut client_state,
414 &mut tunnel_idle_sleep,
415 timeout_config.tunnel_idle_timeout,
416 &mut half_close_idle_sleep,
417 timeout_config.half_close_idle_timeout,
418 server_done,
419 );
420 }
421 Some(activity) = server_activity_rx.recv(), if !server_done => {
422 record_forward_activity(
423 activity,
424 &mut server_state,
425 &mut tunnel_idle_sleep,
426 timeout_config.tunnel_idle_timeout,
427 &mut half_close_idle_sleep,
428 timeout_config.half_close_idle_timeout,
429 client_done,
430 );
431 }
432 _ = &mut tunnel_idle_sleep, if tunnel_idle_enabled && !half_closed => {
433 tracing::debug!(
434 "forward tunnel idle timeout after {:?}",
435 timeout_config.tunnel_idle_timeout
436 );
437 break;
438 }
439 _ = &mut half_close_idle_sleep, if half_close_idle_enabled && half_closed => {
440 tracing::debug!(
441 "forward half-close idle timeout after {:?}",
442 timeout_config.half_close_idle_timeout
443 );
444 break;
445 }
446 }
447 }
448
449 let client_len = client_state.len;
450 let server_len = server_state.len;
451 handle_forward_final_result(client_state.result, client_len, "client->server");
452 handle_forward_final_result(server_state.result, server_len, "server->client");
453}
454
455fn record_forward_activity(
456 activity: ForwardActivity,
457 state: &mut ForwardDirectionState,
458 tunnel_idle_sleep: &mut Pin<&mut tokio::time::Sleep>,
459 tunnel_idle_timeout: Duration,
460 half_close_idle_sleep: &mut Pin<&mut tokio::time::Sleep>,
461 half_close_idle_timeout: Duration,
462 peer_done: bool,
463) {
464 state.len += activity.bytes;
465 reset_sleep_at(tunnel_idle_sleep, tunnel_idle_timeout, activity.at);
466 if peer_done {
467 reset_sleep_at(half_close_idle_sleep, half_close_idle_timeout, activity.at);
468 }
469}
470
471fn reset_sleep(sleep: &mut Pin<&mut tokio::time::Sleep>, timeout: Duration) {
472 if !timeout.is_zero() {
473 sleep.as_mut().reset(Instant::now() + timeout);
474 }
475}
476
477fn reset_sleep_at(
478 sleep: &mut Pin<&mut tokio::time::Sleep>,
479 timeout: Duration,
480 activity_at: Instant,
481) {
482 if !timeout.is_zero() {
483 sleep.as_mut().reset(activity_at + timeout);
484 }
485}
486
487fn handle_forward_final_result(result: Option<Result<usize>>, len: usize, detail: &'static str) {
488 if let Some(result) = result {
489 handle_forward_result(result, detail);
490 } else {
491 tracing::debug!("forward stopped before peer closed; we send {len} bytes,detail:{detail}");
492 }
493}
494
495pub async fn start_datagram_forward<
496 ClientReader: DatagramReader,
497 ClientWriter: DatagramWriter,
498 ServerReader: DatagramReader,
499 ServerWriter: DatagramWriter,
500>(
501 client_reader: ClientReader,
502 client_writer: ClientWriter,
503 server_reader: ServerReader,
504 server_writer: ServerWriter,
505) {
506 let client_to_server = transfer_datagrams("udp->tcp", client_reader, server_writer);
507 let server_to_client = transfer_datagrams("tcp->udp", server_reader, client_writer);
508 tokio::select! {
509 result = client_to_server =>{
510 handle_forward_result( result,"udp->tcp");
511 },
512 result = server_to_client =>{
513 handle_forward_result( result,"tcp->udp");
514 }
515 }
516}
517
518fn handle_forward_result(result: Result<usize>, detail: &'static str) {
519 match result {
520 Ok(len) => tracing::info!("forward finish! we send {len} bytes,detail:{detail}"),
521 Err(e) => {
522 if e.is_expected_disconnect() {
524 tracing::debug!("forward closed by peer:{e},detail:{detail}");
525 } else {
526 tracing::error!("got forward error:{e},detail:{detail}");
527 }
528 }
529 }
530}
531
532impl DatagramReader for UdpStreamReadHalf {
533 async fn recv(&mut self) -> Result<Bytes> {
534 self.recv_datagram()
535 .await
536 .map_err(|e| pb_mapper_core::error::Error::MsgForward {
537 action: "read",
538 source: e,
539 })
540 }
541}
542
543impl DatagramWriter for UdpStreamWriteHalf<'_> {
544 async fn send(&mut self, src: &[u8]) -> Result<()> {
545 self.send_datagram(src)
546 .await
547 .map_err(|e| pb_mapper_core::error::Error::MsgForward {
548 action: "write",
549 source: e,
550 })
551 }
552}
553
554pub trait StreamForward: StreamSplit + Sized {
555 fn forward_local_to_remote<'a, R, W>(
556 codec_key: Option<AesKeyType>,
557 framing_key: AesKeyType,
558 local_reader: Self::ReaderRef<'a>,
559 local_writer: Self::WriterRef<'a>,
560 remote_reader: R,
561 remote_writer: W,
562 ) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>>
563 where
564 R: AsyncReadExt + Unpin + Send + 'a,
565 W: AsyncWriteExt + Unpin + Send + 'a;
566}
567
568impl StreamForward for TcpStreamImpl {
569 fn forward_local_to_remote<'a, R, W>(
570 codec_key: Option<AesKeyType>,
571 framing_key: AesKeyType,
572 local_reader: Self::ReaderRef<'a>,
573 local_writer: Self::WriterRef<'a>,
574 remote_reader: R,
575 remote_writer: W,
576 ) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>>
577 where
578 R: AsyncReadExt + Unpin + Send + 'a,
579 W: AsyncWriteExt + Unpin + Send + 'a,
580 {
581 Box::pin(async move {
582 if let Err(error) = local_reader.as_ref().set_nodelay(true) {
586 tracing::warn!(%error, "failed to disable Nagle on local TCP stream");
587 }
588 let mut local_reader = local_reader;
589 let mut local_writer = local_writer;
590 let mut remote_reader = remote_reader;
591 let mut remote_writer = remote_writer;
592 match codec_key {
593 Some(key) => {
594 start_forward(
595 NormalForwardReader::new(&mut local_reader),
596 NormalForwardWriter::new(&mut local_writer),
597 CodecForwardReader::new(
598 &mut remote_reader,
599 snafu_error_get_or_return_ok!(
600 super::get_decodec(&key),
601 "failed to create decoder when remote forward"
602 ),
603 )
604 .with_checksum_key(framing_key),
605 CodecForwardWriter::new(
606 &mut remote_writer,
607 snafu_error_get_or_return_ok!(
608 super::get_encodec(&key),
609 "failed to create encoder when remote forward"
610 ),
611 )
612 .with_checksum_key(framing_key),
613 )
614 .await;
615 }
616 None => {
617 start_forward(
618 NormalForwardReader::new(&mut local_reader),
619 NormalForwardWriter::new(&mut local_writer),
620 NormalForwardReader::new(&mut remote_reader),
621 NormalForwardWriter::new(&mut remote_writer),
622 )
623 .await;
624 }
625 }
626 Ok(())
627 })
628 }
629}
630
631impl StreamForward for UdpStreamImpl {
632 fn forward_local_to_remote<'a, R, W>(
633 codec_key: Option<AesKeyType>,
634 framing_key: AesKeyType,
635 local_reader: Self::ReaderRef<'a>,
636 local_writer: Self::WriterRef<'a>,
637 remote_reader: R,
638 remote_writer: W,
639 ) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>>
640 where
641 R: AsyncReadExt + Unpin + Send + 'a,
642 W: AsyncWriteExt + Unpin + Send + 'a,
643 {
644 Box::pin(async move {
645 let mut remote_reader = remote_reader;
646 let mut remote_writer = remote_writer;
647 match codec_key {
648 Some(key) => {
649 start_datagram_forward(
650 local_reader,
651 local_writer,
652 CodecDatagramReader::new(
653 &mut remote_reader,
654 snafu_error_get_or_return_ok!(
655 super::get_decodec(&key),
656 "failed to create decoder when datagram forward"
657 ),
658 )
659 .with_checksum_key(framing_key),
660 CodecDatagramWriter::new(
661 &mut remote_writer,
662 snafu_error_get_or_return_ok!(
663 super::get_encodec(&key),
664 "failed to create encoder when datagram forward"
665 ),
666 )
667 .with_checksum_key(framing_key),
668 )
669 .await;
670 }
671 None => {
672 start_datagram_forward(
673 local_reader,
674 local_writer,
675 NormalDatagramReader::new(&mut remote_reader)
676 .with_checksum_key(framing_key),
677 NormalDatagramWriter::new(&mut remote_writer)
678 .with_checksum_key(framing_key),
679 )
680 .await;
681 }
682 }
683 Ok(())
684 })
685 }
686}
687
688#[cfg(test)]
689mod tests {
690 #[tokio::test]
691 async fn tcp_forward_disables_nagle_on_the_local_leg() {
692 use tokio::net::{TcpListener, TcpStream};
693 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
694 let (local, accepted) = tokio::join!(
695 TcpStream::connect(listener.local_addr().unwrap()),
696 listener.accept()
697 );
698 let socket = local.unwrap().into_std().unwrap();
699 let observer = socket.try_clone().unwrap();
700 assert!(!observer.nodelay().unwrap());
701 let mut local = TcpStreamImpl::new(TcpStream::from_std(socket).unwrap());
702 let (mut peer, _) = accepted.unwrap();
703 let (remote, mut echo) = tokio::io::duplex(64);
704 let forwarding = tokio::spawn(async move {
705 let (reader, writer) = local.split();
706 let (remote_reader, remote_writer) = tokio::io::split(remote);
707 TcpStreamImpl::forward_local_to_remote(
708 None,
709 [0; 32],
710 reader,
711 writer,
712 remote_reader,
713 remote_writer,
714 )
715 .await
716 .unwrap();
717 });
718 peer.write_all(b"x").await.unwrap();
719 assert_eq!(
720 tokio::time::timeout(Duration::from_secs(1), echo.read_u8())
721 .await
722 .unwrap()
723 .unwrap(),
724 b'x'
725 );
726 assert!(observer.nodelay().unwrap());
727 forwarding.abort();
728 }
729
730 use std::collections::VecDeque;
731 use std::io;
732 use std::sync::Arc;
733
734 use parking_lot::Mutex;
735 use std::time::Duration;
736
737 use super::*;
738 use pb_mapper_core::config::parse_duration;
739 use pb_mapper_core::error::Error;
740 use tokio::sync::Notify;
741
742 enum ReadAction {
743 Data(Vec<u8>),
744 Eof,
745 Pending,
746 Error(io::ErrorKind),
747 }
748
749 struct ScriptedReader {
750 actions: VecDeque<ReadAction>,
751 current: Vec<u8>,
752 }
753
754 impl ScriptedReader {
755 fn new(actions: impl IntoIterator<Item = ReadAction>) -> Self {
756 Self {
757 actions: actions.into_iter().collect(),
758 current: Vec::new(),
759 }
760 }
761 }
762
763 impl ForwardReader for ScriptedReader {
764 async fn read(&mut self) -> Result<&'_ [u8]> {
765 match self.actions.pop_front().unwrap_or(ReadAction::Pending) {
766 ReadAction::Data(data) => {
767 self.current = data;
768 Ok(&self.current)
769 }
770 ReadAction::Eof => {
771 self.current.clear();
772 Ok(&self.current)
773 }
774 ReadAction::Pending => std::future::pending().await,
775 ReadAction::Error(kind) => Err(Error::MsgForward {
776 action: "read",
777 source: io::Error::new(kind, "scripted read error"),
778 }),
779 }
780 }
781 }
782
783 #[derive(Default)]
784 struct WriterState {
785 chunks: Vec<Vec<u8>>,
786 shutdowns: usize,
787 }
788
789 #[derive(Clone, Default)]
790 struct ScriptedWriter {
791 state: Arc<Mutex<WriterState>>,
792 }
793
794 impl ScriptedWriter {
795 fn chunks(&self) -> Vec<Vec<u8>> {
796 self.state.lock().chunks.clone()
797 }
798
799 fn shutdowns(&self) -> usize {
800 self.state.lock().shutdowns
801 }
802 }
803
804 impl ForwardWriter for ScriptedWriter {
805 async fn write(&mut self, src: &[u8]) -> Result<()> {
806 self.state.lock().chunks.push(src.to_vec());
807 Ok(())
808 }
809
810 async fn shutdown(&mut self) {
811 self.state.lock().shutdowns += 1;
812 }
813 }
814
815 struct EofAfterWriteStartsReader {
816 write_started: Arc<Notify>,
817 returned_eof: bool,
818 empty: Vec<u8>,
819 }
820
821 impl EofAfterWriteStartsReader {
822 fn new(write_started: Arc<Notify>) -> Self {
823 Self {
824 write_started,
825 returned_eof: false,
826 empty: Vec::new(),
827 }
828 }
829 }
830
831 impl ForwardReader for EofAfterWriteStartsReader {
832 async fn read(&mut self) -> Result<&'_ [u8]> {
833 if self.returned_eof {
834 return std::future::pending().await;
835 }
836 self.write_started.notified().await;
837 self.returned_eof = true;
838 Ok(&self.empty)
839 }
840 }
841
842 #[derive(Clone)]
843 struct DelayedWriter {
844 state: Arc<Mutex<WriterState>>,
845 write_started: Arc<Notify>,
846 delay: Duration,
847 }
848
849 impl DelayedWriter {
850 fn new(write_started: Arc<Notify>, delay: Duration) -> Self {
851 Self {
852 state: Arc::new(Mutex::new(WriterState::default())),
853 write_started,
854 delay,
855 }
856 }
857
858 fn chunks(&self) -> Vec<Vec<u8>> {
859 self.state.lock().chunks.clone()
860 }
861 }
862
863 impl ForwardWriter for DelayedWriter {
864 async fn write(&mut self, src: &[u8]) -> Result<()> {
865 self.write_started.notify_one();
866 tokio::time::sleep(self.delay).await;
867 self.state.lock().chunks.push(src.to_vec());
868 Ok(())
869 }
870
871 async fn shutdown(&mut self) {
872 self.state.lock().shutdowns += 1;
873 }
874 }
875
876 #[test]
877 fn parse_duration_accepts_suffixes_and_plain_seconds() {
878 assert_eq!(parse_duration("42"), Some(Duration::from_secs(42)));
879 assert_eq!(parse_duration("500ms"), Some(Duration::from_millis(500)));
880 assert_eq!(parse_duration("2s"), Some(Duration::from_secs(2)));
881 assert_eq!(parse_duration("3m"), Some(Duration::from_secs(180)));
882 assert_eq!(parse_duration("1h"), Some(Duration::from_secs(3600)));
883 assert_eq!(parse_duration(""), None);
884 assert_eq!(parse_duration("bad"), None);
885 assert_eq!(parse_duration("18446744073709551615h"), None);
886 }
887
888 #[tokio::test]
889 async fn half_close_idle_timeout_closes_stalled_peer() {
890 let client_reader = ScriptedReader::new([ReadAction::Eof]);
891 let client_writer = ScriptedWriter::default();
892 let server_reader = ScriptedReader::new([ReadAction::Pending]);
893 let server_writer = ScriptedWriter::default();
894 let server_writer_state = server_writer.clone();
895
896 tokio::time::timeout(
897 Duration::from_millis(200),
898 start_forward_with_config(
899 client_reader,
900 client_writer,
901 server_reader,
902 server_writer,
903 ForwardTimeoutConfig {
904 tunnel_idle_timeout: Duration::from_secs(60 * 60),
905 half_close_idle_timeout: Duration::from_millis(20),
906 },
907 ),
908 )
909 .await
910 .expect("half-closed tunnel did not stop after half-close idle timeout");
911
912 assert_eq!(server_writer_state.shutdowns(), 1);
913 }
914
915 #[tokio::test]
916 async fn expected_disconnect_stops_waiting_for_pending_peer() {
917 let client_reader =
918 ScriptedReader::new([ReadAction::Error(io::ErrorKind::ConnectionReset)]);
919 let client_writer = ScriptedWriter::default();
920 let server_reader = ScriptedReader::new([ReadAction::Pending]);
921 let server_writer = ScriptedWriter::default();
922
923 tokio::time::timeout(
924 Duration::from_millis(200),
925 start_forward_with_config(
926 client_reader,
927 client_writer,
928 server_reader,
929 server_writer,
930 ForwardTimeoutConfig {
931 tunnel_idle_timeout: Duration::from_secs(60 * 60),
932 half_close_idle_timeout: Duration::from_secs(60),
933 },
934 ),
935 )
936 .await
937 .expect("expected disconnect did not stop the tunnel");
938 }
939
940 #[tokio::test]
941 async fn half_closed_tunnel_drains_peer_before_timeout() {
942 let client_reader = ScriptedReader::new([ReadAction::Eof]);
943 let client_writer = ScriptedWriter::default();
944 let client_writer_state = client_writer.clone();
945 let server_reader =
946 ScriptedReader::new([ReadAction::Data(b"response".to_vec()), ReadAction::Eof]);
947 let server_writer = ScriptedWriter::default();
948
949 tokio::time::timeout(
950 Duration::from_millis(200),
951 start_forward_with_config(
952 client_reader,
953 client_writer,
954 server_reader,
955 server_writer,
956 ForwardTimeoutConfig {
957 tunnel_idle_timeout: Duration::from_secs(60 * 60),
958 half_close_idle_timeout: Duration::from_millis(200),
959 },
960 ),
961 )
962 .await
963 .expect("half-closed tunnel failed to drain the peer");
964
965 assert_eq!(client_writer_state.chunks(), vec![b"response".to_vec()]);
966 assert_eq!(client_writer_state.shutdowns(), 1);
967 }
968
969 #[tokio::test]
970 async fn delayed_tail_write_survives_peer_half_close() {
971 let write_started = Arc::new(Notify::new());
972 let client_reader = EofAfterWriteStartsReader::new(write_started.clone());
973 let client_writer = DelayedWriter::new(write_started, Duration::from_millis(20));
974 let client_writer_state = client_writer.clone();
975 let tail = vec![0x5a; 499];
976 let server_reader = ScriptedReader::new([ReadAction::Data(tail.clone()), ReadAction::Eof]);
977 let server_writer = ScriptedWriter::default();
978
979 tokio::time::timeout(
980 Duration::from_millis(300),
981 start_forward_with_config(
982 client_reader,
983 client_writer,
984 server_reader,
985 server_writer,
986 ForwardTimeoutConfig {
987 tunnel_idle_timeout: Duration::from_secs(60 * 60),
988 half_close_idle_timeout: Duration::from_millis(200),
989 },
990 ),
991 )
992 .await
993 .expect("delayed response tail was lost after the peer half-closed");
994
995 assert_eq!(client_writer_state.chunks(), vec![tail]);
996 }
997
998 #[tokio::test]
999 async fn open_tunnel_idle_timeout_closes_inactive_tunnel() {
1000 let client_reader = ScriptedReader::new([ReadAction::Pending]);
1001 let client_writer = ScriptedWriter::default();
1002 let server_reader = ScriptedReader::new([ReadAction::Pending]);
1003 let server_writer = ScriptedWriter::default();
1004
1005 tokio::time::timeout(
1006 Duration::from_millis(200),
1007 start_forward_with_config(
1008 client_reader,
1009 client_writer,
1010 server_reader,
1011 server_writer,
1012 ForwardTimeoutConfig {
1013 tunnel_idle_timeout: Duration::from_millis(20),
1014 half_close_idle_timeout: Duration::from_secs(60),
1015 },
1016 ),
1017 )
1018 .await
1019 .expect("inactive open tunnel did not stop after tunnel idle timeout");
1020 }
1021}