Skip to main content

pb_mapper_protocol/
forward.rs

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    // Gracefully close the write side to support half-closed TCP streams.
30    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    /// SAFETY: Same as [`CodecMessageWriter`]
224    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    // Keep each direction as one pinned future. Recreating a combined read/write future in every
373    // `select!` iteration can cancel it after the read consumed bytes but before the write
374    // completed, silently dropping a response tail when the peer direction finishes first.
375    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            // Treat peer-initiated shutdowns as expected to avoid noisy error logs.
523            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}