diff --git a/easytier/src/peers/peer_conn.rs b/easytier/src/peers/peer_conn.rs index 8b9b6fe6..362f4631 100644 --- a/easytier/src/peers/peer_conn.rs +++ b/easytier/src/peers/peer_conn.rs @@ -363,7 +363,7 @@ impl PeerConn { let throughput = peer_conn_tunnel_filter.filter_output(); let filter_chain = TunnelFilterChain::new(session_filter.clone(), peer_conn_tunnel_filter); let peer_conn_tunnel = TunnelWithFilter::new(tunnel, filter_chain); - let mut mpsc_tunnel = MpscTunnel::new(peer_conn_tunnel, Some(Duration::from_secs(7))); + let mut mpsc_tunnel = MpscTunnel::new_direct(peer_conn_tunnel); let (recv, sink) = (mpsc_tunnel.get_stream(), mpsc_tunnel.get_sink()); diff --git a/easytier/src/tunnel/mpsc.rs b/easytier/src/tunnel/mpsc.rs index e15231ae..5d140fb8 100644 --- a/easytier/src/tunnel/mpsc.rs +++ b/easytier/src/tunnel/mpsc.rs @@ -1,8 +1,9 @@ // this mod wrap tunnel to a mpsc tunnel, based on crossbeam_channel -use std::{pin::Pin, time::Duration}; +use std::{pin::Pin, sync::Arc, time::Duration}; use anyhow::Context; +use tokio::sync::Mutex; use tokio::time::timeout; use crate::proto::common::TunnelInfo; @@ -11,21 +12,38 @@ use super::{Tunnel, TunnelError, ZCPacketSink, ZCPacketStream, packet_def::ZCPac use tokio::sync::mpsc::{Receiver, Sender, channel, error::TrySendError}; use tokio_util::task::AbortOnDropHandle; -// use tachyonix::{channel, Receiver, Sender, TrySendError}; use futures::SinkExt; #[derive(Clone)] -pub struct MpscTunnelSender(Sender); +pub struct MpscTunnelSender { + channel_tx: Option>, + direct_sink: Option>>>>, +} impl MpscTunnelSender { pub async fn send(&self, item: ZCPacket) -> Result<(), TunnelError> { - self.0.send(item).await.with_context(|| "send error")?; - Ok(()) + if let Some(sink) = &self.direct_sink { + let mut guard = sink.lock().await; + guard.feed(item).await?; + guard.flush().await?; + return Ok(()); + } + + let tx = self.channel_tx.as_ref().ok_or(TunnelError::Shutdown)?; + match tx.try_send(item) { + Ok(()) => Ok(()), + Err(TrySendError::Full(item)) => { + tx.send(item).await.with_context(|| "send error")?; + Ok(()) + } + Err(TrySendError::Closed(_)) => Err(TunnelError::Shutdown), + } } pub fn try_send(&self, item: ZCPacket) -> Result<(), TunnelError> { - self.0.try_send(item).map_err(|e| match e { + let tx = self.channel_tx.as_ref().ok_or(TunnelError::Shutdown)?; + tx.try_send(item).map_err(|e| match e { TrySendError::Full(_) => TunnelError::BufferFull, TrySendError::Closed(_) => TunnelError::Shutdown, }) @@ -34,11 +52,12 @@ impl MpscTunnelSender { pub struct MpscTunnel { tx: Option>, + direct_sink: Option>>>>, tunnel: T, stream: Option>>, - task: AbortOnDropHandle<()>, + task: Option>, } impl MpscTunnel { @@ -60,9 +79,21 @@ impl MpscTunnel { Self { tx: Some(tx), + direct_sink: None, tunnel, stream: Some(stream), - task: AbortOnDropHandle::new(task), + task: Some(AbortOnDropHandle::new(task)), + } + } + + pub fn new_direct(tunnel: T) -> Self { + let (stream, sink) = tunnel.split(); + Self { + tx: None, + direct_sink: Some(Arc::new(Mutex::new(sink))), + tunnel, + stream: Some(stream), + task: None, } } @@ -124,12 +155,18 @@ impl MpscTunnel { } pub fn get_sink(&self) -> MpscTunnelSender { - MpscTunnelSender(self.tx.as_ref().unwrap().clone()) + MpscTunnelSender { + channel_tx: self.tx.as_ref().cloned(), + direct_sink: self.direct_sink.clone(), + } } pub fn close(&mut self) { self.tx.take(); - self.task.abort(); + self.direct_sink.take(); + if let Some(task) = self.task.take() { + task.abort(); + } } pub fn tunnel_info(&self) -> Option {