1use 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 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 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 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 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 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, }
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 RawStream::Virtual(_) => !0,
155 }
156 }
157}
158
159#[derive(Debug)]
160struct RawStreamWrapper {
161 pub(crate) stream: RawStream,
162 pub(crate) rx_ts: Option<SystemTime>,
164 #[cfg(target_os = "linux")]
166 pub(crate) enable_rx_ts: bool,
167 #[cfg(target_os = "linux")]
168 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 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 !self.enable_rx_ts {
222 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 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 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 rs_wrapper.rx_ts =
265 SystemTime::UNIX_EPOCH.checked_add(rtime.system.into());
266 }
267 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 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 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 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 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 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
358pub const DEFAULT_L4_READ_BUFFER_SIZE: usize = 64 * 1024;
364
365pub const DEFAULT_L4_WRITE_BUFFER_SIZE: usize = 1460;
371
372#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
378pub struct L4BufferSettings {
379 pub read: Option<usize>,
381 pub write: Option<usize>,
383}
384
385impl L4BufferSettings {
386 pub fn new(read: usize, write: usize) -> Self {
388 Self {
389 read: Some(read),
390 write: Some(write),
391 }
392 }
393
394 pub fn unbuffered() -> Self {
396 Self::new(0, 0)
397 }
398
399 pub fn read(mut self, read: usize) -> Self {
401 self.read = Some(read);
402 self
403 }
404
405 pub fn write(mut self, write: usize) -> Self {
407 self.write = Some(write);
408 self
409 }
410
411 pub fn read_capacity(&self) -> usize {
413 self.read.unwrap_or(DEFAULT_L4_READ_BUFFER_SIZE)
414 }
415
416 pub fn write_capacity(&self) -> usize {
418 self.write.unwrap_or(DEFAULT_L4_WRITE_BUFFER_SIZE)
419 }
420}
421
422#[derive(Debug)]
427pub struct Stream {
428 stream: Option<BufStream<RawStreamWrapper>>,
430 rewind_read_buf: Vec<Vec<u8>>,
432 buffer_write: bool,
433 proxy_digest: Option<Arc<ProxyDigest>>,
434 socket_digest: Option<Arc<SocketDigest>>,
435 pub established_ts: SystemTime,
437 establishment_duration: Option<Duration>,
438 offload_wait_duration: Option<Duration>,
439 pub tracer: Option<Tracer>,
441 read_pending_time: AccumulatedDuration,
442 write_pending_time: AccumulatedDuration,
443 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 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 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 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, ×tamp_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 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 pub(crate) fn set_buffer(&mut self, buffer: L4BufferSettings) {
531 use std::mem;
532 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 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); 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 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 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 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 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(); let mut data_to_read = data_to_read.as_slice();
754 let result = Pin::new(&mut data_to_read).poll_read(cx, buf);
755 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() } 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 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 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 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 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 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 let mut buffer = vec![0u8; 4];
1024 assert!(stream.try_peek(&mut buffer).await.unwrap());
1025 assert_eq!(buffer, message[0..4]);
1026
1027 let mut buffer = vec![0u8; 2];
1029 assert!(stream.try_peek(&mut buffer).await.unwrap());
1030 assert_eq!(buffer, message[0..2]);
1031
1032 let mut buffer = vec![0u8; 1];
1034 stream.read_exact(&mut buffer).await.unwrap();
1035 assert_eq!(buffer, message[0..1]);
1036
1037 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 let mut buffer = vec![];
1046 stream.read_to_end(&mut buffer).await.unwrap();
1047 assert_eq!(buffer, message[2..]);
1048 }
1049}