Skip to content
95 changes: 88 additions & 7 deletions crates/rds-desktop/src/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,45 @@ const FRAME_STREAM_TIMEOUT: std::time::Duration = std::time::Duration::from_secs
const FRAME_READ_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(3);
const KEYFRAME_READ_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(8);

fn reorder_wait(rtt_ms: u64) -> std::time::Duration {
// A missing stream tag may itself be in transit. Two observed round trips
// tolerate WAN reordering; the cap keeps a truly missing reference bounded.
let ms = if rtt_ms == u64::MAX {
250
} else {
rtt_ms.saturating_mul(2).clamp(100, 1000)
};
std::time::Duration::from_millis(ms)
}

#[derive(Default)]
struct HeartbeatProbes {
pending: std::collections::VecDeque<(u64, u64, std::time::Instant)>,
}

impl HeartbeatProbes {
fn sent(&mut self, seq: u64, ts_ms: u64) {
// Caller timestamps are opaque correlation values. In managed mode
// they belong to the viewer, whose clock survives session reconnects.
self.pending.retain(|(s, t, _)| (*s, *t) != (seq, ts_ms));
if self.pending.len() >= 64 {
self.pending.pop_front();
}
self.pending
.push_back((seq, ts_ms, std::time::Instant::now()));
}

fn echoed(&mut self, seq: u64, ts_ms: u64) -> Option<std::time::Duration> {
let index = self
.pending
.iter()
.position(|(s, t, _)| (*s, *t) == (seq, ts_ms))?;
self.pending
.remove(index)
.map(|(_, _, sent)| sent.elapsed())
}
}

/// Decode state owned by one session. Codec reference frames are chain
/// state: a decoder shared across sessions would cross-contaminate
/// streams, so it lives here and is dropped with the session.
Expand Down Expand Up @@ -365,8 +404,13 @@ impl DesktopSession {
// Construct ownership before spawning. Dropping the supervisor aborts
// its JoinSet, which in turn drops streams and the frame-reader JoinSet.
let mut tasks = JoinSet::new();
let probes = Arc::new(tokio::sync::Mutex::new(HeartbeatProbes::default()));
let sending_probes = probes.clone();
tasks.spawn(async move {
while let Some(msg) = ctrl_rx.recv().await {
if let DesktopControl::Heartbeat { seq, ts_ms } = &msg {
sending_probes.lock().await.sent(*seq, *ts_ms);
}
if !matches!(
tokio::time::timeout(FRAME_STREAM_TIMEOUT, write_frame(&mut send.0, &msg))
.await,
Expand All @@ -378,15 +422,13 @@ impl DesktopSession {
});

let rtt_marker = control_rtt_ms.clone();
let event_clock = clock.clone();
tasks.spawn(async move {
loop {
match read_frame::<_, DesktopEvent>(&mut recv).await {
Ok(ev @ DesktopEvent::Heartbeat { ts_ms, .. }) => {
rtt_marker.store(
event_clock.now_ms().saturating_sub(ts_ms),
Ordering::Relaxed,
);
Ok(ev @ DesktopEvent::Heartbeat { seq, ts_ms }) => {
if let Some(rtt) = probes.lock().await.echoed(seq, ts_ms) {
rtt_marker.store(rtt.as_millis() as u64, Ordering::Relaxed);
}
events_tx.send(ev);
}
Ok(ev) => {
Expand All @@ -406,6 +448,7 @@ impl DesktopSession {
ctrl: ctrl_tx.clone(),
clock: clock.clone(),
receiving: receiving.clone(),
control_rtt_ms: control_rtt_ms.clone(),
},
));
let task = tokio::spawn(async move {
Expand Down Expand Up @@ -642,6 +685,7 @@ struct ReceiveContext {
ctrl: mpsc::Sender<DesktopControl>,
clock: SessionClock,
receiving: Arc<AtomicUsize>,
control_rtt_ms: Arc<AtomicU64>,
}

enum FrameRead {
Expand Down Expand Up @@ -726,8 +770,13 @@ async fn receive_frames(mut uni: rds_net::UniStreams, ctx: ReceiveContext) {
tokio::select! {
biased;
_ = repair.tick() => {
if ordered.expire(std::time::Duration::from_millis(100)) {
// An admitted reference has its own bounded reader deadline.
// Before its tag arrives, allow a measured WAN reorder window
// rather than generating an IDR at the old fixed 100 ms.
let wait = reorder_wait(ctx.control_rtt_ms.load(Ordering::Relaxed));
if readers.is_empty() && ordered.expire(wait) {
gaps+=1;
tracing::warn!(expected_seq=ctx.next_seq.load(Ordering::Relaxed),reorder_budget_ms=wait.as_millis(),"desktop reference gap expired");
delivery.invalidate();
delivery.request_idr(&ctx.ctrl,ctx.clock.now_ms());
}
Expand Down Expand Up @@ -867,6 +916,38 @@ pub async fn run_desktop_client(
mod tests {
use super::*;

#[test]
fn heartbeat_measurements_are_bounded_and_require_exact_correlation() {
let mut probes = HeartbeatProbes::default();
for seq in 0..128 {
probes.sent(seq, 1_000_000 + seq);
}
assert_eq!(probes.pending.len(), 64);
assert!(probes.echoed(0, 1_000_000).is_none());
assert!(probes.echoed(127, 0).is_none());
assert!(probes.echoed(126, 1_000_127).is_none());
assert!(probes.echoed(127, 1_000_127).is_some());
assert!(probes.echoed(127, 1_000_127).is_none());
probes.sent(64, 1_000_064);
assert_eq!(
probes
.pending
.iter()
.filter(|(s, t, _)| (*s, *t) == (64, 1_000_064))
.count(),
1
);
}

#[test]
fn reference_wait_uses_rtt_and_caps_untrusted_or_missing_measurements() {
assert_eq!(reorder_wait(0).as_millis(), 100);
assert_eq!(reorder_wait(200).as_millis(), 400);
assert_eq!(reorder_wait(50_000).as_millis(), 1000);
assert_eq!(reorder_wait(u64::MAX - 1).as_millis(), 1000);
assert_eq!(reorder_wait(u64::MAX).as_millis(), 250);
}

#[test]
fn exhausted_control_sequences_never_wrap() {
let sequence = AtomicU64::new(u64::MAX - 1);
Expand Down
Loading
Loading