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            // Both the accepted subscriber socket and the publisher's target
583            // socket pass here. Small RPC frames must not wait for Nagle's
584            // coalescing after the relay legs have already disabled it.
585            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}