mirror of
https://github.com/rustdesk/rustdesk.git
synced 2026-09-11 15:01:02 +03:00
port_forward_mux: credit-windowed relay halves and channel coordinator
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01EZ49AbZJYfm8NTp5yDPMab
This commit is contained in:
@@ -1,9 +1,16 @@
|
|||||||
use hbb_common::{
|
use hbb_common::{
|
||||||
bytes::Bytes,
|
bytes::{BufMut, Bytes, BytesMut},
|
||||||
|
log,
|
||||||
message_proto::*,
|
message_proto::*,
|
||||||
tokio::{self, sync::Notify},
|
tokio::{
|
||||||
|
self,
|
||||||
|
io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt},
|
||||||
|
sync::{mpsc, watch, Notify},
|
||||||
|
time::Instant,
|
||||||
|
},
|
||||||
|
ResultType,
|
||||||
};
|
};
|
||||||
use std::sync::Mutex;
|
use std::sync::{Arc, Mutex};
|
||||||
|
|
||||||
/// On the wire and fixed forever: what the controller may have in flight on a
|
/// On the wire and fixed forever: what the controller may have in flight on a
|
||||||
/// channel before `opened` brings the peer's window.
|
/// channel before `opened` brings the peer's window.
|
||||||
@@ -172,6 +179,203 @@ pub fn window_update_msg(id: i32, add: u32) -> Message {
|
|||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Where a channel's frames go. The controller keeps two queues so
|
||||||
|
/// `window_update` can bypass bulk data; the controlled side has the
|
||||||
|
/// connection's single ordered `inner.tx`. `open` is not control: it rides
|
||||||
|
/// the ordered queue so it can never arrive after the channel's first `data`.
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub enum FrameSink {
|
||||||
|
Queued {
|
||||||
|
data: mpsc::Sender<Message>,
|
||||||
|
control: mpsc::UnboundedSender<Message>,
|
||||||
|
},
|
||||||
|
Direct(mpsc::UnboundedSender<(Instant, Arc<Message>)>),
|
||||||
|
}
|
||||||
|
|
||||||
|
fn writer_gone(what: &str) -> std::io::Error {
|
||||||
|
std::io::Error::new(std::io::ErrorKind::BrokenPipe, format!("{} gone", what))
|
||||||
|
}
|
||||||
|
|
||||||
|
impl FrameSink {
|
||||||
|
pub async fn send_ordered(&self, msg: Message) -> ResultType<()> {
|
||||||
|
match self {
|
||||||
|
FrameSink::Queued { data, .. } => data
|
||||||
|
.send(msg)
|
||||||
|
.await
|
||||||
|
.map_err(|_| writer_gone("tunnel writer").into()),
|
||||||
|
FrameSink::Direct(tx) => tx
|
||||||
|
.send((Instant::now(), Arc::new(msg)))
|
||||||
|
.map_err(|_| writer_gone("connection writer").into()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn send_control(&self, msg: Message) -> ResultType<()> {
|
||||||
|
match self {
|
||||||
|
FrameSink::Queued { control, .. } => control
|
||||||
|
.send(msg)
|
||||||
|
.map_err(|_| writer_gone("tunnel writer").into()),
|
||||||
|
FrameSink::Direct(tx) => tx
|
||||||
|
.send((Instant::now(), Arc::new(msg)))
|
||||||
|
.map_err(|_| writer_gone("connection writer").into()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn is_closed(&self) -> bool {
|
||||||
|
match self {
|
||||||
|
FrameSink::Queued { data, .. } => data.is_closed(),
|
||||||
|
FrameSink::Direct(tx) => tx.is_closed(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub enum Inbound {
|
||||||
|
Data(Bytes),
|
||||||
|
Close,
|
||||||
|
/// The demultiplexer found the peer over its window; the channel task
|
||||||
|
/// closes and tells the peer, so the frame still leaves in order.
|
||||||
|
Violation,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, PartialEq)]
|
||||||
|
pub enum RelayEnd {
|
||||||
|
LocalEof,
|
||||||
|
PeerClosed,
|
||||||
|
Violation,
|
||||||
|
Cancelled,
|
||||||
|
TunnelGone,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Local socket -> tunnel, under the peer's credit. `prebuf` is simply the
|
||||||
|
/// head of the byte stream.
|
||||||
|
async fn relay_socket_to_tunnel<R: AsyncRead + Unpin>(
|
||||||
|
id: i32,
|
||||||
|
reader: R,
|
||||||
|
prebuf: Vec<u8>,
|
||||||
|
credit: Arc<SendCredit>,
|
||||||
|
sink: FrameSink,
|
||||||
|
mut cancel: watch::Receiver<bool>,
|
||||||
|
) -> RelayEnd {
|
||||||
|
let mut reader = std::io::Cursor::new(prebuf).chain(reader);
|
||||||
|
loop {
|
||||||
|
let allow = tokio::select! {
|
||||||
|
n = credit.take(MAX_FRAME) => n,
|
||||||
|
_ = cancel.changed() => return RelayEnd::Cancelled,
|
||||||
|
};
|
||||||
|
let mut buf = BytesMut::with_capacity(allow);
|
||||||
|
let mut limited = (&mut buf).limit(allow);
|
||||||
|
let got = tokio::select! {
|
||||||
|
r = reader.read_buf(&mut limited) => match r {
|
||||||
|
Ok(n) => n,
|
||||||
|
Err(_) => 0,
|
||||||
|
},
|
||||||
|
_ = cancel.changed() => {
|
||||||
|
credit.add(allow as u32);
|
||||||
|
return RelayEnd::Cancelled;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
let spent = if got == 0 { 0 } else { charge(got) };
|
||||||
|
if (spent as usize) < allow {
|
||||||
|
credit.add(allow as u32 - spent);
|
||||||
|
}
|
||||||
|
if got == 0 {
|
||||||
|
return RelayEnd::LocalEof;
|
||||||
|
}
|
||||||
|
if sink.send_ordered(data_msg(id, buf.freeze())).await.is_err() {
|
||||||
|
return RelayEnd::TunnelGone;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Tunnel -> local socket. `initial` is written before anything from the
|
||||||
|
/// queue (the controlled side's bytes buffered while connecting).
|
||||||
|
async fn relay_tunnel_to_socket<W: AsyncWrite + Unpin>(
|
||||||
|
id: i32,
|
||||||
|
mut writer: W,
|
||||||
|
initial: Vec<Bytes>,
|
||||||
|
mut inbound: mpsc::UnboundedReceiver<Inbound>,
|
||||||
|
window: Arc<Mutex<RecvWindow>>,
|
||||||
|
sink: FrameSink,
|
||||||
|
mut cancel: watch::Receiver<bool>,
|
||||||
|
) -> RelayEnd {
|
||||||
|
let mut pending: std::collections::VecDeque<Bytes> = initial.into();
|
||||||
|
loop {
|
||||||
|
let chunk = match pending.pop_front() {
|
||||||
|
Some(c) => c,
|
||||||
|
None => {
|
||||||
|
let next = tokio::select! {
|
||||||
|
n = inbound.recv() => n,
|
||||||
|
_ = cancel.changed() => return RelayEnd::Cancelled,
|
||||||
|
};
|
||||||
|
match next {
|
||||||
|
Some(Inbound::Data(c)) => c,
|
||||||
|
Some(Inbound::Close) => return RelayEnd::PeerClosed,
|
||||||
|
Some(Inbound::Violation) => return RelayEnd::Violation,
|
||||||
|
None => return RelayEnd::TunnelGone,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
};
|
||||||
|
let written = tokio::select! {
|
||||||
|
r = writer.write_all(&chunk) => r.is_ok(),
|
||||||
|
_ = cancel.changed() => return RelayEnd::Cancelled,
|
||||||
|
};
|
||||||
|
if !written {
|
||||||
|
return RelayEnd::LocalEof;
|
||||||
|
}
|
||||||
|
let update = window.lock().unwrap().drained(chunk.len());
|
||||||
|
if let Some(add) = update {
|
||||||
|
if sink.send_control(window_update_msg(id, add)).is_err() {
|
||||||
|
return RelayEnd::TunnelGone;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Runs both halves as independent tasks; whichever ends first cancels the
|
||||||
|
/// other. Sends `close` once, after the last data, and only when the channel
|
||||||
|
/// ended for a local reason — the peer's own `close` is never echoed.
|
||||||
|
pub async fn run_channel<R, W>(
|
||||||
|
id: i32,
|
||||||
|
reader: R,
|
||||||
|
writer: W,
|
||||||
|
prebuf: Vec<u8>,
|
||||||
|
initial_out: Vec<Bytes>,
|
||||||
|
credit: Arc<SendCredit>,
|
||||||
|
window: Arc<Mutex<RecvWindow>>,
|
||||||
|
inbound: mpsc::UnboundedReceiver<Inbound>,
|
||||||
|
sink: FrameSink,
|
||||||
|
) where
|
||||||
|
R: AsyncRead + Unpin + Send + 'static,
|
||||||
|
W: AsyncWrite + Unpin + Send + 'static,
|
||||||
|
{
|
||||||
|
let (cancel_tx, cancel_rx) = watch::channel(false);
|
||||||
|
let mut to_tunnel = tokio::spawn(relay_socket_to_tunnel(
|
||||||
|
id, reader, prebuf, credit, sink.clone(), cancel_rx.clone(),
|
||||||
|
));
|
||||||
|
let mut to_socket = tokio::spawn(relay_tunnel_to_socket(
|
||||||
|
id, writer, initial_out, inbound, window, sink.clone(), cancel_rx,
|
||||||
|
));
|
||||||
|
let (first, second) = tokio::select! {
|
||||||
|
r = &mut to_tunnel => {
|
||||||
|
let _ = cancel_tx.send(true);
|
||||||
|
(r.unwrap_or(RelayEnd::Cancelled), to_socket.await.unwrap_or(RelayEnd::Cancelled))
|
||||||
|
}
|
||||||
|
r = &mut to_socket => {
|
||||||
|
let _ = cancel_tx.send(true);
|
||||||
|
(r.unwrap_or(RelayEnd::Cancelled), to_tunnel.await.unwrap_or(RelayEnd::Cancelled))
|
||||||
|
}
|
||||||
|
};
|
||||||
|
let peer_closed = first == RelayEnd::PeerClosed || second == RelayEnd::PeerClosed;
|
||||||
|
let tunnel_gone = first == RelayEnd::TunnelGone || second == RelayEnd::TunnelGone;
|
||||||
|
let local_reason = matches!(first, RelayEnd::LocalEof | RelayEnd::Violation)
|
||||||
|
|| matches!(second, RelayEnd::LocalEof | RelayEnd::Violation);
|
||||||
|
if !peer_closed && !tunnel_gone && local_reason {
|
||||||
|
if let Err(e) = sink.send_ordered(close_msg(id)).await {
|
||||||
|
log::debug!("port forward channel {} close not sent: {}", id, e);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
log::debug!("port forward channel {} ended: {:?} / {:?}", id, first, second);
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
@@ -308,4 +512,181 @@ mod tests {
|
|||||||
other => panic!("unexpected {:?}", other),
|
other => panic!("unexpected {:?}", other),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
use hbb_common::message_proto::{message, port_forward_channel};
|
||||||
|
use hbb_common::tokio::{self, io::AsyncReadExt, io::AsyncWriteExt, sync::mpsc};
|
||||||
|
use std::sync::{Arc, Mutex};
|
||||||
|
|
||||||
|
struct Harness {
|
||||||
|
data_rx: mpsc::Receiver<Message>,
|
||||||
|
control_rx: mpsc::UnboundedReceiver<Message>,
|
||||||
|
inbound_tx: mpsc::UnboundedSender<Inbound>,
|
||||||
|
credit: Arc<SendCredit>,
|
||||||
|
window: Arc<Mutex<RecvWindow>>,
|
||||||
|
local: tokio::io::DuplexStream,
|
||||||
|
task: tokio::task::JoinHandle<()>,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A channel whose "local socket" is one end of a duplex pipe and whose
|
||||||
|
/// "tunnel" is a pair of queues the test reads directly.
|
||||||
|
fn harness(id: i32, prebuf: Vec<u8>, initial_out: Vec<Bytes>) -> Harness {
|
||||||
|
let (data_tx, data_rx) = mpsc::channel(DATA_QUEUE_FRAMES);
|
||||||
|
let (control_tx, control_rx) = mpsc::unbounded_channel();
|
||||||
|
let (inbound_tx, inbound_rx) = mpsc::unbounded_channel();
|
||||||
|
let (local, remote) = tokio::io::duplex(1 << 20);
|
||||||
|
let (r, w) = tokio::io::split(remote);
|
||||||
|
let credit = Arc::new(SendCredit::new(INITIAL_WINDOW));
|
||||||
|
let window = Arc::new(Mutex::new(RecvWindow::new(CHANNEL_WINDOW)));
|
||||||
|
let sink = FrameSink::Queued { data: data_tx, control: control_tx };
|
||||||
|
let task = tokio::spawn(run_channel(
|
||||||
|
id, r, w, prebuf, initial_out, credit.clone(), window.clone(), inbound_rx, sink,
|
||||||
|
));
|
||||||
|
Harness { data_rx, control_rx, inbound_tx, credit, window, local, task }
|
||||||
|
}
|
||||||
|
|
||||||
|
// `PortForwardData.data` is generated as `bytes::Bytes` (hbb_common builds
|
||||||
|
// rust-protobuf with the bytes feature), so it converts with `to_vec()`, not
|
||||||
|
// `clone()`, and needs no wrapping when it becomes an `Inbound::Data`.
|
||||||
|
fn frame_kind(m: &Message) -> (&'static str, i32, Vec<u8>) {
|
||||||
|
match &m.union {
|
||||||
|
Some(message::Union::PortForwardChannel(ch)) => match &ch.union {
|
||||||
|
Some(port_forward_channel::Union::Data(d)) => ("data", d.channel_id, d.data.to_vec()),
|
||||||
|
Some(port_forward_channel::Union::Close(c)) => ("close", c.channel_id, vec![]),
|
||||||
|
Some(port_forward_channel::Union::WindowUpdate(u)) => {
|
||||||
|
("window_update", u.channel_id, u.add.to_le_bytes().to_vec())
|
||||||
|
}
|
||||||
|
Some(port_forward_channel::Union::Open(o)) => ("open", o.channel_id, vec![]),
|
||||||
|
Some(port_forward_channel::Union::Opened(o)) => ("opened", o.channel_id, vec![]),
|
||||||
|
None => ("none", 0, vec![]),
|
||||||
|
// `port_forward_channel::Union` is `#[non_exhaustive]` in the
|
||||||
|
// generated protobuf code, so it needs a catch-all here even
|
||||||
|
// though every current variant is already matched above.
|
||||||
|
_ => ("other", 0, vec![]),
|
||||||
|
},
|
||||||
|
_ => ("other", 0, vec![]),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn local_bytes_become_data_frames_capped_at_max_frame() {
|
||||||
|
rt().block_on(async {
|
||||||
|
let mut h = harness(1, vec![], vec![]);
|
||||||
|
let payload = vec![7u8; MAX_FRAME + 10];
|
||||||
|
h.local.write_all(&payload).await.unwrap();
|
||||||
|
// INITIAL_WINDOW equals MAX_FRAME, so the first frame exhausts
|
||||||
|
// it exactly; grant one minimum charge for the 10-byte tail.
|
||||||
|
h.credit.add(MIN_FRAME_CHARGE);
|
||||||
|
let mut got = Vec::new();
|
||||||
|
while got.len() < payload.len() {
|
||||||
|
let m = h.data_rx.recv().await.unwrap();
|
||||||
|
let (kind, id, bytes) = frame_kind(&m);
|
||||||
|
assert_eq!((kind, id), ("data", 1));
|
||||||
|
assert!(bytes.len() <= MAX_FRAME);
|
||||||
|
got.extend(bytes);
|
||||||
|
}
|
||||||
|
assert_eq!(got, payload);
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn prebuf_is_the_head_of_the_send_stream() {
|
||||||
|
rt().block_on(async {
|
||||||
|
let mut h = harness(2, b"head".to_vec(), vec![]);
|
||||||
|
h.local.write_all(b"tail").await.unwrap();
|
||||||
|
let mut got = Vec::new();
|
||||||
|
while got.len() < 8 {
|
||||||
|
let m = h.data_rx.recv().await.unwrap();
|
||||||
|
got.extend(frame_kind(&m).2);
|
||||||
|
}
|
||||||
|
assert_eq!(got, b"headtail".to_vec());
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn send_side_stops_at_credit_and_resumes_on_add() {
|
||||||
|
rt().block_on(async {
|
||||||
|
let mut h = harness(3, vec![], vec![]);
|
||||||
|
let payload = vec![1u8; INITIAL_WINDOW as usize + 5];
|
||||||
|
h.local.write_all(&payload).await.unwrap();
|
||||||
|
let mut got = 0usize;
|
||||||
|
while got < INITIAL_WINDOW as usize {
|
||||||
|
got += frame_kind(&h.data_rx.recv().await.unwrap()).2.len();
|
||||||
|
}
|
||||||
|
assert_eq!(got, INITIAL_WINDOW as usize);
|
||||||
|
assert!(tokio::time::timeout(
|
||||||
|
std::time::Duration::from_millis(50),
|
||||||
|
h.data_rx.recv()
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.is_err());
|
||||||
|
// One minimum charge is enough to send the 5-byte tail.
|
||||||
|
h.credit.add(MIN_FRAME_CHARGE);
|
||||||
|
assert_eq!(frame_kind(&h.data_rx.recv().await.unwrap()).2.len(), 5);
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn inbound_data_is_written_and_window_update_follows_threshold() {
|
||||||
|
rt().block_on(async {
|
||||||
|
let mut h = harness(4, vec![], vec![Bytes::from_static(b"first")]);
|
||||||
|
let mut buf = [0u8; 5];
|
||||||
|
h.local.read_exact(&mut buf).await.unwrap();
|
||||||
|
assert_eq!(&buf, b"first");
|
||||||
|
let chunk = Bytes::from(vec![9u8; UPDATE_THRESHOLD as usize]);
|
||||||
|
assert!(h.window.lock().unwrap().accept(chunk.len()));
|
||||||
|
h.inbound_tx.send(Inbound::Data(chunk.clone())).unwrap();
|
||||||
|
let mut sink = vec![0u8; chunk.len()];
|
||||||
|
h.local.read_exact(&mut sink).await.unwrap();
|
||||||
|
let m = h.control_rx.recv().await.unwrap();
|
||||||
|
let (kind, id, add) = frame_kind(&m);
|
||||||
|
assert_eq!((kind, id), ("window_update", 4));
|
||||||
|
let add = u32::from_le_bytes([add[0], add[1], add[2], add[3]]);
|
||||||
|
// The 5-byte `initial` chunk drained a whole minimum charge.
|
||||||
|
assert_eq!(add, UPDATE_THRESHOLD + MIN_FRAME_CHARGE);
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn local_eof_sends_close_exactly_once_after_the_data() {
|
||||||
|
rt().block_on(async {
|
||||||
|
let mut h = harness(5, vec![], vec![]);
|
||||||
|
h.local.write_all(b"bye").await.unwrap();
|
||||||
|
drop(h.local);
|
||||||
|
assert_eq!(frame_kind(&h.data_rx.recv().await.unwrap()).0, "data");
|
||||||
|
assert_eq!(frame_kind(&h.data_rx.recv().await.unwrap()), ("close", 5, vec![]));
|
||||||
|
h.task.await.unwrap();
|
||||||
|
assert!(h.data_rx.try_recv().is_err());
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn peer_close_ends_the_channel_without_echoing_close() {
|
||||||
|
rt().block_on(async {
|
||||||
|
let mut h = harness(6, vec![], vec![]);
|
||||||
|
h.inbound_tx.send(Inbound::Close).unwrap();
|
||||||
|
h.task.await.unwrap();
|
||||||
|
assert!(h.data_rx.try_recv().is_err());
|
||||||
|
assert!(h.control_rx.try_recv().is_err());
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn violation_signalled_by_the_demux_sends_close() {
|
||||||
|
rt().block_on(async {
|
||||||
|
let mut h = harness(7, vec![], vec![]);
|
||||||
|
h.inbound_tx.send(Inbound::Violation).unwrap();
|
||||||
|
assert_eq!(frame_kind(&h.data_rx.recv().await.unwrap()), ("close", 7, vec![]));
|
||||||
|
h.task.await.unwrap();
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn dropped_tunnel_ends_the_channel_silently() {
|
||||||
|
rt().block_on(async {
|
||||||
|
let mut h = harness(8, vec![], vec![]);
|
||||||
|
drop(h.inbound_tx);
|
||||||
|
h.task.await.unwrap();
|
||||||
|
assert!(h.data_rx.try_recv().is_err());
|
||||||
|
});
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user