Skip to main content

pingora_core/protocols/l4/
stream.rs

1// Copyright 2026 Cloudflare, Inc.
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7// http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! Transport layer connection
16
17use async_trait::async_trait;
18use futures::FutureExt;
19use log::{debug, error};
20
21use pingora_error::{ErrorType::*, OrErr, Result};
22#[cfg(target_os = "linux")]
23use std::io::IoSliceMut;
24#[cfg(unix)]
25use std::os::unix::io::AsRawFd;
26#[cfg(windows)]
27use std::os::windows::io::AsRawSocket;
28use std::pin::Pin;
29use std::sync::Arc;
30use std::task::{Context, Poll};
31use std::time::{Duration, Instant, SystemTime};
32#[cfg(target_os = "linux")]
33use tokio::io::Interest;
34use tokio::io::{self, AsyncRead, AsyncWrite, AsyncWriteExt, BufStream, ReadBuf};
35use tokio::net::TcpStream;
36#[cfg(unix)]
37use tokio::net::UnixStream;
38
39use crate::protocols::l4::ext::{set_tcp_keepalive, TcpKeepalive};
40use crate::protocols::l4::virt;
41use crate::protocols::raw_connect::ProxyDigest;
42use crate::protocols::{
43    GetProxyDigest, GetSocketDigest, GetTimingDigest, Peek, Shutdown, SocketDigest, Ssl,
44    TimingDigest, UniqueID, UniqueIDType,
45};
46use crate::upstreams::peer::Tracer;
47
48#[derive(Debug)]
49enum RawStream {
50    Tcp(TcpStream),
51    #[cfg(unix)]
52    Unix(UnixStream),
53    Virtual(virt::VirtualSocketStream),
54}
55
56impl AsyncRead for RawStream {
57    fn poll_read(
58        self: Pin<&mut Self>,
59        cx: &mut Context<'_>,
60        buf: &mut ReadBuf<'_>,
61    ) -> Poll<io::Result<()>> {
62        // Safety: Basic enum pin projection
63        unsafe {
64            match &mut Pin::get_unchecked_mut(self) {
65                RawStream::Tcp(s) => Pin::new_unchecked(s).poll_read(cx, buf),
66                #[cfg(unix)]
67                RawStream::Unix(s) => Pin::new_unchecked(s).poll_read(cx, buf),
68                RawStream::Virtual(s) => Pin::new_unchecked(s).poll_read(cx, buf),
69            }
70        }
71    }
72}
73
74impl AsyncWrite for RawStream {
75    fn poll_write(self: Pin<&mut Self>, cx: &mut Context, buf: &[u8]) -> Poll<io::Result<usize>> {
76        // Safety: Basic enum pin projection
77        unsafe {
78            match &mut Pin::get_unchecked_mut(self) {
79                RawStream::Tcp(s) => Pin::new_unchecked(s).poll_write(cx, buf),
80                #[cfg(unix)]
81                RawStream::Unix(s) => Pin::new_unchecked(s).poll_write(cx, buf),
82                RawStream::Virtual(s) => Pin::new_unchecked(s).poll_write(cx, buf),
83            }
84        }
85    }
86
87    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
88        // Safety: Basic enum pin projection
89        unsafe {
90            match &mut Pin::get_unchecked_mut(self) {
91                RawStream::Tcp(s) => Pin::new_unchecked(s).poll_flush(cx),
92                #[cfg(unix)]
93                RawStream::Unix(s) => Pin::new_unchecked(s).poll_flush(cx),
94                RawStream::Virtual(s) => Pin::new_unchecked(s).poll_flush(cx),
95            }
96        }
97    }
98
99    fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
100        // Safety: Basic enum pin projection
101        unsafe {
102            match &mut Pin::get_unchecked_mut(self) {
103                RawStream::Tcp(s) => Pin::new_unchecked(s).poll_shutdown(cx),
104                #[cfg(unix)]
105                RawStream::Unix(s) => Pin::new_unchecked(s).poll_shutdown(cx),
106                RawStream::Virtual(s) => Pin::new_unchecked(s).poll_shutdown(cx),
107            }
108        }
109    }
110
111    fn poll_write_vectored(
112        self: Pin<&mut Self>,
113        cx: &mut Context<'_>,
114        bufs: &[std::io::IoSlice<'_>],
115    ) -> Poll<io::Result<usize>> {
116        // Safety: Basic enum pin projection
117        unsafe {
118            match &mut Pin::get_unchecked_mut(self) {
119                RawStream::Tcp(s) => Pin::new_unchecked(s).poll_write_vectored(cx, bufs),
120                #[cfg(unix)]
121                RawStream::Unix(s) => Pin::new_unchecked(s).poll_write_vectored(cx, bufs),
122                RawStream::Virtual(s) => Pin::new_unchecked(s).poll_write_vectored(cx, bufs),
123            }
124        }
125    }
126
127    fn is_write_vectored(&self) -> bool {
128        match self {
129            RawStream::Tcp(s) => s.is_write_vectored(),
130            #[cfg(unix)]
131            RawStream::Unix(s) => s.is_write_vectored(),
132            RawStream::Virtual(s) => s.is_write_vectored(),
133        }
134    }
135}
136
137#[cfg(unix)]
138impl AsRawFd for RawStream {
139    fn as_raw_fd(&self) -> std::os::unix::io::RawFd {
140        match self {
141            RawStream::Tcp(s) => s.as_raw_fd(),
142            RawStream::Unix(s) => s.as_raw_fd(),
143            RawStream::Virtual(_) => -1, // Virtual stream does not have a real fd
144        }
145    }
146}
147
148#[cfg(windows)]
149impl AsRawSocket for RawStream {
150    fn as_raw_socket(&self) -> std::os::windows::io::RawSocket {
151        match self {
152            RawStream::Tcp(s) => s.as_raw_socket(),
153            // Virtual stream does not have a real socket, return INVALID_SOCKET (!0)
154            RawStream::Virtual(_) => !0,
155        }
156    }
157}
158
159#[derive(Debug)]
160struct RawStreamWrapper {
161    pub(crate) stream: RawStream,
162    /// store the last rx timestamp of the stream.
163    pub(crate) rx_ts: Option<SystemTime>,
164    /// enable reading rx timestamp
165    #[cfg(target_os = "linux")]
166    pub(crate) enable_rx_ts: bool,
167    #[cfg(target_os = "linux")]
168    /// This can be reused across multiple recvmsg calls. The cmsg buffer may
169    /// come from old sockets created by older version of pingora and so,
170    /// this vector can only grow.
171    reusable_cmsg_space: Vec<u8>,
172}
173
174impl RawStreamWrapper {
175    pub fn new(stream: RawStream) -> Self {
176        RawStreamWrapper {
177            stream,
178            rx_ts: None,
179            #[cfg(target_os = "linux")]
180            enable_rx_ts: false,
181            #[cfg(target_os = "linux")]
182            reusable_cmsg_space: nix::cmsg_space!(nix::sys::socket::Timestamps),
183        }
184    }
185
186    #[cfg(target_os = "linux")]
187    pub fn enable_rx_ts(&mut self, enable_rx_ts: bool) {
188        self.enable_rx_ts = enable_rx_ts;
189    }
190}
191
192impl AsyncRead for RawStreamWrapper {
193    #[cfg(not(target_os = "linux"))]
194    fn poll_read(
195        self: Pin<&mut Self>,
196        cx: &mut Context<'_>,
197        buf: &mut ReadBuf<'_>,
198    ) -> Poll<io::Result<()>> {
199        // Safety: Basic enum pin projection
200        unsafe {
201            let rs_wrapper = Pin::get_unchecked_mut(self);
202            match &mut rs_wrapper.stream {
203                RawStream::Tcp(s) => Pin::new_unchecked(s).poll_read(cx, buf),
204                #[cfg(unix)]
205                RawStream::Unix(s) => Pin::new_unchecked(s).poll_read(cx, buf),
206                RawStream::Virtual(s) => Pin::new_unchecked(s).poll_read(cx, buf),
207            }
208        }
209    }
210
211    #[cfg(target_os = "linux")]
212    fn poll_read(
213        self: Pin<&mut Self>,
214        cx: &mut Context<'_>,
215        buf: &mut ReadBuf<'_>,
216    ) -> Poll<io::Result<()>> {
217        use futures::ready;
218        use nix::sys::socket::{recvmsg, ControlMessageOwned, MsgFlags, SockaddrStorage};
219
220        // if we do not need rx timestamp, then use the standard path
221        if !self.enable_rx_ts {
222            // Safety: Basic enum pin projection
223            unsafe {
224                let rs_wrapper = Pin::get_unchecked_mut(self);
225                match &mut rs_wrapper.stream {
226                    RawStream::Tcp(s) => return Pin::new_unchecked(s).poll_read(cx, buf),
227                    RawStream::Unix(s) => return Pin::new_unchecked(s).poll_read(cx, buf),
228                    RawStream::Virtual(s) => return Pin::new_unchecked(s).poll_read(cx, buf),
229                }
230            }
231        }
232
233        // Safety: Basic pin projection to get mutable stream
234        let rs_wrapper = unsafe { Pin::get_unchecked_mut(self) };
235        match &mut rs_wrapper.stream {
236            RawStream::Tcp(s) => {
237                loop {
238                    ready!(s.poll_read_ready(cx))?;
239                    // Safety: maybe uninitialized bytes will only be passed to recvmsg
240                    let b = unsafe {
241                        &mut *(buf.unfilled_mut() as *mut [std::mem::MaybeUninit<u8>]
242                            as *mut [u8])
243                    };
244                    let mut iov = [IoSliceMut::new(b)];
245
246                    rs_wrapper.reusable_cmsg_space.fill(0);
247
248                    match s.try_io(Interest::READABLE, || {
249                        recvmsg::<SockaddrStorage>(
250                            s.as_raw_fd(),
251                            &mut iov,
252                            Some(&mut rs_wrapper.reusable_cmsg_space),
253                            MsgFlags::empty(),
254                        )
255                        .map_err(|errno| errno.into())
256                    }) {
257                        Ok(r) => {
258                            if let Some(ControlMessageOwned::ScmTimestampsns(rtime)) = r
259                                .cmsgs()?
260                                .find(|i| matches!(i, ControlMessageOwned::ScmTimestampsns(_)))
261                            {
262                                // The returned timestamp is a real (i.e. not monotonic) timestamp
263                                // https://docs.kernel.org/networking/timestamping.html
264                                rs_wrapper.rx_ts =
265                                    SystemTime::UNIX_EPOCH.checked_add(rtime.system.into());
266                            }
267                            // Safety: We trust `recvmsg` to have filled up `r.bytes` bytes in the buffer.
268                            unsafe {
269                                buf.assume_init(r.bytes);
270                            }
271                            buf.advance(r.bytes);
272                            return Poll::Ready(Ok(()));
273                        }
274                        Err(ref e) if e.kind() == io::ErrorKind::WouldBlock => continue,
275                        Err(e) => return Poll::Ready(Err(e)),
276                    }
277                }
278            }
279            // Unix RX timestamp only works with datagram for now, so we do not care about it
280            RawStream::Unix(s) => unsafe { Pin::new_unchecked(s).poll_read(cx, buf) },
281            RawStream::Virtual(s) => unsafe { Pin::new_unchecked(s).poll_read(cx, buf) },
282        }
283    }
284}
285
286impl AsyncWrite for RawStreamWrapper {
287    fn poll_write(self: Pin<&mut Self>, cx: &mut Context, buf: &[u8]) -> Poll<io::Result<usize>> {
288        // Safety: Basic enum pin projection
289        unsafe {
290            match &mut Pin::get_unchecked_mut(self).stream {
291                RawStream::Tcp(s) => Pin::new_unchecked(s).poll_write(cx, buf),
292                #[cfg(unix)]
293                RawStream::Unix(s) => Pin::new_unchecked(s).poll_write(cx, buf),
294                RawStream::Virtual(s) => Pin::new_unchecked(s).poll_write(cx, buf),
295            }
296        }
297    }
298
299    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
300        // Safety: Basic enum pin projection
301        unsafe {
302            match &mut Pin::get_unchecked_mut(self).stream {
303                RawStream::Tcp(s) => Pin::new_unchecked(s).poll_flush(cx),
304                #[cfg(unix)]
305                RawStream::Unix(s) => Pin::new_unchecked(s).poll_flush(cx),
306                RawStream::Virtual(s) => Pin::new_unchecked(s).poll_flush(cx),
307            }
308        }
309    }
310
311    fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
312        // Safety: Basic enum pin projection
313        unsafe {
314            match &mut Pin::get_unchecked_mut(self).stream {
315                RawStream::Tcp(s) => Pin::new_unchecked(s).poll_shutdown(cx),
316                #[cfg(unix)]
317                RawStream::Unix(s) => Pin::new_unchecked(s).poll_shutdown(cx),
318                RawStream::Virtual(s) => Pin::new_unchecked(s).poll_shutdown(cx),
319            }
320        }
321    }
322
323    fn poll_write_vectored(
324        self: Pin<&mut Self>,
325        cx: &mut Context<'_>,
326        bufs: &[std::io::IoSlice<'_>],
327    ) -> Poll<io::Result<usize>> {
328        // Safety: Basic enum pin projection
329        unsafe {
330            match &mut Pin::get_unchecked_mut(self).stream {
331                RawStream::Tcp(s) => Pin::new_unchecked(s).poll_write_vectored(cx, bufs),
332                #[cfg(unix)]
333                RawStream::Unix(s) => Pin::new_unchecked(s).poll_write_vectored(cx, bufs),
334                RawStream::Virtual(s) => Pin::new_unchecked(s).poll_write_vectored(cx, bufs),
335            }
336        }
337    }
338
339    fn is_write_vectored(&self) -> bool {
340        self.stream.is_write_vectored()
341    }
342}
343
344#[cfg(unix)]
345impl AsRawFd for RawStreamWrapper {
346    fn as_raw_fd(&self) -> std::os::unix::io::RawFd {
347        self.stream.as_raw_fd()
348    }
349}
350
351#[cfg(windows)]
352impl AsRawSocket for RawStreamWrapper {
353    fn as_raw_socket(&self) -> std::os::windows::io::RawSocket {
354        self.stream.as_raw_socket()
355    }
356}
357
358/// The default L4 read buffer size.
359///
360/// Large read buffering helps reducing syscalls with little trade-off. The SSL
361/// layer always does "small" reads in 16k chunks (TLS record size), so L4 read
362/// buffering helps a lot.
363pub const DEFAULT_L4_READ_BUFFER_SIZE: usize = 64 * 1024;
364
365/// The default L4 write buffer size.
366///
367/// Small write buffering matches a typical MSS. Too large a write buffer delays
368/// real-time communication. This buffering effectively implements something
369/// similar to Nagle's algorithm, but user space can control when to flush.
370pub const DEFAULT_L4_WRITE_BUFFER_SIZE: usize = 1460;
371
372/// L4 [`BufStream`] buffer sizing.
373///
374/// Leaving either side as `None` preserves Pingora's default for that side.
375/// Setting either side to `Some(0)` disables `BufStream` buffering for that
376/// direction.
377#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
378pub struct L4BufferSettings {
379    /// Read buffer size in bytes. `None` uses [`DEFAULT_L4_READ_BUFFER_SIZE`].
380    pub read: Option<usize>,
381    /// Write buffer size in bytes. `None` uses [`DEFAULT_L4_WRITE_BUFFER_SIZE`].
382    pub write: Option<usize>,
383}
384
385impl L4BufferSettings {
386    /// Create settings with both read and write buffer sizes set explicitly.
387    pub fn new(read: usize, write: usize) -> Self {
388        Self {
389            read: Some(read),
390            write: Some(write),
391        }
392    }
393
394    /// Create settings that disable both read and write `BufStream` buffering.
395    pub fn unbuffered() -> Self {
396        Self::new(0, 0)
397    }
398
399    /// Set the read buffer size.
400    pub fn read(mut self, read: usize) -> Self {
401        self.read = Some(read);
402        self
403    }
404
405    /// Set the write buffer size.
406    pub fn write(mut self, write: usize) -> Self {
407        self.write = Some(write);
408        self
409    }
410
411    /// Resolved read buffer size after applying defaults.
412    pub fn read_capacity(&self) -> usize {
413        self.read.unwrap_or(DEFAULT_L4_READ_BUFFER_SIZE)
414    }
415
416    /// Resolved write buffer size after applying defaults.
417    pub fn write_capacity(&self) -> usize {
418        self.write.unwrap_or(DEFAULT_L4_WRITE_BUFFER_SIZE)
419    }
420}
421
422// NOTE: with writer buffering, users need to call flush() to make sure the data is actually
423// sent. Otherwise data could be stuck in the buffer forever or get lost when stream is closed.
424
425/// A concrete type for transport layer connection + extra fields for logging
426#[derive(Debug)]
427pub struct Stream {
428    // Use `Option` to be able to swap to adjust the buffer size. Always safe to unwrap
429    stream: Option<BufStream<RawStreamWrapper>>,
430    // the data put back at the front of the read buffer, in order to replay the read
431    rewind_read_buf: Vec<Vec<u8>>,
432    buffer_write: bool,
433    proxy_digest: Option<Arc<ProxyDigest>>,
434    socket_digest: Option<Arc<SocketDigest>>,
435    /// When this connection is established
436    pub established_ts: SystemTime,
437    establishment_duration: Option<Duration>,
438    offload_wait_duration: Option<Duration>,
439    /// The distributed tracing object for this stream
440    pub tracer: Option<Tracer>,
441    read_pending_time: AccumulatedDuration,
442    write_pending_time: AccumulatedDuration,
443    /// Last rx timestamp associated with the last recvmsg call.
444    pub rx_ts: Option<SystemTime>,
445}
446
447impl Stream {
448    fn stream(&self) -> &BufStream<RawStreamWrapper> {
449        self.stream.as_ref().expect("stream should always be set")
450    }
451
452    fn stream_mut(&mut self) -> &mut BufStream<RawStreamWrapper> {
453        self.stream.as_mut().expect("stream should always be set")
454    }
455
456    /// Record transport connection timing for this stream.
457    pub(crate) fn set_establishment_timing(
458        &mut self,
459        establishment_duration: Duration,
460        offload_wait_duration: Option<Duration>,
461    ) {
462        self.establishment_duration = Some(establishment_duration);
463        self.offload_wait_duration = offload_wait_duration;
464    }
465
466    /// set TCP nodelay for this connection if `self` is TCP
467    pub fn set_nodelay(&mut self) -> Result<()> {
468        match &self.stream_mut().get_mut().stream {
469            RawStream::Tcp(s) => {
470                s.set_nodelay(true)
471                    .or_err(ConnectError, "failed to set_nodelay")?;
472            }
473            RawStream::Virtual(s) => {
474                s.set_socket_option(virt::VirtualSockOpt::NoDelay)
475                    .or_err(ConnectError, "failed to set_nodelay on virtual socket")?;
476            }
477            _ => (),
478        }
479        Ok(())
480    }
481
482    /// set TCP keepalive settings for this connection if `self` is TCP
483    pub fn set_keepalive(&mut self, ka: &TcpKeepalive) -> Result<()> {
484        match &self.stream_mut().get_mut().stream {
485            RawStream::Tcp(s) => {
486                debug!("Setting tcp keepalive");
487                set_tcp_keepalive(s, ka)?;
488            }
489            RawStream::Virtual(s) => {
490                s.set_socket_option(virt::VirtualSockOpt::KeepAlive(ka.clone()))
491                    .or_err(ConnectError, "failed to set_keepalive on virtual socket")?;
492            }
493            _ => (),
494        }
495        Ok(())
496    }
497
498    #[cfg(target_os = "linux")]
499    pub fn set_rx_timestamp(&mut self) -> Result<()> {
500        use nix::sys::socket::{setsockopt, sockopt, TimestampingFlag};
501
502        if let RawStream::Tcp(s) = &self.stream_mut().get_mut().stream {
503            let timestamp_options = TimestampingFlag::SOF_TIMESTAMPING_RX_SOFTWARE
504                | TimestampingFlag::SOF_TIMESTAMPING_SOFTWARE;
505            setsockopt(&s, sockopt::Timestamping, &timestamp_options)
506                .or_err(InternalError, "failed to set SOF_TIMESTAMPING_RX_SOFTWARE")?;
507            self.stream_mut().get_mut().enable_rx_ts(true);
508        }
509
510        Ok(())
511    }
512
513    #[cfg(not(target_os = "linux"))]
514    pub fn set_rx_timestamp(&mut self) -> io::Result<()> {
515        Ok(())
516    }
517
518    /// Put some data back to the head of the stream to be read again.
519    ///
520    /// This is useful when you've read data to detect a protocol (e.g., PROXY protocol)
521    /// but the data wasn't what you expected, so you need to "unread" it.
522    pub fn rewind(&mut self, data: &[u8]) {
523        if !data.is_empty() {
524            self.rewind_read_buf.push(data.to_vec());
525        }
526    }
527
528    /// Set the buffer of BufStream
529    /// It is only set later because of the malloc overhead in critical accept() path
530    pub(crate) fn set_buffer(&mut self, buffer: L4BufferSettings) {
531        use std::mem;
532        // Since BufStream doesn't provide an API to adjust the buf directly,
533        // we take the raw stream out of it and put it in a new BufStream with the size we want
534        let stream = mem::take(&mut self.stream);
535        let stream = stream.map(|s| {
536            BufStream::with_capacity(
537                buffer.read_capacity(),
538                buffer.write_capacity(),
539                s.into_inner(),
540            )
541        });
542        let _ = mem::replace(&mut self.stream, stream);
543    }
544}
545
546impl From<TcpStream> for Stream {
547    fn from(s: TcpStream) -> Self {
548        Stream {
549            stream: Some(BufStream::with_capacity(
550                0,
551                0,
552                RawStreamWrapper::new(RawStream::Tcp(s)),
553            )),
554            rewind_read_buf: Vec::new(),
555            buffer_write: true,
556            established_ts: SystemTime::now(),
557            establishment_duration: None,
558            offload_wait_duration: None,
559            proxy_digest: None,
560            socket_digest: None,
561            tracer: None,
562            read_pending_time: AccumulatedDuration::new(),
563            write_pending_time: AccumulatedDuration::new(),
564            rx_ts: None,
565        }
566    }
567}
568
569impl From<virt::VirtualSocketStream> for Stream {
570    fn from(s: virt::VirtualSocketStream) -> Self {
571        Stream {
572            stream: Some(BufStream::with_capacity(
573                0,
574                0,
575                RawStreamWrapper::new(RawStream::Virtual(s)),
576            )),
577            rewind_read_buf: Vec::new(),
578            buffer_write: true,
579            established_ts: SystemTime::now(),
580            establishment_duration: None,
581            offload_wait_duration: None,
582            proxy_digest: None,
583            socket_digest: None,
584            tracer: None,
585            read_pending_time: AccumulatedDuration::new(),
586            write_pending_time: AccumulatedDuration::new(),
587            rx_ts: None,
588        }
589    }
590}
591
592#[cfg(unix)]
593impl From<UnixStream> for Stream {
594    fn from(s: UnixStream) -> Self {
595        Stream {
596            stream: Some(BufStream::with_capacity(
597                0,
598                0,
599                RawStreamWrapper::new(RawStream::Unix(s)),
600            )),
601            rewind_read_buf: Vec::new(),
602            buffer_write: true,
603            established_ts: SystemTime::now(),
604            establishment_duration: None,
605            offload_wait_duration: None,
606            proxy_digest: None,
607            socket_digest: None,
608            tracer: None,
609            read_pending_time: AccumulatedDuration::new(),
610            write_pending_time: AccumulatedDuration::new(),
611            rx_ts: None,
612        }
613    }
614}
615
616#[cfg(unix)]
617impl AsRawFd for Stream {
618    fn as_raw_fd(&self) -> std::os::unix::io::RawFd {
619        self.stream().get_ref().as_raw_fd()
620    }
621}
622
623#[cfg(windows)]
624impl AsRawSocket for Stream {
625    fn as_raw_socket(&self) -> std::os::windows::io::RawSocket {
626        self.stream().get_ref().as_raw_socket()
627    }
628}
629
630#[cfg(unix)]
631impl UniqueID for Stream {
632    fn id(&self) -> UniqueIDType {
633        self.as_raw_fd()
634    }
635}
636
637#[cfg(windows)]
638impl UniqueID for Stream {
639    fn id(&self) -> usize {
640        self.as_raw_socket() as usize
641    }
642}
643
644impl Ssl for Stream {}
645
646#[async_trait]
647impl Peek for Stream {
648    async fn try_peek(&mut self, buf: &mut [u8]) -> std::io::Result<bool> {
649        use tokio::io::AsyncReadExt;
650        self.read_exact(buf).await?;
651        // rewind regardless of what is read
652        self.rewind(buf);
653        Ok(true)
654    }
655}
656
657#[async_trait]
658impl Shutdown for Stream {
659    async fn shutdown(&mut self) {
660        AsyncWriteExt::shutdown(self).await.unwrap_or_else(|e| {
661            debug!("Failed to shutdown connection: {:?}", e);
662        });
663    }
664}
665
666impl GetTimingDigest for Stream {
667    fn get_timing_digest(&self) -> Vec<Option<TimingDigest>> {
668        let mut digest = Vec::with_capacity(2); // expect to have both L4 stream and TLS layer
669        digest.push(Some(TimingDigest {
670            established_ts: self.established_ts,
671            establishment_duration: self.establishment_duration,
672            offload_wait_duration: self.offload_wait_duration,
673        }));
674        digest
675    }
676
677    fn get_read_pending_time(&self) -> Duration {
678        self.read_pending_time.total
679    }
680
681    fn get_write_pending_time(&self) -> Duration {
682        self.write_pending_time.total
683    }
684}
685
686impl GetProxyDigest for Stream {
687    fn get_proxy_digest(&self) -> Option<Arc<ProxyDigest>> {
688        self.proxy_digest.clone()
689    }
690
691    fn set_proxy_digest(&mut self, digest: ProxyDigest) {
692        self.proxy_digest = Some(Arc::new(digest));
693    }
694}
695
696impl GetSocketDigest for Stream {
697    fn get_socket_digest(&self) -> Option<Arc<SocketDigest>> {
698        self.socket_digest.clone()
699    }
700
701    fn set_socket_digest(&mut self, socket_digest: SocketDigest) {
702        self.socket_digest = Some(Arc::new(socket_digest))
703    }
704}
705
706impl Drop for Stream {
707    fn drop(&mut self) {
708        if let Some(t) = self.tracer.as_ref() {
709            t.0.on_disconnected();
710        }
711        /* use nodelay/local_addr function to detect socket status */
712        let ret = match &self.stream().get_ref().stream {
713            RawStream::Tcp(s) => s.nodelay().err(),
714            #[cfg(unix)]
715            RawStream::Unix(s) => s.local_addr().err(),
716            RawStream::Virtual(_) => {
717                // TODO: should this do something?
718                None
719            }
720        };
721        if let Some(e) = ret {
722            match e.kind() {
723                tokio::io::ErrorKind::Other => {
724                    if let Some(ecode) = e.raw_os_error() {
725                        if ecode == 9 {
726                            // Or we could panic here
727                            error!("Crit: socket {:?} is being double closed", self.stream);
728                        }
729                    }
730                }
731                _ => {
732                    debug!("Socket is already broken {:?}", e);
733                }
734            }
735        } else {
736            // try flush the write buffer. We use now_or_never() because
737            // 1. Drop cannot be async
738            // 2. write should usually be ready, unless the buf is full.
739            let _ = self.flush().now_or_never();
740        }
741        debug!("Dropping socket {:?}", self.stream);
742    }
743}
744
745impl AsyncRead for Stream {
746    fn poll_read(
747        mut self: Pin<&mut Self>,
748        cx: &mut Context<'_>,
749        buf: &mut ReadBuf<'_>,
750    ) -> Poll<io::Result<()>> {
751        let result = if !self.rewind_read_buf.is_empty() {
752            let data_to_read = self.rewind_read_buf.pop().unwrap(); // safe
753            let mut data_to_read = data_to_read.as_slice();
754            let result = Pin::new(&mut data_to_read).poll_read(cx, buf);
755            // return the remaining data back to the head of rewind_read_buf
756            if !data_to_read.is_empty() {
757                let remaining_buf = Vec::from(data_to_read);
758                self.rewind_read_buf.push(remaining_buf);
759            }
760            result
761        } else {
762            Pin::new(&mut self.stream_mut()).poll_read(cx, buf)
763        };
764        self.read_pending_time.poll_time(&result);
765        self.rx_ts = self.stream().get_ref().rx_ts;
766        result
767    }
768}
769
770impl AsyncWrite for Stream {
771    fn poll_write(
772        mut self: Pin<&mut Self>,
773        cx: &mut Context,
774        buf: &[u8],
775    ) -> Poll<io::Result<usize>> {
776        let result = if self.buffer_write {
777            Pin::new(&mut self.stream_mut()).poll_write(cx, buf)
778        } else {
779            Pin::new(&mut self.stream_mut().get_mut()).poll_write(cx, buf)
780        };
781        self.write_pending_time.poll_write_time(&result, buf.len());
782        result
783    }
784
785    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
786        let result = Pin::new(&mut self.stream_mut()).poll_flush(cx);
787        self.write_pending_time.poll_time(&result);
788        result
789    }
790
791    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
792        Pin::new(&mut self.stream_mut()).poll_shutdown(cx)
793    }
794
795    fn poll_write_vectored(
796        mut self: Pin<&mut Self>,
797        cx: &mut Context<'_>,
798        bufs: &[std::io::IoSlice<'_>],
799    ) -> Poll<io::Result<usize>> {
800        let total_size = bufs.iter().fold(0, |acc, s| acc + s.len());
801
802        let result = if self.buffer_write {
803            Pin::new(&mut self.stream_mut()).poll_write_vectored(cx, bufs)
804        } else {
805            Pin::new(&mut self.stream_mut().get_mut()).poll_write_vectored(cx, bufs)
806        };
807
808        self.write_pending_time.poll_write_time(&result, total_size);
809        result
810    }
811
812    fn is_write_vectored(&self) -> bool {
813        if self.buffer_write {
814            self.stream().is_write_vectored() // it is true
815        } else {
816            self.stream().get_ref().is_write_vectored()
817        }
818    }
819}
820
821#[derive(Debug)]
822struct AccumulatedDuration {
823    total: Duration,
824    last_start: Option<Instant>,
825}
826
827impl AccumulatedDuration {
828    fn new() -> Self {
829        AccumulatedDuration {
830            total: Duration::ZERO,
831            last_start: None,
832        }
833    }
834
835    fn start(&mut self) {
836        if self.last_start.is_none() {
837            self.last_start = Some(Instant::now());
838        }
839    }
840
841    fn stop(&mut self) {
842        if let Some(start) = self.last_start.take() {
843            self.total += start.elapsed();
844        }
845    }
846
847    fn poll_write_time(&mut self, result: &Poll<io::Result<usize>>, buf_size: usize) {
848        match result {
849            Poll::Ready(Ok(n)) => {
850                if *n == buf_size {
851                    self.stop();
852                } else {
853                    // partial write
854                    self.start();
855                }
856            }
857            Poll::Ready(Err(_)) => {
858                self.stop();
859            }
860            _ => self.start(),
861        }
862    }
863
864    fn poll_time(&mut self, result: &Poll<io::Result<()>>) {
865        match result {
866            Poll::Ready(_) => {
867                self.stop();
868            }
869            _ => self.start(),
870        }
871    }
872}
873
874#[cfg(test)]
875#[cfg(target_os = "linux")]
876mod tests {
877    use super::*;
878    use std::sync::Arc;
879    use tokio::io::AsyncReadExt;
880    use tokio::io::AsyncWriteExt;
881    use tokio::net::TcpListener;
882    use tokio::sync::Notify;
883
884    #[cfg(target_os = "linux")]
885    #[tokio::test]
886    async fn test_rx_timestamp() {
887        let message = "hello world".as_bytes();
888        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
889        let addr = listener.local_addr().unwrap();
890        let notify = Arc::new(Notify::new());
891        let notify2 = notify.clone();
892
893        tokio::spawn(async move {
894            let (mut stream, _) = listener.accept().await.unwrap();
895            notify2.notified().await;
896            stream.write_all(message).await.unwrap();
897        });
898
899        let mut stream: Stream = TcpStream::connect(addr).await.unwrap().into();
900        stream.set_rx_timestamp().unwrap();
901        // Receive the message
902        // setsockopt for SO_TIMESTAMPING is asynchronous so sleep a little bit
903        // to let kernel do the work
904        std::thread::sleep(Duration::from_micros(100));
905        notify.notify_one();
906
907        let mut buffer = vec![0u8; message.len()];
908        let n = stream.read(buffer.as_mut_slice()).await.unwrap();
909        assert_eq!(n, message.len());
910        assert!(stream.rx_ts.is_some());
911    }
912
913    #[cfg(target_os = "linux")]
914    #[tokio::test]
915    async fn test_rx_timestamp_standard_path() {
916        let message = "hello world".as_bytes();
917        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
918        let addr = listener.local_addr().unwrap();
919        let notify = Arc::new(Notify::new());
920        let notify2 = notify.clone();
921
922        tokio::spawn(async move {
923            let (mut stream, _) = listener.accept().await.unwrap();
924            notify2.notified().await;
925            stream.write_all(message).await.unwrap();
926        });
927
928        let mut stream: Stream = TcpStream::connect(addr).await.unwrap().into();
929        std::thread::sleep(Duration::from_micros(100));
930        notify.notify_one();
931
932        let mut buffer = vec![0u8; message.len()];
933        let n = stream.read(buffer.as_mut_slice()).await.unwrap();
934        assert_eq!(n, message.len());
935        assert!(stream.rx_ts.is_none());
936    }
937
938    #[tokio::test]
939    async fn test_stream_rewind() {
940        let message = b"hello world";
941        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
942        let addr = listener.local_addr().unwrap();
943        let notify = Arc::new(Notify::new());
944        let notify2 = notify.clone();
945
946        tokio::spawn(async move {
947            let (mut stream, _) = listener.accept().await.unwrap();
948            notify2.notified().await;
949            stream.write_all(message).await.unwrap();
950        });
951
952        let mut stream: Stream = TcpStream::connect(addr).await.unwrap().into();
953
954        let rewind_test = b"this is Sparta!";
955        stream.rewind(rewind_test);
956
957        // partially read rewind_test because of the buffer size limit
958        let mut buffer = vec![0u8; message.len()];
959        let n = stream.read(buffer.as_mut_slice()).await.unwrap();
960        assert_eq!(n, message.len());
961        assert_eq!(buffer, rewind_test[..message.len()]);
962
963        // read the rest of rewind_test
964        let n = stream.read(buffer.as_mut_slice()).await.unwrap();
965        assert_eq!(n, rewind_test.len() - message.len());
966        assert_eq!(buffer[..n], rewind_test[message.len()..]);
967
968        // read the actual data
969        notify.notify_one();
970        let n = stream.read(buffer.as_mut_slice()).await.unwrap();
971        assert_eq!(n, message.len());
972        assert_eq!(buffer, message);
973    }
974
975    #[tokio::test]
976    async fn test_stream_peek() {
977        let message = b"hello world";
978        dbg!("try peek");
979        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
980        let addr = listener.local_addr().unwrap();
981        let notify = Arc::new(Notify::new());
982        let notify2 = notify.clone();
983
984        tokio::spawn(async move {
985            let (mut stream, _) = listener.accept().await.unwrap();
986            notify2.notified().await;
987            stream.write_all(message).await.unwrap();
988            drop(stream);
989        });
990
991        notify.notify_one();
992
993        let mut stream: Stream = TcpStream::connect(addr).await.unwrap().into();
994        let mut buffer = vec![0u8; 5];
995        assert!(stream.try_peek(&mut buffer).await.unwrap());
996        assert_eq!(buffer, message[0..5]);
997        let mut buffer = vec![];
998        stream.read_to_end(&mut buffer).await.unwrap();
999        assert_eq!(buffer, message);
1000    }
1001
1002    #[tokio::test]
1003    async fn test_stream_two_subsequent_peek_calls_before_read() {
1004        let message = b"abcdefghijklmn";
1005
1006        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1007        let addr = listener.local_addr().unwrap();
1008        let notify = Arc::new(Notify::new());
1009        let notify2 = notify.clone();
1010
1011        tokio::spawn(async move {
1012            let (mut stream, _) = listener.accept().await.unwrap();
1013            notify2.notified().await;
1014            stream.write_all(message).await.unwrap();
1015            drop(stream);
1016        });
1017
1018        notify.notify_one();
1019
1020        let mut stream: Stream = TcpStream::connect(addr).await.unwrap().into();
1021
1022        // Peek 4 bytes
1023        let mut buffer = vec![0u8; 4];
1024        assert!(stream.try_peek(&mut buffer).await.unwrap());
1025        assert_eq!(buffer, message[0..4]);
1026
1027        // Peek 2 bytes
1028        let mut buffer = vec![0u8; 2];
1029        assert!(stream.try_peek(&mut buffer).await.unwrap());
1030        assert_eq!(buffer, message[0..2]);
1031
1032        // Read 1 byte: ['a']
1033        let mut buffer = vec![0u8; 1];
1034        stream.read_exact(&mut buffer).await.unwrap();
1035        assert_eq!(buffer, message[0..1]);
1036
1037        // Read as many bytes as possible, return 1 byte ['b']
1038        //  from the first retry buffer chunk
1039        let mut buffer = vec![0u8; 100];
1040        let n = stream.read(&mut buffer).await.unwrap();
1041        assert_eq!(n, 1);
1042        assert_eq!(buffer[..n], message[1..2]);
1043
1044        // Read the rest ['cdefghijklmn']
1045        let mut buffer = vec![];
1046        stream.read_to_end(&mut buffer).await.unwrap();
1047        assert_eq!(buffer, message[2..]);
1048    }
1049}