fix: 修正 RustDesk UDP 注册状态判定

- 移除无法证明服务器可达的 connected 状态
- 引入注册跟踪器统一管理响应、超时与退避
- 连续注册超时后正确降级为未连接
 - 分离本地服务状态与 HBBS 注册状态
- 补充注册状态机单元测试
This commit is contained in:
mofeng-git
2026-07-30 19:43:28 +08:00
parent 9fb23476ac
commit 8b09824e04
2 changed files with 237 additions and 51 deletions

View File

@@ -276,7 +276,6 @@ impl RustDeskService {
}, },
)); ));
let status = self.status.clone();
let handle = tokio::spawn(async move { let handle = tokio::spawn(async move {
loop { loop {
match mediator.start().await { match mediator.start().await {
@@ -286,9 +285,7 @@ impl RustDeskService {
} }
Err(e) => { Err(e) => {
error!("Rendezvous mediator error: {}", e); error!("Rendezvous mediator error: {}", e);
*status.write() = ServiceStatus::Error(e.to_string());
tokio::time::sleep(std::time::Duration::from_secs(5)).await; tokio::time::sleep(std::time::Duration::from_secs(5)).await;
*status.write() = ServiceStatus::Starting;
} }
} }
} }

View File

@@ -20,11 +20,13 @@ use super::protocol::{
rendezvous_message, NatType, RendezvousMessage, rendezvous_message, NatType, RendezvousMessage,
}; };
const REG_INTERVAL_MS: u64 = 12_000; const REG_INTERVAL: Duration = Duration::from_secs(12);
const MIN_REG_TIMEOUT_MS: u64 = 3_000; const MIN_REG_TIMEOUT: Duration = Duration::from_secs(3);
const MAX_REG_TIMEOUT_MS: u64 = 30_000; const MAX_REG_TIMEOUT: Duration = Duration::from_secs(30);
const OFFLINE_AFTER_TIMEOUTS: u32 = 4;
const TIMER_INTERVAL_MS: u64 = 300; const TIMER_INTERVAL_MS: u64 = 300;
@@ -32,7 +34,6 @@ const TIMER_INTERVAL_MS: u64 = 300;
pub enum RendezvousStatus { pub enum RendezvousStatus {
Disconnected, Disconnected,
Connecting, Connecting,
Connected,
Registered, Registered,
Error(String), Error(String),
} }
@@ -42,13 +43,86 @@ impl std::fmt::Display for RendezvousStatus {
match self { match self {
Self::Disconnected => write!(f, "disconnected"), Self::Disconnected => write!(f, "disconnected"),
Self::Connecting => write!(f, "connecting"), Self::Connecting => write!(f, "connecting"),
Self::Connected => write!(f, "connected"),
Self::Registered => write!(f, "registered"), Self::Registered => write!(f, "registered"),
Self::Error(e) => write!(f, "error: {}", e), Self::Error(e) => write!(f, "error: {}", e),
} }
} }
} }
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum RegistrationDecision {
Wait,
Send,
Retry { consecutive_timeouts: u32 },
}
/// Tracks the request/response lifecycle for HBBS registration.
///
/// UDP `connect` only selects a peer; it does not prove reachability. Registration
/// health is therefore derived exclusively from acknowledged registration requests.
#[derive(Debug)]
struct RegistrationTracker {
last_sent: Option<Instant>,
last_response: Option<Instant>,
response_timeout: Duration,
consecutive_timeouts: u32,
}
impl RegistrationTracker {
fn new() -> Self {
Self {
last_sent: None,
last_response: None,
response_timeout: MIN_REG_TIMEOUT,
consecutive_timeouts: 0,
}
}
fn poll(&mut self, now: Instant) -> RegistrationDecision {
if let Some(sent_at) = self.last_sent {
if now.saturating_duration_since(sent_at) < self.response_timeout {
return RegistrationDecision::Wait;
}
self.consecutive_timeouts = self.consecutive_timeouts.saturating_add(1);
self.response_timeout = (self.response_timeout + MIN_REG_TIMEOUT).min(MAX_REG_TIMEOUT);
return RegistrationDecision::Retry {
consecutive_timeouts: self.consecutive_timeouts,
};
}
let registration_expired = self
.last_response
.map(|response_at| now.saturating_duration_since(response_at) >= REG_INTERVAL)
.unwrap_or(true);
if registration_expired {
RegistrationDecision::Send
} else {
RegistrationDecision::Wait
}
}
fn mark_sent(&mut self, now: Instant) {
self.last_sent = Some(now);
}
fn mark_response(&mut self, now: Instant) {
self.last_sent = None;
self.last_response = Some(now);
self.response_timeout = MIN_REG_TIMEOUT;
self.consecutive_timeouts = 0;
}
fn status_after_timeout(consecutive_timeouts: u32) -> RendezvousStatus {
if consecutive_timeouts >= OFFLINE_AFTER_TIMEOUTS {
RendezvousStatus::Disconnected
} else {
RendezvousStatus::Connecting
}
}
}
pub type RelayCallback = Arc<dyn Fn(String, String, String, Vec<u8>, String) + Send + Sync>; pub type RelayCallback = Arc<dyn Fn(String, String, String, Vec<u8>, String) + Send + Sync>;
pub type PunchCallback = pub type PunchCallback =
@@ -147,6 +221,14 @@ impl RendezvousMediator {
self.status.read().clone() self.status.read().clone()
} }
fn set_status(&self, next: RendezvousStatus) {
let mut current = self.status.write();
if *current != next {
info!("Rendezvous status changed: {} -> {}", *current, next);
*current = next;
}
}
pub fn update_config(&self, config: RustDeskConfig) { pub fn update_config(&self, config: RustDeskConfig) {
*self.config.write() = config; *self.config.write() = config;
self.increment_serial(); self.increment_serial();
@@ -209,11 +291,21 @@ impl RendezvousMediator {
"Rendezvous mediator not starting: enabled={}, server='{}'", "Rendezvous mediator not starting: enabled={}, server='{}'",
config.enabled, effective_server config.enabled, effective_server
); );
self.set_status(RendezvousStatus::Disconnected);
return Ok(()); return Ok(());
} }
*self.status.write() = RendezvousStatus::Connecting; self.set_status(RendezvousStatus::Connecting);
let result = self.run(config).await;
match &result {
Ok(()) => self.set_status(RendezvousStatus::Disconnected),
Err(err) => self.set_status(RendezvousStatus::Error(err.to_string())),
}
result
}
async fn run(&self, config: RustDeskConfig) -> anyhow::Result<()> {
let addr = config.rendezvous_addr(); let addr = config.rendezvous_addr();
info!( info!(
"Starting rendezvous mediator for {} to {}", "Starting rendezvous mediator for {} to {}",
@@ -233,8 +325,7 @@ impl RendezvousMediator {
let socket = UdpSocket::from_std(std_socket)?; let socket = UdpSocket::from_std(std_socket)?;
socket.connect(server_addr).await?; socket.connect(server_addr).await?;
info!("Connected to rendezvous server at {}", server_addr); info!("RustDesk UDP transport ready for {}", server_addr);
*self.status.write() = RendezvousStatus::Connected;
self.registration_loop(socket).await self.registration_loop(socket).await
} }
@@ -242,10 +333,7 @@ impl RendezvousMediator {
async fn registration_loop(&self, socket: UdpSocket) -> anyhow::Result<()> { async fn registration_loop(&self, socket: UdpSocket) -> anyhow::Result<()> {
let mut timer = interval(Duration::from_millis(TIMER_INTERVAL_MS)); let mut timer = interval(Duration::from_millis(TIMER_INTERVAL_MS));
let mut recv_buf = vec![0u8; 65535]; let mut recv_buf = vec![0u8; 65535];
let mut last_register_sent: Option<Instant> = None; let mut registration = RegistrationTracker::new();
let mut last_register_resp: Option<Instant> = None;
let mut reg_timeout = MIN_REG_TIMEOUT_MS;
let mut fails = 0;
let mut shutdown_rx = self.shutdown_tx.subscribe(); let mut shutdown_rx = self.shutdown_tx.subscribe();
loop { loop {
@@ -254,39 +342,36 @@ impl RendezvousMediator {
match result { match result {
Ok(len) => { Ok(len) => {
if let Ok(msg) = decode_rendezvous_message(&recv_buf[..len]) { if let Ok(msg) = decode_rendezvous_message(&recv_buf[..len]) {
self.handle_response(&socket, msg, &mut last_register_resp, &mut fails, &mut reg_timeout).await?; self.handle_response(&socket, msg, &mut registration).await?;
} else { } else {
debug!("Failed to decode rendezvous message"); debug!("Failed to decode rendezvous message");
} }
} }
Err(e) => { Err(e) => {
error!("Failed to receive from socket: {}", e); return Err(anyhow::anyhow!("Failed to receive from socket: {}", e));
*self.status.write() = RendezvousStatus::Error(e.to_string());
break;
} }
} }
} }
_ = timer.tick() => { _ = timer.tick() => {
let now = Instant::now(); let now = Instant::now();
let expired = last_register_resp match registration.poll(now) {
.map(|x| x.elapsed().as_millis() as u64 >= REG_INTERVAL_MS) RegistrationDecision::Wait => {}
.unwrap_or(true); RegistrationDecision::Send => {
let timeout = last_register_sent
.map(|x| x.elapsed().as_millis() as u64 >= reg_timeout)
.unwrap_or(false);
if timeout && reg_timeout < MAX_REG_TIMEOUT_MS {
reg_timeout += MIN_REG_TIMEOUT_MS;
fails += 1;
if fails >= 4 {
warn!("Registration timeout, {} consecutive failures", fails);
}
}
if timeout || (last_register_sent.is_none() && expired) {
self.send_register(&socket).await?; self.send_register(&socket).await?;
last_register_sent = Some(now); registration.mark_sent(now);
}
RegistrationDecision::Retry { consecutive_timeouts } => {
let next_status =
RegistrationTracker::status_after_timeout(consecutive_timeouts);
self.set_status(next_status);
warn!(
"RustDesk registration timed out ({} consecutive timeouts)",
consecutive_timeouts
);
self.send_register(&socket).await?;
registration.mark_sent(now);
}
} }
} }
@@ -297,7 +382,6 @@ impl RendezvousMediator {
} }
} }
*self.status.write() = RendezvousStatus::Disconnected;
Ok(()) Ok(())
} }
@@ -384,48 +468,49 @@ impl RendezvousMediator {
&self, &self,
socket: &UdpSocket, socket: &UdpSocket,
msg: RendezvousMessage, msg: RendezvousMessage,
last_resp: &mut Option<Instant>, registration: &mut RegistrationTracker,
fails: &mut i32,
reg_timeout: &mut u64,
) -> anyhow::Result<()> { ) -> anyhow::Result<()> {
*last_resp = Some(Instant::now());
*fails = 0;
*reg_timeout = MIN_REG_TIMEOUT_MS;
match msg.union { match msg.union {
Some(rendezvous_message::Union::RegisterPeerResponse(rpr)) => { Some(rendezvous_message::Union::RegisterPeerResponse(rpr)) => {
registration.mark_response(Instant::now());
if rpr.request_pk { if rpr.request_pk {
info!("Server requested public key registration"); info!("Server requested public key registration");
*self.key_confirmed.write() = false; *self.key_confirmed.write() = false;
self.set_status(RendezvousStatus::Connecting);
self.send_register_pk(socket).await?; self.send_register_pk(socket).await?;
registration.mark_sent(Instant::now());
} else {
self.set_status(RendezvousStatus::Registered);
} }
*self.status.write() = RendezvousStatus::Registered;
} }
Some(rendezvous_message::Union::RegisterPkResponse(rpr)) => { Some(rendezvous_message::Union::RegisterPkResponse(rpr)) => {
registration.mark_response(Instant::now());
info!("Received RegisterPkResponse: result={:?}", rpr.result); info!("Received RegisterPkResponse: result={:?}", rpr.result);
match rpr.result.value() { match rpr.result.value() {
0 => { 0 => {
info!("✓ Public key registered successfully with server"); info!("✓ Public key registered successfully with server");
*self.key_confirmed.write() = true; *self.key_confirmed.write() = true;
self.increment_serial(); self.increment_serial();
*self.status.write() = RendezvousStatus::Registered; self.set_status(RendezvousStatus::Registered);
} }
2 => { 2 => {
warn!("UUID mismatch, need to re-register"); warn!("UUID mismatch, need to re-register");
*self.key_confirmed.write() = false; *self.key_confirmed.write() = false;
self.set_status(RendezvousStatus::Connecting);
} }
3 => { 3 => {
error!("Device ID already exists on server"); error!("Device ID already exists on server");
*self.status.write() = self.set_status(RendezvousStatus::Error(
RendezvousStatus::Error("Device ID already exists".to_string()); "Device ID already exists".to_string(),
));
} }
4 => { 4 => {
warn!("Registration too frequent"); warn!("Registration too frequent");
self.set_status(RendezvousStatus::Connecting);
} }
5 => { 5 => {
error!("Invalid device ID format"); error!("Invalid device ID format");
*self.status.write() = self.set_status(RendezvousStatus::Error("Invalid ID format".to_string()));
RendezvousStatus::Error("Invalid ID format".to_string());
} }
_ => { _ => {
error!("Unknown RegisterPkResponse result: {:?}", rpr.result); error!("Unknown RegisterPkResponse result: {:?}", rpr.result);
@@ -797,7 +882,111 @@ fn get_local_addresses() -> Vec<std::net::IpAddr> {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::{normalize_relay_server, select_relay_server}; use std::time::{Duration, Instant};
use super::{
normalize_relay_server, select_relay_server, RegistrationDecision, RegistrationTracker,
RendezvousStatus, REG_INTERVAL,
};
#[test]
fn registration_tracker_requires_an_acknowledged_response() {
let started_at = Instant::now();
let mut tracker = RegistrationTracker::new();
assert_eq!(tracker.poll(started_at), RegistrationDecision::Send);
tracker.mark_sent(started_at);
assert_eq!(
tracker.poll(started_at + Duration::from_secs(2)),
RegistrationDecision::Wait
);
assert_eq!(
tracker.poll(started_at + Duration::from_secs(3)),
RegistrationDecision::Retry {
consecutive_timeouts: 1
}
);
assert_eq!(
RegistrationTracker::status_after_timeout(1),
RendezvousStatus::Connecting
);
}
#[test]
fn registration_tracker_marks_four_timeouts_offline() {
let mut now = Instant::now();
let mut tracker = RegistrationTracker::new();
assert_eq!(tracker.poll(now), RegistrationDecision::Send);
tracker.mark_sent(now);
for (failure, timeout) in [(1, 3), (2, 6), (3, 9), (4, 12)] {
now += Duration::from_secs(timeout);
assert_eq!(
tracker.poll(now),
RegistrationDecision::Retry {
consecutive_timeouts: failure
}
);
tracker.mark_sent(now);
}
assert_eq!(
RegistrationTracker::status_after_timeout(4),
RendezvousStatus::Disconnected
);
}
#[test]
fn registration_response_clears_in_flight_retry_state() {
let started_at = Instant::now();
let mut tracker = RegistrationTracker::new();
tracker.mark_sent(started_at);
let retry_at = started_at + Duration::from_secs(3);
assert_eq!(
tracker.poll(retry_at),
RegistrationDecision::Retry {
consecutive_timeouts: 1
}
);
tracker.mark_sent(retry_at);
let response_at = retry_at + Duration::from_millis(100);
tracker.mark_response(response_at);
assert_eq!(tracker.last_sent, None);
assert_eq!(tracker.consecutive_timeouts, 0);
assert_eq!(
tracker.poll(response_at + REG_INTERVAL - Duration::from_millis(1)),
RegistrationDecision::Wait
);
assert_eq!(
tracker.poll(response_at + REG_INTERVAL),
RegistrationDecision::Send
);
}
#[test]
fn registration_timeout_count_continues_after_backoff_reaches_its_cap() {
let mut now = Instant::now();
let mut tracker = RegistrationTracker::new();
tracker.mark_sent(now);
for failure in 1..=12 {
now += tracker.response_timeout;
assert_eq!(
tracker.poll(now),
RegistrationDecision::Retry {
consecutive_timeouts: failure
}
);
tracker.mark_sent(now);
}
assert_eq!(tracker.response_timeout, Duration::from_secs(30));
assert_eq!(tracker.consecutive_timeouts, 12);
}
#[test] #[test]
fn test_normalize_relay_server() { fn test_normalize_relay_server() {