diff --git a/easytier/src/common/error.rs b/easytier/src/common/error.rs index 87ef7ff4..906499c8 100644 --- a/easytier/src/common/error.rs +++ b/easytier/src/common/error.rs @@ -1,5 +1,4 @@ use std::{io, result}; - use thiserror::Error; use crate::tunnel; @@ -55,4 +54,6 @@ pub enum Error { pub type Result = result::Result; +pub type ErrorCollection = crate::utils::error::ErrorCollection; + // impl From for std:: diff --git a/easytier/src/gateway/kcp_proxy.rs b/easytier/src/gateway/kcp_proxy.rs index 6025d307..b1dceacc 100644 --- a/easytier/src/gateway/kcp_proxy.rs +++ b/easytier/src/gateway/kcp_proxy.rs @@ -4,7 +4,7 @@ use std::{ time::Duration, }; -use anyhow::Context; +use anyhow::{Context, anyhow, bail}; use bytes::Bytes; use dashmap::DashMap; use guarden::defer; @@ -15,12 +15,13 @@ use kcp_sys::{ stream::KcpStream, }; use prost::Message; -use tokio::{select, task::JoinSet}; +use tokio::task::JoinSet; use super::{ CidrSet, tcp_proxy::{NatDstConnector, NatDstTcpConnector, TcpProxy}, }; +use crate::utils::task::HedgeExt; use crate::{ common::{ acl_processor::PacketInfo, @@ -114,72 +115,57 @@ pub struct NatDstKcpConnector { impl NatDstConnector for NatDstKcpConnector { type DstStream = KcpStream; - async fn connect(&self, src: SocketAddr, nat_dst: SocketAddr) -> Result { + async fn connect( + &self, + src: SocketAddr, + nat_dst: SocketAddr, + ) -> anyhow::Result { + let peer_mgr = self + .peer_mgr + .upgrade() + .ok_or_else(|| anyhow!("peer manager is not available"))?; + + let dst_peer = { + let SocketAddr::V4(addr) = nat_dst else { + bail!("ipv6 is not supported"); + }; + peer_mgr + .get_peer_map() + .get_peer_id_by_ipv4(addr.ip()) + .await + .ok_or_else(|| anyhow!("no peer found for nat dst: {}", nat_dst))? + }; + + tracing::trace!(?nat_dst, ?dst_peer, "kcp nat"); + let conn_data = KcpConnData { src: Some(src.into()), dst: Some(nat_dst.into()), }; - let Some(peer_mgr) = self.peer_mgr.upgrade() else { - return Err(anyhow::anyhow!("peer manager is not available").into()); - }; + let stream = (0..5) + .map(|_| { + let kcp_endpoint = self.kcp_endpoint.clone(); + let my_peer_id = peer_mgr.my_peer_id(); - let dst_peer_id = match nat_dst { - SocketAddr::V4(addr) => peer_mgr.get_peer_map().get_peer_id_by_ipv4(addr.ip()).await, - SocketAddr::V6(_) => return Err(anyhow::anyhow!("ipv6 is not supported").into()), - }; + async move { + let conn_id = kcp_endpoint + .connect( + Duration::from_secs(10), + my_peer_id, + dst_peer, + Bytes::from(conn_data.encode_to_vec()), + ) + .await?; - let Some(dst_peer) = dst_peer_id else { - return Err(anyhow::anyhow!("no peer found for nat dst: {}", nat_dst).into()); - }; - - tracing::trace!("kcp nat dst: {:?}, dst peers: {:?}", nat_dst, dst_peer); - - let mut connect_tasks: JoinSet> = JoinSet::new(); - let mut retry_remain = 5; - loop { - select! { - Some(Ok(Ok(ret))) = connect_tasks.join_next() => { - // just wait for the previous connection to finish - let stream = KcpStream::new(&self.kcp_endpoint, ret) - .ok_or(anyhow::anyhow!("failed to create kcp stream"))?; - return Ok(stream); + KcpStream::new(&kcp_endpoint, conn_id).context("failed to create kcp stream") } - _ = tokio::time::sleep(Duration::from_millis(200)), if !connect_tasks.is_empty() && retry_remain > 0 => { - // no successful connection yet, trigger another connection attempt - } - else => { - // got error in connect_tasks, continue to retry - if retry_remain == 0 && connect_tasks.is_empty() { - break; - } - } - } + }) + .hedge(Duration::from_millis(200)) + .await + .context("failed to connect to peer")?; - // create a new connection task - if retry_remain == 0 { - continue; - } - retry_remain -= 1; - - let kcp_endpoint = self.kcp_endpoint.clone(); - let my_peer_id = peer_mgr.my_peer_id(); - let conn_data_clone = conn_data; - - connect_tasks.spawn(async move { - kcp_endpoint - .connect( - Duration::from_secs(10), - my_peer_id, - dst_peer, - Bytes::from(conn_data_clone.encode_to_vec()), - ) - .await - .with_context(|| format!("failed to connect to nat dst: {}", nat_dst)) - }); - } - - Err(anyhow::anyhow!("failed to connect to nat dst: {}", nat_dst).into()) + Ok(stream) } fn check_packet_from_peer_fast(&self, _cidr_set: &CidrSet, _global_ctx: &GlobalCtx) -> bool { diff --git a/easytier/src/gateway/quic_proxy.rs b/easytier/src/gateway/quic_proxy.rs index 0019aaf6..7e0767aa 100644 --- a/easytier/src/gateway/quic_proxy.rs +++ b/easytier/src/gateway/quic_proxy.rs @@ -18,7 +18,8 @@ use crate::tunnel::packet_def::{ PacketType, PeerManagerHeader, TAIL_RESERVED_SIZE, ZCPacket, ZCPacketType, }; use crate::tunnel::quic::{client_config, endpoint_config, server_config}; -use anyhow::{Context, Error, anyhow}; +use crate::utils::task::HedgeExt; +use anyhow::{Context, Error, anyhow, bail, ensure}; use atomic_refcell::AtomicRefCell; use bytes::{BufMut, Bytes, BytesMut}; use dashmap::DashMap; @@ -29,7 +30,8 @@ use moka::future::Cache; use prost::Message; use quinn::udp::{EcnCodepoint, RecvMeta, Transmit}; use quinn::{ - AsyncUdpSocket, Endpoint, RecvStream, SendStream, StreamId, UdpPoller, default_runtime, + AsyncUdpSocket, Connection, ConnectionError, Endpoint, RecvStream, SendStream, StreamId, + UdpPoller, WriteError, default_runtime, }; use std::cmp::min; use std::future::Future; @@ -280,7 +282,7 @@ impl From<(SendStream, RecvStream)> for QuicStream { pub struct NatDstQuicConnector { pub(crate) endpoint: Endpoint, pub(crate) peer_mgr: Weak, - pub(crate) conn_map: Cache, + pub(crate) conn_map: Cache, } #[async_trait::async_trait] @@ -291,20 +293,25 @@ impl NatDstConnector for NatDstQuicConnector { &self, src: SocketAddr, nat_dst: SocketAddr, - ) -> crate::common::error::Result { - let Some(peer_mgr) = self.peer_mgr.upgrade() else { - return Err(anyhow::anyhow!("peer manager is not available").into()); + ) -> anyhow::Result { + let peer_mgr = self + .peer_mgr + .upgrade() + .ok_or_else(|| anyhow!("peer manager is not available"))?; + + let dst_peer = { + let SocketAddr::V4(addr) = nat_dst else { + bail!("ipv6 is not supported"); + }; + peer_mgr + .get_peer_map() + .get_peer_id_by_ipv4(addr.ip()) + .await + .ok_or_else(|| anyhow!("no peer found for nat dst: {}", nat_dst))? }; - let Some(dst_peer_id) = (match nat_dst { - SocketAddr::V4(addr) => peer_mgr.get_peer_map().get_peer_id_by_ipv4(addr.ip()).await, - SocketAddr::V6(_) => return Err(anyhow::anyhow!("ipv6 is not supported").into()), - }) else { - return Err(anyhow::anyhow!("no peer found for nat dst: {}", nat_dst).into()); - }; + tracing::trace!(?nat_dst, ?dst_peer, "quic nat"); - trace!("quic nat dst: {:?}, dst peers: {:?}", nat_dst, dst_peer_id); - let addr = QuicAddr::new(dst_peer_id, PacketType::QuicSrc).into(); let header = { let conn_data = QuicConnData { src: Some(src.into()), @@ -312,77 +319,91 @@ impl NatDstConnector for NatDstQuicConnector { }; let len = conn_data.encoded_len(); - if len > (u16::MAX as usize) { - return Err(anyhow!("conn data too large: {:?}", len).into()); - } + ensure!(len <= u16::MAX as usize, "conn data too large: {len}"); let mut buf = BytesMut::with_capacity(2 + len); buf.put_u16(len as u16); - conn_data.encode(&mut buf).unwrap(); + conn_data.encode(&mut buf)?; buf.freeze() }; - for attempt in 0..2 { - let endpoint = self.endpoint.clone(); + let reconnect = || async move { + self.conn_map.invalidate(&dst_peer).await; - let connection = match self - .conn_map - .try_get_with(dst_peer_id, async move { - endpoint - .connect(addr, "") - .map_err(|e| anyhow!("quic connect: {:#}", e))? - .await - .map_err(|e| anyhow!("quic connection: {:#}", e)) - }) - .await - { - Ok(conn) => conn, - Err(e) => { - if attempt == 0 { - debug!("quic connect failed, retrying: {:#}", e); - tokio::time::sleep(Duration::from_millis(300)).await; - continue; + let connect = (0..5) + .map(|_| { + let endpoint = self.endpoint.clone(); + async move { + endpoint + .connect(QuicAddr::new(dst_peer, PacketType::QuicSrc).into(), "") + .context("failed to create connection")? + .await + .context("connection failed") } - return Err(anyhow!("{:#}", e).into()); - } - }; + }) + .hedge(Duration::from_millis(200)); - let stream: Result = async { + self.conn_map + .try_get_with(dst_peer, connect) + .await + .context("failed to connect to peer") + }; + + let mut reconnected = false; + + let mut connection = if let Some(connection) = self.conn_map.get(&dst_peer).await + && connection.close_reason().is_none() + { + connection + } else { + reconnected = true; + reconnect().await? + }; + + loop { + let is_retryable = |error: &ConnectionError| { + matches!( + error, + ConnectionError::ConnectionClosed(_) + | ConnectionError::ApplicationClosed(_) + | ConnectionError::Reset + | ConnectionError::TimedOut + ) + }; + let mut retry = !reconnected; + let header = header.clone(); + let result = async { let mut stream: QuicStream = connection .open_bi() .await - .map_err(|e| anyhow!("open bi: {:#}", e))? + .inspect_err(|error| retry &= is_retryable(error))? .into(); - stream.writer_mut().write_chunk(header.clone()).await?; + stream + .writer_mut() + .write_chunk(header) + .await + .inspect_err(|error| { + retry &= matches!(error, WriteError::ConnectionLost(error) if is_retryable(error)) + })?; Ok(stream.into()) } - .await; + .await; - match stream { - Ok(stream) => return Ok(stream), - Err(error) => { - debug!( - ?dst_peer_id, - attempt, - ?error, - "quic connect: stream setup failed" - ); + if let Err(error) = &result { + if retry { + debug!(?error, "failed to open quic stream, retrying..."); + reconnected = true; + connection = reconnect().await?; + continue; + } else { + self.conn_map.invalidate(&dst_peer).await; } } - // Evict stale connection; - self.conn_map.invalidate(&dst_peer_id).await; + break result; } - - Err(anyhow!( - "quic connect: failed after {} attempts, dst_peer_id={}, nat_dst={}", - 2, - dst_peer_id, - nat_dst - ) - .into()) } #[inline] @@ -839,7 +860,7 @@ impl QuicProxy { Arc::new(socket), default_runtime().unwrap(), ) - .unwrap(); + .unwrap(); // TODO: maybe a different transport config endpoint.set_default_client_config(client_config()); self.endpoint = Some(endpoint.clone()); @@ -863,26 +884,15 @@ impl QuicProxy { return; } - let conn_map = Cache::builder() - .max_capacity(u8::MAX.into()) // same with max_concurrent_bidi_streams, can be increased - .time_to_idle(Duration::from_secs(600)) - .build(); - - let conn_map_bg = conn_map.clone(); - self.tasks.spawn(async move { - let mut interval = tokio::time::interval(Duration::from_secs(60)); - loop { - interval.tick().await; - conn_map_bg.run_pending_tasks().await; - } - }); - let tcp_proxy = TcpProxyForQuicSrc(TcpProxy::new( peer_mgr.clone(), NatDstQuicConnector { endpoint: endpoint.clone(), peer_mgr: Arc::downgrade(&peer_mgr), - conn_map, + conn_map: Cache::builder() + .max_capacity(u8::MAX.into()) // cf. quinn transport config (max_concurrent_bidi_streams) + .time_to_idle(Duration::from_secs(600)) // cf. quinn transport config (max_idle_timeout) + .build(), }, )); diff --git a/easytier/src/gateway/socks5.rs b/easytier/src/gateway/socks5.rs index fcad40bf..0cd5d07d 100644 --- a/easytier/src/gateway/socks5.rs +++ b/easytier/src/gateway/socks5.rs @@ -240,7 +240,7 @@ impl AsyncTcpConnector for Socks5KcpConnector { let ret = c .connect(self.src_addr, addr) .await - .map_err(|e| super::fast_socks5::SocksError::Other(e.into()))?; + .map_err(super::fast_socks5::SocksError::Other)?; Ok(SocksTcpStream::Kcp(ret)) } } diff --git a/easytier/src/gateway/tcp_proxy.rs b/easytier/src/gateway/tcp_proxy.rs index 3b4075c6..6e252268 100644 --- a/easytier/src/gateway/tcp_proxy.rs +++ b/easytier/src/gateway/tcp_proxy.rs @@ -44,7 +44,7 @@ use super::tokio_smoltcp::{self, Net, NetConfig, channel_device}; pub(crate) trait NatDstConnector: Send + Sync + Clone + 'static { type DstStream: AsyncRead + AsyncWrite + Unpin + Send; - async fn connect(&self, src: SocketAddr, dst: SocketAddr) -> Result; + async fn connect(&self, src: SocketAddr, dst: SocketAddr) -> anyhow::Result; fn check_packet_from_peer_fast(&self, cidr_set: &CidrSet, global_ctx: &GlobalCtx) -> bool; fn check_packet_from_peer( &self, @@ -63,14 +63,13 @@ pub struct NatDstTcpConnector; #[async_trait::async_trait] impl NatDstConnector for NatDstTcpConnector { type DstStream = TcpStream; - async fn connect(&self, _src: SocketAddr, nat_dst: SocketAddr) -> Result { - let socket = match TcpSocket::new_v4() { - Ok(s) => s, - Err(error) => { - log::error!(?error, "create v4 socket failed"); - return Err(error.into()); - } - }; + async fn connect( + &self, + _src: SocketAddr, + nat_dst: SocketAddr, + ) -> anyhow::Result { + let socket = TcpSocket::new_v4() + .inspect_err(|error| log::error!(?error, "create v4 socket failed"))?; let stream = timeout(Duration::from_secs(10), socket.connect(nat_dst)) .await? diff --git a/easytier/src/proto/utils.rs b/easytier/src/proto/utils.rs index c9ab016e..951a9b2c 100644 --- a/easytier/src/proto/utils.rs +++ b/easytier/src/proto/utils.rs @@ -1,6 +1,6 @@ use delegate::delegate; use derivative::Derivative; -use derive_more::{Deref, DerefMut, From, IntoIterator}; +use derive_more::{AsMut, AsRef, Deref, DerefMut, From, IntoIterator}; use prost::Message; use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; @@ -45,11 +45,15 @@ where From, Deref, DerefMut, + AsRef, + AsMut, Serialize, Deserialize, IntoIterator, )] #[derivative(Default(bound = ""))] +#[as_ref(forward)] +#[as_mut(forward)] #[serde(transparent)] #[into_iterator(owned, ref, ref_mut)] pub struct RepeatedMessageModel(Vec); @@ -74,22 +78,6 @@ impl Extend for RepeatedMessageModel { } } -impl AsRef<[Model]> for RepeatedMessageModel { - delegate! { - to self.0 { - fn as_ref(&self) -> &[Model]; - } - } -} - -impl AsMut<[Model]> for RepeatedMessageModel { - delegate! { - to self.0 { - fn as_mut(&mut self) -> &mut [Model]; - } - } -} - impl<'m, Message, Model> TryFrom<&'m [Message]> for RepeatedMessageModel where Message: prost::Message, diff --git a/easytier/src/utils/error.rs b/easytier/src/utils/error.rs new file mode 100644 index 00000000..75711d82 --- /dev/null +++ b/easytier/src/utils/error.rs @@ -0,0 +1,58 @@ +use delegate::delegate; +use derivative::Derivative; +use derive_more::{AsMut, AsRef, Deref, DerefMut, From, Into, IntoIterator}; +use std::fmt; +use std::fmt::Display; +use thiserror::Error; + +#[derive(Derivative, Debug, From, Into, Deref, DerefMut, AsRef, AsMut, IntoIterator, Error)] +#[derivative(Default(bound = ""))] +#[as_ref(forward)] +#[as_mut(forward)] +#[into_iterator(owned, ref, ref_mut)] +pub struct ErrorCollection { + pub errors: Vec, +} + +impl ErrorCollection { + delegate! { + to Vec { + #[into] + pub fn new() -> Self; + #[into] + pub fn with_capacity(capacity: usize) -> Self; + } + } +} + +impl> FromIterator for ErrorCollection { + fn from_iter>(iter: I) -> Self { + Self { + errors: iter.into_iter().map(Into::into).collect(), + } + } +} + +impl Extend for ErrorCollection { + delegate! { + to self.errors { + fn extend>(&mut self, iter: T); + } + } +} + +impl Display for ErrorCollection { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + if self.errors.is_empty() { + return write!(f, "No errors"); + } + + write!(f, "{} error(s) occurred:", self.errors.len())?; + for (i, err) in self.errors.iter().enumerate() { + writeln!(f)?; + write!(f, " {}. {}", i + 1, err)?; + } + + Ok(()) + } +} diff --git a/easytier/src/utils/mod.rs b/easytier/src/utils/mod.rs index 1339c10d..280cd307 100644 --- a/easytier/src/utils/mod.rs +++ b/easytier/src/utils/mod.rs @@ -1,3 +1,4 @@ +pub mod error; pub mod panic; pub mod string; pub mod task; diff --git a/easytier/src/utils/task.rs b/easytier/src/utils/task.rs index b0ac0b40..ce34df56 100644 --- a/easytier/src/utils/task.rs +++ b/easytier/src/utils/task.rs @@ -1,9 +1,13 @@ +use crate::utils::error::ErrorCollection; +use futures::StreamExt; +use futures::stream::FuturesUnordered; use std::future::Future; use std::io; use std::pin::Pin; use std::task::{Context, Poll}; use std::time::Duration; use tokio::task::JoinHandle; +use tokio::time::sleep; use tokio_util::sync::CancellationToken; use tokio_util::task::AbortOnDropHandle; @@ -78,3 +82,61 @@ impl Future for CancellableTask { } // endregion + +// region HedgeExt + +pub(crate) trait HedgeExt: Iterator + Sized { + async fn hedge(self, delay: Duration) -> Result> + where + Self::Item: Future>; +} + +impl HedgeExt for I +where + I: Iterator, +{ + async fn hedge(mut self, delay: Duration) -> Result> + where + Self::Item: Future>, + { + let mut tasks = FuturesUnordered::new(); + let mut errors = ErrorCollection::new(); + let mut exhausted = false; + + macro_rules! spawn { + () => { + if let Some(fut) = self.next() { + tasks.push(fut); + } else { + exhausted = true; + } + }; + } + + spawn!(); + + while !tasks.is_empty() { + tokio::select! { + res = tasks.next() => { + match res { + Some(Ok(v)) => return Ok(v), + Some(Err(e)) => errors.push(e), + None => unreachable!(), + } + + if !exhausted { + spawn!(); + } + } + + _ = sleep(delay), if !exhausted => { + spawn!(); + } + } + } + + Err(errors) + } +} + +// endregion