From 6f97250234b2377bedd4c7a96d8e9a8a890c6f99 Mon Sep 17 00:00:00 2001 From: rustdesk Date: Fri, 4 Sep 2026 00:33:52 +0800 Subject: [PATCH] server: PortForwardMux channel table and per-channel tasks Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01EZ49AbZJYfm8NTp5yDPMab --- src/server.rs | 1 + src/server/connection.rs | 2 +- src/server/port_forward_mux.rs | 496 +++++++++++++++++++++++++++++++++ 3 files changed, 498 insertions(+), 1 deletion(-) create mode 100644 src/server/port_forward_mux.rs diff --git a/src/server.rs b/src/server.rs index 5af982772..753ad5903 100644 --- a/src/server.rs +++ b/src/server.rs @@ -70,6 +70,7 @@ pub mod input_service { mod connection; mod login_failure_check; +pub(crate) mod port_forward_mux; pub mod display_service; #[cfg(windows)] pub mod portable_service; diff --git a/src/server/connection.rs b/src/server/connection.rs index 28c4a25fa..ec7fef120 100644 --- a/src/server/connection.rs +++ b/src/server/connection.rs @@ -1645,7 +1645,7 @@ impl Connection { } } - fn normalize_port_forward_target(pf: &mut PortForward) -> (String, bool) { + pub(super) fn normalize_port_forward_target(pf: &mut PortForward) -> (String, bool) { let mut is_rdp = false; if pf.host == "RDP" && pf.port == 0 { pf.host = "localhost".to_owned(); diff --git a/src/server/port_forward_mux.rs b/src/server/port_forward_mux.rs new file mode 100644 index 000000000..2e659bdab --- /dev/null +++ b/src/server/port_forward_mux.rs @@ -0,0 +1,496 @@ +use super::connection::{Connection, Sender}; +use crate::port_forward_mux::{ + charge, close_msg, effective_window, opened_msg, run_channel, FrameSink, Inbound, RecvWindow, + SendCredit, CHANNEL_WINDOW, INITIAL_WINDOW, MAX_CHANNELS, +}; +use hbb_common::{ + bytes::Bytes, + log, + message_proto::*, + timeout, + tokio::{self, net::TcpStream, sync::mpsc}, +}; +use std::{ + collections::{BTreeSet, HashMap}, + sync::{Arc, Mutex}, +}; + +const CONNECT_TIMEOUT_MS: u64 = 3000; + +/// Before `opened` the controller may only have used `INITIAL_WINDOW`. +/// `charged` is the running total of `charge(len)`, not of raw lengths. +fn pending_fits(charged: usize, add_len: usize) -> bool { + charged.saturating_add(charge(add_len) as usize) <= INITIAL_WINDOW as usize +} + +struct Entry { + inbound: mpsc::UnboundedSender, + credit: Arc, + window: Arc>, + target: String, +} + +/// The controlled side of one multiplexed tunnel. The main loop owns it and +/// forwards every `PortForwardChannel` frame here; each channel is a task. +pub struct PortForwardMux { + channels: HashMap, + tx: Sender, + login_target: String, + last_label: String, +} + +impl PortForwardMux { + pub fn new(tx: Sender, login_target: String) -> Self { + let last_label = login_target.clone(); + Self { + channels: HashMap::new(), + tx, + login_target, + last_label, + } + } + + pub fn handle(&mut self, frame: PortForwardChannel, permitted: bool) { + match frame.union { + Some(port_forward_channel::Union::Open(open)) => self.on_open(open, permitted), + Some(port_forward_channel::Union::Data(d)) => { + let len = d.data.len(); + let Some(entry) = self.channels.get(&d.channel_id) else { + log::debug!("port forward data for unknown channel {}", d.channel_id); + return; + }; + let accepted = entry.window.lock().unwrap().accept(len); + let msg = if accepted { + Inbound::Data(d.data) + } else { + log::warn!("port forward channel {} overran its window", d.channel_id); + Inbound::Violation + }; + if entry.inbound.send(msg).is_err() { + self.channels.remove(&d.channel_id); + } + } + Some(port_forward_channel::Union::Close(c)) => { + if let Some(entry) = self.channels.remove(&c.channel_id) { + entry.inbound.send(Inbound::Close).ok(); + } else { + log::debug!("port forward close for unknown channel {}", c.channel_id); + } + } + Some(port_forward_channel::Union::WindowUpdate(u)) => { + match self.channels.get(&u.channel_id) { + Some(entry) => entry.credit.add(u.add), + None => log::debug!( + "port forward window update for unknown channel {}", + u.channel_id + ), + } + } + Some(port_forward_channel::Union::Opened(o)) => { + log::debug!("ignoring opened for channel {} on the controlled side", o.channel_id); + } + _ => {} + } + } + + fn on_open(&mut self, open: PortForwardOpen, permitted: bool) { + let id = open.channel_id; + self.channels.retain(|_, e| !e.inbound.is_closed()); + if !permitted { + self.reply(opened_msg(id, false, "No permission of IP tunneling", 0)); + return; + } + if self.channels.len() >= MAX_CHANNELS { + self.reply(opened_msg(id, false, "Too many port forward channels", 0)); + return; + } + if self.channels.contains_key(&id) { + log::debug!("ignoring open for live channel {}", id); + return; + } + let mut pf = PortForward { + host: open.host, + port: open.port, + ..Default::default() + }; + let (addr, _is_rdp) = Connection::normalize_port_forward_target(&mut pf); + let (inbound_tx, inbound_rx) = mpsc::unbounded_channel(); + let credit = Arc::new(SendCredit::new(effective_window(open.window))); + let window = Arc::new(Mutex::new(RecvWindow::new(CHANNEL_WINDOW))); + self.channels.insert( + id, + Entry { + inbound: inbound_tx, + credit: credit.clone(), + window: window.clone(), + target: addr.clone(), + }, + ); + tokio::spawn(run_controlled_channel( + id, + addr, + credit, + window, + inbound_rx, + FrameSink::Direct(self.tx.clone()), + )); + } + + fn reply(&self, msg: Message) { + self.tx + .send((tokio::time::Instant::now(), Arc::new(msg))) + .ok(); + } + + /// Drops dead entries and recomputes the CM label from the survivors. + /// Returns the label only when it differs from the last one returned. + pub fn sweep(&mut self) -> Option { + self.channels.retain(|_, e| !e.inbound.is_closed()); + let targets: BTreeSet<&str> = self.channels.values().map(|e| e.target.as_str()).collect(); + let label = if targets.is_empty() { + self.login_target.clone() + } else { + let first = if targets.contains(self.login_target.as_str()) { + self.login_target.as_str() + } else { + targets.iter().next().copied().unwrap_or(self.login_target.as_str()) + }; + match targets.len() - 1 { + 0 => first.to_owned(), + n => format!("{} +{}", first, n), + } + }; + if label == self.last_label { + return None; + } + self.last_label = label.clone(); + Some(label) + } + + #[cfg(test)] + pub fn live_channels(&self) -> usize { + self.channels.len() + } + + /// Dropping every sender ends every task; each drops its target socket. + pub fn close_all(&mut self) { + self.channels.clear(); + } +} + +/// Owns the whole channel lifecycle: connect under a `select!` in which a +/// queued command always wins over the connect, buffer what arrives +/// meanwhile, then relay. +async fn run_controlled_channel( + id: i32, + addr: String, + credit: Arc, + window: Arc>, + mut inbound: mpsc::UnboundedReceiver, + sink: FrameSink, +) { + let mut pending: Vec = Vec::new(); + let mut pending_len = 0usize; + let connect = timeout(CONNECT_TIMEOUT_MS, TcpStream::connect(&addr)); + tokio::pin!(connect); + let socket = loop { + tokio::select! { + // Biased with the command arm first: a `close` that is already + // queued must win over a connect that completed on the same poll, + // or `opened` would go out for a channel the controller has dropped. + biased; + cmd = inbound.recv() => match cmd { + Some(Inbound::Data(b)) => { + if !pending_fits(pending_len, b.len()) { + log::warn!("port forward channel {} sent more than INITIAL_WINDOW before opened", id); + sink.send_ordered(close_msg(id)).await.ok(); + return; + } + pending_len += charge(b.len()) as usize; + pending.push(b); + } + Some(Inbound::Close) | None => return, + Some(Inbound::Violation) => { + sink.send_ordered(close_msg(id)).await.ok(); + return; + } + }, + res = &mut connect => match res { + Ok(Ok(s)) => break s, + Ok(Err(e)) => { + log::debug!("port forward channel {} connect {} failed: {}", id, addr, e); + sink.send_ordered(opened_msg(id, false, &format!("Failed to access remote {}", addr), 0)).await.ok(); + return; + } + Err(_) => { + log::debug!("port forward channel {} connect {} timed out", id, addr); + sink.send_ordered(opened_msg(id, false, &format!("Failed to access remote {}", addr), 0)).await.ok(); + return; + } + }, + } + }; + if sink + .send_ordered(opened_msg(id, true, "", CHANNEL_WINDOW)) + .await + .is_err() + { + return; + } + let (reader, writer) = socket.into_split(); + run_channel(id, reader, writer, Vec::new(), pending, credit, window, inbound, sink).await; +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::port_forward_mux::{CHANNEL_WINDOW, INITIAL_WINDOW, MAX_CHANNELS, MIN_FRAME_CHARGE}; + use hbb_common::{ + message_proto::{message, port_forward_channel}, + tokio::{ + self, + io::{AsyncReadExt, AsyncWriteExt}, + net::TcpListener, + sync::mpsc, + time::Instant, + }, + }; + + fn rt() -> tokio::runtime::Runtime { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap() + } + + /// An echo server standing in for the forward target. + async fn echo_target() -> u16 { + let l = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let port = l.local_addr().unwrap().port(); + tokio::spawn(async move { + loop { + let (mut s, _) = l.accept().await.unwrap(); + tokio::spawn(async move { + let mut buf = [0u8; 4096]; + loop { + let n = s.read(&mut buf).await.unwrap_or(0); + if n == 0 || s.write_all(&buf[..n]).await.is_err() { + return; + } + } + }); + } + }); + port + } + + fn open(id: i32, port: u16) -> PortForwardChannel { + let mut ch = PortForwardChannel::new(); + ch.set_open(PortForwardOpen { + channel_id: id, + host: "127.0.0.1".to_owned(), + port: port as i32, + window: CHANNEL_WINDOW, + ..Default::default() + }); + ch + } + + fn data(id: i32, bytes: &[u8]) -> PortForwardChannel { + let mut ch = PortForwardChannel::new(); + ch.set_data(PortForwardData { + channel_id: id, + data: Bytes::copy_from_slice(bytes), + ..Default::default() + }); + ch + } + + fn close(id: i32) -> PortForwardChannel { + let mut ch = PortForwardChannel::new(); + ch.set_close(PortForwardClose { channel_id: id, ..Default::default() }); + ch + } + + async fn next_frame(rx: &mut mpsc::UnboundedReceiver<(Instant, Arc)>) -> PortForwardChannel { + let (_, m) = rx.recv().await.unwrap(); + match &m.union { + Some(message::Union::PortForwardChannel(ch)) => ch.clone(), + other => panic!("unexpected {:?}", other), + } + } + + fn opened(ch: &PortForwardChannel) -> (i32, bool) { + match &ch.union { + Some(port_forward_channel::Union::Opened(o)) => (o.channel_id, o.success), + other => panic!("expected opened, got {:?}", other), + } + } + + fn data_of(ch: &PortForwardChannel) -> (i32, Vec) { + match &ch.union { + Some(port_forward_channel::Union::Data(d)) => (d.channel_id, d.data.to_vec()), + other => panic!("expected data, got {:?}", other), + } + } + + #[test] + fn open_connects_and_echoes_pipelined_data() { + rt().block_on(async { + let port = echo_target().await; + let (tx, mut rx) = mpsc::unbounded_channel(); + let mut mux = PortForwardMux::new(tx, "localhost:1".to_owned()); + mux.handle(open(1, port), true); + mux.handle(data(1, b"ping"), true); + assert_eq!(opened(&next_frame(&mut rx).await), (1, true)); + assert_eq!(data_of(&next_frame(&mut rx).await), (1, b"ping".to_vec())); + mux.handle(close(1), true); + }); + } + + #[test] + fn unreachable_target_fails_open_and_discards_pipelined_data() { + rt().block_on(async { + let l = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let port = l.local_addr().unwrap().port(); + drop(l); + let (tx, mut rx) = mpsc::unbounded_channel(); + let mut mux = PortForwardMux::new(tx, "localhost:1".to_owned()); + mux.handle(open(1, port), true); + mux.handle(data(1, b"lost"), true); + assert_eq!(opened(&next_frame(&mut rx).await), (1, false)); + assert!(tokio::time::timeout(std::time::Duration::from_millis(50), rx.recv()).await.is_err()); + }); + } + + #[test] + fn permission_denied_refuses_without_spawning() { + rt().block_on(async { + let port = echo_target().await; + let (tx, mut rx) = mpsc::unbounded_channel(); + let mut mux = PortForwardMux::new(tx, "localhost:1".to_owned()); + mux.handle(open(1, port), false); + assert_eq!(opened(&next_frame(&mut rx).await), (1, false)); + assert_eq!(mux.live_channels(), 0); + }); + } + + #[test] + fn close_while_connecting_sends_no_opened() { + rt().block_on(async { + // `close` is queued before the task is first polled. Its `select!` is + // biased towards the command arm, so even a connect that completes on + // that same poll loses: no `opened` may ever be sent. + let port = echo_target().await; + let (tx, mut rx) = mpsc::unbounded_channel(); + let mut mux = PortForwardMux::new(tx, "localhost:1".to_owned()); + mux.handle(open(1, port), true); + mux.handle(close(1), true); + assert!(tokio::time::timeout(std::time::Duration::from_millis(200), rx.recv()).await.is_err()); + assert_eq!(mux.live_channels(), 0); + }); + } + + #[test] + fn over_window_data_closes_only_that_channel() { + rt().block_on(async { + let port = echo_target().await; + let (tx, mut rx) = mpsc::unbounded_channel(); + let mut mux = PortForwardMux::new(tx, "localhost:1".to_owned()); + mux.handle(open(1, port), true); + mux.handle(open(2, port), true); + let mut seen = 0; + while seen < 2 { + opened(&next_frame(&mut rx).await); + seen += 1; + } + let too_much = vec![0u8; CHANNEL_WINDOW as usize + 1]; + mux.handle(data(1, &too_much), true); + let ch = next_frame(&mut rx).await; + match &ch.union { + Some(port_forward_channel::Union::Close(c)) => assert_eq!(c.channel_id, 1), + other => panic!("expected close, got {:?}", other), + } + mux.handle(data(2, b"still fine"), true); + assert_eq!(data_of(&next_frame(&mut rx).await), (2, b"still fine".to_vec())); + }); + } + + #[test] + fn pending_bytes_are_bounded_by_initial_window_before_opened() { + // A loopback connect completes before a task can observe "connecting", + // so the bound is pinned on the pure predicate the task uses. + assert!(pending_fits(0, INITIAL_WINDOW as usize)); + assert!(pending_fits( + INITIAL_WINDOW as usize - MIN_FRAME_CHARGE as usize, + 1 + )); + // A 1-byte frame costs a whole minimum charge here too. + assert!(!pending_fits( + INITIAL_WINDOW as usize - MIN_FRAME_CHARGE as usize + 1, + 1 + )); + assert!(!pending_fits(usize::MAX, 1)); + } + + /// A target that accepts and hangs up at once, so every channel ends on + /// the target's EOF — the case where only the sweep can free the entry. + async fn drop_target() -> u16 { + let l = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let port = l.local_addr().unwrap().port(); + tokio::spawn(async move { + loop { + let (s, _) = l.accept().await.unwrap(); + drop(s); + } + }); + port + } + + #[test] + fn sweep_frees_dead_entries_so_the_cap_counts_live_channels() { + rt().block_on(async { + let port = drop_target().await; + let (tx, mut rx) = mpsc::unbounded_channel(); + let mut mux = PortForwardMux::new(tx, "localhost:1".to_owned()); + for id in 1..=(MAX_CHANNELS as i32 * 2) { + mux.handle(open(id, port), true); + assert_eq!(opened(&next_frame(&mut rx).await), (id, true)); + // The task sends `close` on the target's EOF and exits; the + // entry is dead until the next `open` sweeps it. + let ch = next_frame(&mut rx).await; + match &ch.union { + Some(port_forward_channel::Union::Close(c)) => assert_eq!(c.channel_id, id), + other => panic!("expected close, got {:?}", other), + } + tokio::task::yield_now().await; + } + // A task has exited by the time its `close` is read, so one sweep + // must leave nothing behind. + mux.sweep(); + assert_eq!(mux.live_channels(), 0); + }); + } + + #[test] + fn label_counts_distinct_targets_only() { + rt().block_on(async { + let a = echo_target().await; + let b = echo_target().await; + let (tx, mut rx) = mpsc::unbounded_channel(); + let mut mux = PortForwardMux::new(tx, format!("127.0.0.1:{}", a)); + assert_eq!(mux.sweep(), None); + mux.handle(open(1, a), true); + mux.handle(open(2, a), true); + opened(&next_frame(&mut rx).await); + opened(&next_frame(&mut rx).await); + assert_eq!(mux.sweep(), None); + mux.handle(open(3, b), true); + opened(&next_frame(&mut rx).await); + assert_eq!(mux.sweep(), Some(format!("127.0.0.1:{} +1", a))); + mux.handle(close(3), true); + tokio::task::yield_now().await; + assert_eq!(mux.sweep(), Some(format!("127.0.0.1:{}", a))); + }); + } +}