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 let mut local_reader = local_reader;
583 let mut local_writer = local_writer;
584 let mut remote_reader = remote_reader;
585 let mut remote_writer = remote_writer;
586 match codec_key {
587 Some(key) => {
588 start_forward(
589 NormalForwardReader::new(&mut local_reader),
590 NormalForwardWriter::new(&mut local_writer),
591 CodecForwardReader::new(
592 &mut remote_reader,
593 snafu_error_get_or_return_ok!(
594 super::get_decodec(&key),
595 "failed to create decoder when remote forward"
596 ),
597 )
598 .with_checksum_key(framing_key),
599 CodecForwardWriter::new(
600 &mut remote_writer,
601 snafu_error_get_or_return_ok!(
602 super::get_encodec(&key),
603 "failed to create encoder when remote forward"
604 ),
605 )
606 .with_checksum_key(framing_key),
607 )
608 .await;
609 }
610 None => {
611 start_forward(
612 NormalForwardReader::new(&mut local_reader),
613 NormalForwardWriter::new(&mut local_writer),
614 NormalForwardReader::new(&mut remote_reader),
615 NormalForwardWriter::new(&mut remote_writer),
616 )
617 .await;
618 }
619 }
620 Ok(())
621 })
622 }
623}
624
625impl StreamForward for UdpStreamImpl {
626 fn forward_local_to_remote<'a, R, W>(
627 codec_key: Option<AesKeyType>,
628 framing_key: AesKeyType,
629 local_reader: Self::ReaderRef<'a>,
630 local_writer: Self::WriterRef<'a>,
631 remote_reader: R,
632 remote_writer: W,
633 ) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>>
634 where
635 R: AsyncReadExt + Unpin + Send + 'a,
636 W: AsyncWriteExt + Unpin + Send + 'a,
637 {
638 Box::pin(async move {
639 let mut remote_reader = remote_reader;
640 let mut remote_writer = remote_writer;
641 match codec_key {
642 Some(key) => {
643 start_datagram_forward(
644 local_reader,
645 local_writer,
646 CodecDatagramReader::new(
647 &mut remote_reader,
648 snafu_error_get_or_return_ok!(
649 super::get_decodec(&key),
650 "failed to create decoder when datagram forward"
651 ),
652 )
653 .with_checksum_key(framing_key),
654 CodecDatagramWriter::new(
655 &mut remote_writer,
656 snafu_error_get_or_return_ok!(
657 super::get_encodec(&key),
658 "failed to create encoder when datagram forward"
659 ),
660 )
661 .with_checksum_key(framing_key),
662 )
663 .await;
664 }
665 None => {
666 start_datagram_forward(
667 local_reader,
668 local_writer,
669 NormalDatagramReader::new(&mut remote_reader)
670 .with_checksum_key(framing_key),
671 NormalDatagramWriter::new(&mut remote_writer)
672 .with_checksum_key(framing_key),
673 )
674 .await;
675 }
676 }
677 Ok(())
678 })
679 }
680}
681
682#[cfg(test)]
683mod tests {
684 use std::collections::VecDeque;
685 use std::io;
686 use std::sync::Arc;
687
688 use parking_lot::Mutex;
689 use std::time::Duration;
690
691 use super::*;
692 use pb_mapper_core::config::parse_duration;
693 use pb_mapper_core::error::Error;
694 use tokio::sync::Notify;
695
696 enum ReadAction {
697 Data(Vec<u8>),
698 Eof,
699 Pending,
700 Error(io::ErrorKind),
701 }
702
703 struct ScriptedReader {
704 actions: VecDeque<ReadAction>,
705 current: Vec<u8>,
706 }
707
708 impl ScriptedReader {
709 fn new(actions: impl IntoIterator<Item = ReadAction>) -> Self {
710 Self {
711 actions: actions.into_iter().collect(),
712 current: Vec::new(),
713 }
714 }
715 }
716
717 impl ForwardReader for ScriptedReader {
718 async fn read(&mut self) -> Result<&'_ [u8]> {
719 match self.actions.pop_front().unwrap_or(ReadAction::Pending) {
720 ReadAction::Data(data) => {
721 self.current = data;
722 Ok(&self.current)
723 }
724 ReadAction::Eof => {
725 self.current.clear();
726 Ok(&self.current)
727 }
728 ReadAction::Pending => std::future::pending().await,
729 ReadAction::Error(kind) => Err(Error::MsgForward {
730 action: "read",
731 source: io::Error::new(kind, "scripted read error"),
732 }),
733 }
734 }
735 }
736
737 #[derive(Default)]
738 struct WriterState {
739 chunks: Vec<Vec<u8>>,
740 shutdowns: usize,
741 }
742
743 #[derive(Clone, Default)]
744 struct ScriptedWriter {
745 state: Arc<Mutex<WriterState>>,
746 }
747
748 impl ScriptedWriter {
749 fn chunks(&self) -> Vec<Vec<u8>> {
750 self.state.lock().chunks.clone()
751 }
752
753 fn shutdowns(&self) -> usize {
754 self.state.lock().shutdowns
755 }
756 }
757
758 impl ForwardWriter for ScriptedWriter {
759 async fn write(&mut self, src: &[u8]) -> Result<()> {
760 self.state.lock().chunks.push(src.to_vec());
761 Ok(())
762 }
763
764 async fn shutdown(&mut self) {
765 self.state.lock().shutdowns += 1;
766 }
767 }
768
769 struct EofAfterWriteStartsReader {
770 write_started: Arc<Notify>,
771 returned_eof: bool,
772 empty: Vec<u8>,
773 }
774
775 impl EofAfterWriteStartsReader {
776 fn new(write_started: Arc<Notify>) -> Self {
777 Self {
778 write_started,
779 returned_eof: false,
780 empty: Vec::new(),
781 }
782 }
783 }
784
785 impl ForwardReader for EofAfterWriteStartsReader {
786 async fn read(&mut self) -> Result<&'_ [u8]> {
787 if self.returned_eof {
788 return std::future::pending().await;
789 }
790 self.write_started.notified().await;
791 self.returned_eof = true;
792 Ok(&self.empty)
793 }
794 }
795
796 #[derive(Clone)]
797 struct DelayedWriter {
798 state: Arc<Mutex<WriterState>>,
799 write_started: Arc<Notify>,
800 delay: Duration,
801 }
802
803 impl DelayedWriter {
804 fn new(write_started: Arc<Notify>, delay: Duration) -> Self {
805 Self {
806 state: Arc::new(Mutex::new(WriterState::default())),
807 write_started,
808 delay,
809 }
810 }
811
812 fn chunks(&self) -> Vec<Vec<u8>> {
813 self.state.lock().chunks.clone()
814 }
815 }
816
817 impl ForwardWriter for DelayedWriter {
818 async fn write(&mut self, src: &[u8]) -> Result<()> {
819 self.write_started.notify_one();
820 tokio::time::sleep(self.delay).await;
821 self.state.lock().chunks.push(src.to_vec());
822 Ok(())
823 }
824
825 async fn shutdown(&mut self) {
826 self.state.lock().shutdowns += 1;
827 }
828 }
829
830 #[test]
831 fn parse_duration_accepts_suffixes_and_plain_seconds() {
832 assert_eq!(parse_duration("42"), Some(Duration::from_secs(42)));
833 assert_eq!(parse_duration("500ms"), Some(Duration::from_millis(500)));
834 assert_eq!(parse_duration("2s"), Some(Duration::from_secs(2)));
835 assert_eq!(parse_duration("3m"), Some(Duration::from_secs(180)));
836 assert_eq!(parse_duration("1h"), Some(Duration::from_secs(3600)));
837 assert_eq!(parse_duration(""), None);
838 assert_eq!(parse_duration("bad"), None);
839 assert_eq!(parse_duration("18446744073709551615h"), None);
840 }
841
842 #[tokio::test]
843 async fn half_close_idle_timeout_closes_stalled_peer() {
844 let client_reader = ScriptedReader::new([ReadAction::Eof]);
845 let client_writer = ScriptedWriter::default();
846 let server_reader = ScriptedReader::new([ReadAction::Pending]);
847 let server_writer = ScriptedWriter::default();
848 let server_writer_state = server_writer.clone();
849
850 tokio::time::timeout(
851 Duration::from_millis(200),
852 start_forward_with_config(
853 client_reader,
854 client_writer,
855 server_reader,
856 server_writer,
857 ForwardTimeoutConfig {
858 tunnel_idle_timeout: Duration::from_secs(60 * 60),
859 half_close_idle_timeout: Duration::from_millis(20),
860 },
861 ),
862 )
863 .await
864 .expect("half-closed tunnel did not stop after half-close idle timeout");
865
866 assert_eq!(server_writer_state.shutdowns(), 1);
867 }
868
869 #[tokio::test]
870 async fn expected_disconnect_stops_waiting_for_pending_peer() {
871 let client_reader =
872 ScriptedReader::new([ReadAction::Error(io::ErrorKind::ConnectionReset)]);
873 let client_writer = ScriptedWriter::default();
874 let server_reader = ScriptedReader::new([ReadAction::Pending]);
875 let server_writer = ScriptedWriter::default();
876
877 tokio::time::timeout(
878 Duration::from_millis(200),
879 start_forward_with_config(
880 client_reader,
881 client_writer,
882 server_reader,
883 server_writer,
884 ForwardTimeoutConfig {
885 tunnel_idle_timeout: Duration::from_secs(60 * 60),
886 half_close_idle_timeout: Duration::from_secs(60),
887 },
888 ),
889 )
890 .await
891 .expect("expected disconnect did not stop the tunnel");
892 }
893
894 #[tokio::test]
895 async fn half_closed_tunnel_drains_peer_before_timeout() {
896 let client_reader = ScriptedReader::new([ReadAction::Eof]);
897 let client_writer = ScriptedWriter::default();
898 let client_writer_state = client_writer.clone();
899 let server_reader =
900 ScriptedReader::new([ReadAction::Data(b"response".to_vec()), ReadAction::Eof]);
901 let server_writer = ScriptedWriter::default();
902
903 tokio::time::timeout(
904 Duration::from_millis(200),
905 start_forward_with_config(
906 client_reader,
907 client_writer,
908 server_reader,
909 server_writer,
910 ForwardTimeoutConfig {
911 tunnel_idle_timeout: Duration::from_secs(60 * 60),
912 half_close_idle_timeout: Duration::from_millis(200),
913 },
914 ),
915 )
916 .await
917 .expect("half-closed tunnel failed to drain the peer");
918
919 assert_eq!(client_writer_state.chunks(), vec![b"response".to_vec()]);
920 assert_eq!(client_writer_state.shutdowns(), 1);
921 }
922
923 #[tokio::test]
924 async fn delayed_tail_write_survives_peer_half_close() {
925 let write_started = Arc::new(Notify::new());
926 let client_reader = EofAfterWriteStartsReader::new(write_started.clone());
927 let client_writer = DelayedWriter::new(write_started, Duration::from_millis(20));
928 let client_writer_state = client_writer.clone();
929 let tail = vec![0x5a; 499];
930 let server_reader = ScriptedReader::new([ReadAction::Data(tail.clone()), ReadAction::Eof]);
931 let server_writer = ScriptedWriter::default();
932
933 tokio::time::timeout(
934 Duration::from_millis(300),
935 start_forward_with_config(
936 client_reader,
937 client_writer,
938 server_reader,
939 server_writer,
940 ForwardTimeoutConfig {
941 tunnel_idle_timeout: Duration::from_secs(60 * 60),
942 half_close_idle_timeout: Duration::from_millis(200),
943 },
944 ),
945 )
946 .await
947 .expect("delayed response tail was lost after the peer half-closed");
948
949 assert_eq!(client_writer_state.chunks(), vec![tail]);
950 }
951
952 #[tokio::test]
953 async fn open_tunnel_idle_timeout_closes_inactive_tunnel() {
954 let client_reader = ScriptedReader::new([ReadAction::Pending]);
955 let client_writer = ScriptedWriter::default();
956 let server_reader = ScriptedReader::new([ReadAction::Pending]);
957 let server_writer = ScriptedWriter::default();
958
959 tokio::time::timeout(
960 Duration::from_millis(200),
961 start_forward_with_config(
962 client_reader,
963 client_writer,
964 server_reader,
965 server_writer,
966 ForwardTimeoutConfig {
967 tunnel_idle_timeout: Duration::from_millis(20),
968 half_close_idle_timeout: Duration::from_secs(60),
969 },
970 ),
971 )
972 .await
973 .expect("inactive open tunnel did not stop after tunnel idle timeout");
974 }
975}