Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 19 additions & 10 deletions aether/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -964,7 +964,7 @@ async fn establish_wg(
obfuscate: bool,
keepalive: u16,
label: &'static str,
) -> Result<netstack::StackHandle> {
) -> Result<(netstack::StackHandle, AbortOnDrop)> {
let private_key = identity.private_key_bytes()?;
let peer_public = identity.peer_public_key_bytes()?;

Expand Down Expand Up @@ -1000,19 +1000,27 @@ async fn establish_wg(
let tunnel = wireguard::WgTunnel::from_established(session, std::sync::Arc::new(profile), inbound_tx, ipv4);
let stack = netstack::spawn(&identity.ipv4, &identity.ipv6, mtu, inbound_rx, outbound_tx)?;

tokio::spawn(async move {
let handle = tokio::spawn(async move {
if let Err(e) = tunnel.run(outbound_rx).await {
log::error!("[{label}] wireguard tunnel exited: {e}");
}
});

Ok(stack)
Ok((stack, AbortOnDrop(handle)))
}

struct AbortOnDrop(tokio::task::JoinHandle<()>);

impl Drop for AbortOnDrop {
fn drop(&mut self) {
self.0.abort();
}
}

async fn spawn_udp_forwarder(
outer: &netstack::StackHandle,
remote: SocketAddr,
) -> Result<SocketAddr> {
) -> Result<(SocketAddr, AbortOnDrop, AbortOnDrop)> {
let sock = std::sync::Arc::new(tokio::net::UdpSocket::bind("127.0.0.1:0").await?);
let local = sock.local_addr()?;

Expand All @@ -1024,7 +1032,7 @@ async fn spawn_udp_forwarder(

let up_sock = sock.clone();
let up_peer = inner_peer.clone();
tokio::spawn(async move {
let up = tokio::spawn(async move {
let mut buf = vec![0u8; 65536];
loop {
match up_sock.recv_from(&mut buf).await {
Expand All @@ -1041,7 +1049,7 @@ async fn spawn_udp_forwarder(

let down_sock = sock.clone();
let down_peer = inner_peer.clone();
tokio::spawn(async move {
let down = tokio::spawn(async move {
while let Some((_src, data)) = udp_rx.recv().await {
let dst = *down_peer.lock().await;
if let Some(dst) = dst {
Expand All @@ -1050,7 +1058,7 @@ async fn spawn_udp_forwarder(
}
});

Ok(local)
Ok((local, AbortOnDrop(up), AbortOnDrop(down)))
}

async fn run_warp_in_warp(
Expand All @@ -1060,15 +1068,16 @@ async fn run_warp_in_warp(
listen: SocketAddr,
) -> Result<()> {
log::info!("[*] establishing outer WARP tunnel to {peer}...");
let outer_stack = establish_wg(&primary, peer, TUNNEL_MTU, true, 5, "outer").await?;
let (outer_stack, _outer_tunnel) = establish_wg(&primary, peer, TUNNEL_MTU, true, 5, "outer").await?;

let forwarder = spawn_udp_forwarder(&outer_stack, peer).await?;
let (forwarder, _fwd_up, _fwd_down) = spawn_udp_forwarder(&outer_stack, peer).await?;
log::info!("[+] inner endpoint tunneled through outer warp via {forwarder}");

log::info!("[*] establishing inner WARP tunnel (warp-in-warp)...");
let inner_stack = establish_wg(&secondary, forwarder, INNER_MTU, false, 20, "inner").await?;
let (inner_stack, _inner_tunnel) = establish_wg(&secondary, forwarder, INNER_MTU, false, 20, "inner").await?;

log::info!("[+] socks5 server listening on {listen}");
// Guards above abort outer/inner tunnels + forwarders when this returns.
socks::serve(listen, inner_stack).await
}

Expand Down
7 changes: 7 additions & 0 deletions aether/src/masque.rs
Original file line number Diff line number Diff line change
Expand Up @@ -120,12 +120,19 @@ pub struct CapsuleParser {
buf: Vec<u8>,
}

const MAX_CAPSULE_BUF: usize = 256 * 1024;

impl CapsuleParser {
pub fn new() -> Self {
Self { buf: Vec::new() }
}

pub fn push(&mut self, data: &[u8]) {
if self.buf.len().saturating_add(data.len()) > MAX_CAPSULE_BUF {
log::warn!("capsule parser buffer exceeded {MAX_CAPSULE_BUF} bytes; resetting");
self.buf.clear();
return;
}
self.buf.extend_from_slice(data);
}

Expand Down
133 changes: 120 additions & 13 deletions aether/src/netstack.rs
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,10 @@ fn app_queue() -> usize {
const MAX_INGEST_PER_TICK: usize = 512;
const MAX_RECV_CHUNKS: usize = 128;

fn max_tcp_pending() -> usize {
tcp_buf().saturating_mul(2).max(64 * 1024)
}

type OpenTcpResp = oneshot::Sender<std::result::Result<TcpConn, String>>;
type OpenUdpResp = oneshot::Sender<std::result::Result<UdpConn, String>>;

Expand Down Expand Up @@ -115,6 +119,7 @@ pub struct TcpConn {
pub id: usize,
pub from_stack: mpsc::Receiver<Vec<u8>>,
data_in: mpsc::Sender<DataIn>,
split: bool,
}

impl TcpConn {
Expand All @@ -129,17 +134,32 @@ impl TcpConn {
let _ = self.data_in.send(DataIn::TcpClose(self.id)).await;
}

pub fn into_split(self) -> (TcpSender, mpsc::Receiver<Vec<u8>>) {
pub fn into_split(mut self) -> (TcpSender, mpsc::Receiver<Vec<u8>>) {
self.split = true;
(
TcpSender {
id: self.id,
data_in: self.data_in,
data_in: self.data_in.clone(),
},
self.from_stack,
std::mem::replace(
&mut self.from_stack,
{
let (_tx, rx) = mpsc::channel(1);
rx
},
),
)
}
}

impl Drop for TcpConn {
fn drop(&mut self) {
if !self.split {
let _ = self.data_in.try_send(DataIn::TcpClose(self.id));
}
}
}

pub struct TcpSender {
id: usize,
data_in: mpsc::Sender<DataIn>,
Expand All @@ -158,10 +178,17 @@ impl TcpSender {
}
}

impl Drop for TcpSender {
fn drop(&mut self) {
let _ = self.data_in.try_send(DataIn::TcpClose(self.id));
}
}

pub struct UdpConn {
pub id: usize,
pub from_stack: mpsc::Receiver<(SocketAddr, Vec<u8>)>,
data_in: mpsc::Sender<DataIn>,
split: bool,
}

impl UdpConn {
Expand All @@ -176,17 +203,32 @@ impl UdpConn {
let _ = self.data_in.send(DataIn::UdpClose(self.id)).await;
}

pub fn into_split(self) -> (UdpSender, mpsc::Receiver<(SocketAddr, Vec<u8>)>) {
pub fn into_split(mut self) -> (UdpSender, mpsc::Receiver<(SocketAddr, Vec<u8>)>) {
self.split = true;
(
UdpSender {
id: self.id,
data_in: self.data_in,
data_in: self.data_in.clone(),
},
self.from_stack,
std::mem::replace(
&mut self.from_stack,
{
let (_tx, rx) = mpsc::channel(1);
rx
},
),
)
}
}

impl Drop for UdpConn {
fn drop(&mut self) {
if !self.split {
let _ = self.data_in.try_send(DataIn::UdpClose(self.id));
}
}
}

pub struct UdpSender {
id: usize,
data_in: mpsc::Sender<DataIn>,
Expand All @@ -205,6 +247,12 @@ impl UdpSender {
}
}

impl Drop for UdpSender {
fn drop(&mut self) {
let _ = self.data_in.try_send(DataIn::UdpClose(self.id));
}
}

#[derive(Clone)]
pub struct StackHandle {
cmd_tx: mpsc::Sender<Cmd>,
Expand Down Expand Up @@ -423,6 +471,8 @@ async fn run(
mut inbound_rx: mpsc::Receiver<Vec<u8>>,
outbound_tx: mpsc::Sender<Vec<u8>>,
) -> Result<()> {
let mut deferred: VecDeque<DataIn> = VecDeque::new();

loop {
let now = Instant::now();
let poll_outcome = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
Expand All @@ -436,10 +486,22 @@ async fn run(
service_udp(&mut s).await;
flush_tx(&mut s, &outbound_tx).await;

while let Some(d) = deferred.pop_front() {
if let Some(back) = try_handle_data(&mut s, d) {
deferred.push_front(back);
break;
}
}

let delay = s
.iface
.poll_delay(Instant::now(), &s.sockets)
.map(|d| std::time::Duration::from_micros(d.total_micros()));
let delay = if deferred.is_empty() {
delay
} else {
Some(delay.unwrap_or(std::time::Duration::from_millis(5)))
};

tokio::select! {
biased;
Expand Down Expand Up @@ -467,11 +529,21 @@ async fn run(
}
}

maybe = data_in_rx.recv() => {
maybe = data_in_rx.recv(), if deferred.is_empty() => {
if let Some(d) = maybe {
handle_data(&mut s, d);
while let Ok(d2) = data_in_rx.try_recv() {
handle_data(&mut s, d2);
if let Some(back) = try_handle_data(&mut s, d) {
deferred.push_back(back);
} else {
while deferred.is_empty() {
match data_in_rx.try_recv() {
Ok(d2) => {
if let Some(back) = try_handle_data(&mut s, d2) {
deferred.push_back(back);
}
}
Err(_) => break,
}
}
}
}
}
Expand Down Expand Up @@ -547,6 +619,7 @@ fn handle_cmd(s: &mut NetStack, cmd: Cmd) {
id,
from_stack: to_app_rx,
data_in: s.data_in_tx.clone(),
split: false,
};
let _ = resp.send(Ok(conn));
}
Expand All @@ -557,28 +630,43 @@ fn handle_cmd(s: &mut NetStack, cmd: Cmd) {
}
}

fn handle_data(s: &mut NetStack, d: DataIn) {
/// Returns `Some(d)` when the datagram must be deferred (TCP pending full).
fn try_handle_data(s: &mut NetStack, d: DataIn) -> Option<DataIn> {
match d {
DataIn::Tcp(id, data) => {
if let Some(st) = s.tcp_conns.get_mut(&id) {
st.pending.extend_from_slice(&data);
let max = max_tcp_pending();
if st.pending.len() >= max {
return Some(DataIn::Tcp(id, data));
}
let space = max - st.pending.len();
if data.len() <= space {
st.pending.extend_from_slice(&data);
} else {
st.pending.extend_from_slice(&data[..space]);
return Some(DataIn::Tcp(id, data[space..].to_vec()));
}
}
None
}
DataIn::TcpClose(id) => {
if let Some(st) = s.tcp_conns.get_mut(&id) {
st.half_closed = true;
}
None
}
DataIn::Udp(id, dst, data) => {
if let Some(st) = s.udp_conns.get(&id) {
let sock = s.sockets.get_mut::<udp::Socket>(st.handle);
let _ = sock.send_slice(&data, to_ip_endpoint(dst));
}
None
}
DataIn::UdpClose(id) => {
if let Some(st) = s.udp_conns.remove(&id) {
s.sockets.remove(st.handle);
}
None
}
}
}
Expand All @@ -603,6 +691,7 @@ async fn service_tcp(s: &mut NetStack) {
id,
from_stack: rx,
data_in: data_in_tx.clone(),
split: false,
};
let _ = resp.send(Ok(conn));
}
Expand Down Expand Up @@ -630,6 +719,9 @@ async fn service_tcp(s: &mut NetStack) {
let sent = socket.send_slice(&st.pending).unwrap_or(0);
if sent > 0 {
st.pending.drain(0..sent);
if st.pending.len() * 4 < st.pending.capacity() {
st.pending.shrink_to(max_tcp_pending().min(st.pending.capacity()));
}
}
}
}
Expand Down Expand Up @@ -673,7 +765,15 @@ async fn service_tcp(s: &mut NetStack) {
if matches!(st_state, tcp::State::CloseWait) {
s.sockets.get_mut::<tcp::Socket>(handle).close();
}
if matches!(st_state, tcp::State::Closed) && s.tcp_conns[&id].established {
if matches!(st_state, tcp::State::TimeWait) {
if let Some(st) = s.tcp_conns.get_mut(&id) {
st.pending.clear();
st.pending.shrink_to_fit();
}
}
if matches!(st_state, tcp::State::Closed | tcp::State::TimeWait)
&& s.tcp_conns[&id].established
{
s.sockets.remove(handle);
s.tcp_conns.remove(&id);
}
Expand Down Expand Up @@ -703,11 +803,18 @@ async fn service_udp(s: &mut NetStack) {
}

let to_app = s.udp_conns[&id].to_app.clone();
let mut app_gone = to_app.is_closed();
for p in packets {
if to_app.send(p).await.is_err() {
app_gone = true;
break;
}
}
if app_gone {
if let Some(st) = s.udp_conns.remove(&id) {
s.sockets.remove(st.handle);
}
}
}
}

Expand Down
Loading