feat(rustdesk): 强化会话管理并优化流媒体传输

This commit is contained in:
mofeng-git
2026-08-30 08:18:46 +08:00
parent e678ec394d
commit 95a1fdf42d
11 changed files with 654 additions and 233 deletions

View File

@@ -22,6 +22,9 @@ use super::ConfigApplyOptions;
pub struct RustDeskRuntimeStatus {
pub service_status: String,
pub rendezvous_status: Option<String>,
pub connection_count: usize,
pub listening: bool,
pub listen_port: Option<u16>,
}
pub struct RemoteAccessCoordinator {
@@ -145,10 +148,16 @@ impl RemoteAccessCoordinator {
Some(service) => RustDeskRuntimeStatus {
service_status: service.status().to_string(),
rendezvous_status: service.rendezvous_status().map(|status| status.to_string()),
connection_count: service.connection_count(),
listening: service.is_listening(),
listen_port: service.is_listening().then(|| service.listen_port()),
},
None => RustDeskRuntimeStatus {
service_status: "not_initialized".to_string(),
rendezvous_status: None,
connection_count: 0,
listening: false,
listen_port: None,
},
}
}

View File

@@ -1,7 +1,7 @@
//! Variable-length TCP framing (RustDesk wire format).
use bytes::{Buf, BufMut, Bytes, BytesMut};
use std::io;
use std::io::{self, IoSlice};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
const MAX_PACKET_LENGTH: usize = 0x3FFFFFFF;
@@ -53,6 +53,18 @@ fn decode_header(first_byte: u8, header_bytes: &[u8]) -> (usize, usize) {
}
pub async fn read_frame<R: AsyncRead + Unpin>(reader: &mut R) -> io::Result<BytesMut> {
read_frame_with_limit(reader, MAX_PACKET_LENGTH).await
}
/// Read one framed message while enforcing a caller-selected allocation limit.
///
/// Network-facing protocol stages should use a substantially smaller limit than
/// the wire format's theoretical maximum so an untrusted peer cannot force a
/// huge allocation by sending only a length header.
pub async fn read_frame_with_limit<R: AsyncRead + Unpin>(
reader: &mut R,
max_packet_length: usize,
) -> io::Result<BytesMut> {
let mut first_byte = [0u8; 1];
reader.read_exact(&mut first_byte).await?;
@@ -65,10 +77,10 @@ pub async fn read_frame<R: AsyncRead + Unpin>(reader: &mut R) -> io::Result<Byte
let (_, msg_len) = decode_header(first_byte[0], &header_rest);
if msg_len > MAX_PACKET_LENGTH {
if msg_len > max_packet_length {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"Message too large",
format!("Message too large: {msg_len} bytes exceeds {max_packet_length}-byte limit"),
));
}
@@ -86,6 +98,53 @@ pub async fn write_frame<W: AsyncWrite + Unpin>(writer: &mut W, data: &[u8]) ->
Ok(())
}
/// Write a frame without copying its payload into a second contiguous buffer.
/// TCP writers normally send the small header and payload in one vectored write.
pub async fn write_frame_vectored<W: AsyncWrite + Unpin>(
writer: &mut W,
data: &[u8],
) -> io::Result<()> {
let len = data.len();
let mut header = [0u8; 4];
let header_len = if len <= 0x3F {
header[0] = (len << 2) as u8;
1
} else if len <= 0x3FFF {
header[..2].copy_from_slice(&(((len << 2) as u16) | 0x1).to_le_bytes());
2
} else if len <= 0x3FFFFF {
let value = ((len << 2) as u32) | 0x2;
header[0] = (value & 0xFF) as u8;
header[1] = ((value >> 8) & 0xFF) as u8;
header[2] = ((value >> 16) & 0xFF) as u8;
3
} else if len <= MAX_PACKET_LENGTH {
header.copy_from_slice(&(((len << 2) as u32) | 0x3).to_le_bytes());
4
} else {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"Message too large",
));
};
let slices = [IoSlice::new(&header[..header_len]), IoSlice::new(data)];
let written = writer.write_vectored(&slices).await?;
if written == 0 {
return Err(io::Error::new(
io::ErrorKind::WriteZero,
"failed to write RustDesk frame",
));
}
if written < header_len {
writer.write_all(&header[written..header_len]).await?;
writer.write_all(data).await?;
} else {
writer.write_all(&data[written - header_len..]).await?;
}
Ok(())
}
pub async fn write_frame_buffered<W: AsyncWrite + Unpin>(
writer: &mut W,
data: &[u8],
@@ -281,4 +340,32 @@ mod tests {
let decoded = codec.decode(&mut buf).unwrap().unwrap();
assert_eq!(decoded.len(), 100000);
}
#[tokio::test]
async fn read_limit_rejects_length_before_allocating_payload() {
let encoded = encode_frame(&vec![0u8; 1024]).unwrap();
let (mut writer, mut reader) = tokio::io::duplex(encoded.len());
tokio::spawn(async move {
writer.write_all(&encoded).await.unwrap();
});
let error = read_frame_with_limit(&mut reader, 128).await.unwrap_err();
assert_eq!(error.kind(), io::ErrorKind::InvalidData);
}
#[tokio::test]
async fn vectored_writer_round_trips() {
let payload = vec![0x5a; 100_000];
let (mut writer, mut reader) = tokio::io::duplex(payload.len() + 4);
let expected = payload.clone();
let send = tokio::spawn(async move {
write_frame_vectored(&mut writer, &payload).await.unwrap();
});
let decoded = read_frame_with_limit(&mut reader, expected.len())
.await
.unwrap();
send.await.unwrap();
assert_eq!(decoded.as_ref(), expected.as_slice());
}
}

File diff suppressed because it is too large Load Diff

View File

@@ -91,6 +91,12 @@ impl VideoFrameAdapter {
return data;
}
// Parameter sets are relevant only on random-access frames. Avoid a
// full Annex-B/AVCC scan on every delta frame in every client session.
if !is_keyframe {
return data;
}
let (sps, pps) = crate::video::codec::h264_bitstream::extract_sps_pps(&data);
let mut has_sps = false;
let mut has_pps = false;

View File

@@ -17,7 +17,7 @@ use std::time::Duration;
use parking_lot::RwLock;
use protobuf::Message;
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::broadcast;
use tokio::sync::{broadcast, Semaphore};
use tokio::task::JoinHandle;
use tracing::{debug, error, info, warn};
@@ -33,6 +33,7 @@ use self::rendezvous::{AddrMangle, RendezvousMediator, RendezvousStatus};
const RELAY_CONNECT_TIMEOUT_MS: u64 = 10_000;
const SERVICE_SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(2);
const MAX_PENDING_CONNECTION_ATTEMPTS: usize = 8;
#[derive(Debug, Clone, PartialEq)]
pub enum ServiceStatus {
@@ -59,6 +60,7 @@ pub struct RustDeskService {
rendezvous: Arc<RwLock<Option<Arc<RendezvousMediator>>>>,
rendezvous_handle: Arc<RwLock<Option<JoinHandle<()>>>>,
tcp_listener_handle: Arc<RwLock<Option<Vec<JoinHandle<()>>>>>,
listener_start_lock: Arc<tokio::sync::Mutex<()>>,
listen_port: Arc<RwLock<u16>>,
connection_manager: Arc<ConnectionManager>,
video_manager: Arc<VideoStreamManager>,
@@ -84,6 +86,7 @@ impl RustDeskService {
rendezvous: Arc::new(RwLock::new(None)),
rendezvous_handle: Arc::new(RwLock::new(None)),
tcp_listener_handle: Arc::new(RwLock::new(None)),
listener_start_lock: Arc::new(tokio::sync::Mutex::new(())),
listen_port: Arc::new(RwLock::new(direct_access_port)),
connection_manager,
video_manager,
@@ -130,7 +133,7 @@ impl RustDeskService {
self.status() == ServiceStatus::Running
}
pub async fn start(&self) -> anyhow::Result<()> {
pub async fn start(self: &Arc<Self>) -> anyhow::Result<()> {
let config = self.config.read().clone();
if !config.enabled {
@@ -168,14 +171,16 @@ impl RustDeskService {
.set_video_manager(self.video_manager.clone());
if config.mode == RustDeskMode::DirectIp {
let (tcp_handles, listen_port) = match self.start_tcp_listener_with_port().await {
let listen_port = match self
.ensure_tcp_listener(config.direct_access_port, false)
.await
{
Ok(result) => result,
Err(err) => {
*self.status.write() = ServiceStatus::Error(err.to_string());
return Err(err);
}
};
*self.tcp_listener_handle.write() = Some(tcp_handles);
*self.listen_port.write() = listen_port;
*self.status.write() = ServiceStatus::Running;
return Ok(());
@@ -193,14 +198,23 @@ impl RustDeskService {
let connection_manager = self.connection_manager.clone();
let service_config = self.config.clone();
let connection_attempts = Arc::new(Semaphore::new(MAX_PENDING_CONNECTION_ATTEMPTS));
mediator.set_punch_callback(Arc::new({
let connection_manager = connection_manager.clone();
let service_config = service_config.clone();
let connection_attempts = connection_attempts.clone();
move |peer_addr, rendezvous_addr, relay_server, uuid, socket_addr, device_id| {
let conn_mgr = connection_manager.clone();
let config = service_config.clone();
let attempts = connection_attempts.clone();
tokio::spawn(async move {
let Ok(_permit) = attempts.try_acquire_owned() else {
warn!(
"Dropping RustDesk punch request: too many pending connection attempts"
);
return;
};
if let Some(addr) = peer_addr {
info!("Attempting P2P direct connection to {}", addr);
match punch::try_direct_connection(addr).await {
@@ -238,10 +252,18 @@ impl RustDeskService {
mediator.set_relay_callback(Arc::new({
let connection_manager = connection_manager.clone();
let service_config = service_config.clone();
let connection_attempts = connection_attempts.clone();
move |rendezvous_addr, relay_server, uuid, socket_addr, device_id| {
let conn_mgr = connection_manager.clone();
let config = service_config.clone();
let attempts = connection_attempts.clone();
tokio::spawn(async move {
let Ok(_permit) = attempts.try_acquire_owned() else {
warn!(
"Dropping RustDesk relay request: too many pending connection attempts"
);
return;
};
let relay_key = rustdesk_relay_key(&config);
if let Err(e) = handle_relay_request(
&rendezvous_addr,
@@ -260,19 +282,39 @@ impl RustDeskService {
}
}));
let connection_manager2 = self.connection_manager.clone();
let weak_service = Arc::downgrade(self);
let intranet_attempts = connection_attempts.clone();
mediator.set_intranet_callback(Arc::new(
move |rendezvous_addr, peer_socket_addr, local_addr, relay_server, device_id| {
let conn_mgr = connection_manager2.clone();
move |rendezvous_addr, peer_socket_addr, local_ip, relay_server, device_id| {
let weak_service = weak_service.clone();
let attempts = intranet_attempts.clone();
tokio::spawn(async move {
let Ok(_permit) = attempts.try_acquire_owned() else {
warn!("Dropping RustDesk intranet request: too many pending connection attempts");
return;
};
let Some(service) = weak_service.upgrade() else {
return;
};
let preferred_port = service.config.read().direct_access_port;
let listen_port = match service
.ensure_tcp_listener(preferred_port, true)
.await
{
Ok(port) => port,
Err(error) => {
error!("Failed to start on-demand RustDesk listener: {}", error);
return;
}
};
let local_addr = SocketAddr::new(local_ip, listen_port);
if let Err(e) = handle_intranet_request(
&rendezvous_addr,
&peer_socket_addr,
local_addr,
&relay_server,
&device_id,
conn_mgr,
service.connection_manager.clone(),
)
.await
{
@@ -304,9 +346,30 @@ impl RustDeskService {
Ok(())
}
async fn start_tcp_listener_with_port(&self) -> anyhow::Result<(Vec<JoinHandle<()>>, u16)> {
let direct_access_port = self.config.read().direct_access_port;
let (listeners, listen_port) = self.bind_direct_listeners(direct_access_port)?;
async fn ensure_tcp_listener(
self: &Arc<Self>,
preferred_port: u16,
allow_ephemeral_fallback: bool,
) -> anyhow::Result<u16> {
let _guard = self.listener_start_lock.lock().await;
if self.tcp_listener_handle.read().is_some() {
return Ok(*self.listen_port.read());
}
if self.status() == ServiceStatus::Stopped {
anyhow::bail!("RustDesk service stopped before listener could start");
}
let (listeners, listen_port) = match self.bind_direct_listeners(preferred_port) {
Ok(result) => result,
Err(error) if allow_ephemeral_fallback => {
warn!(
"RustDesk port {} unavailable for on-demand listening: {}; using an ephemeral port",
preferred_port, error
);
self.bind_direct_listeners(0)?
}
Err(error) => return Err(error),
};
*self.listen_port.write() = listen_port;
@@ -326,15 +389,13 @@ impl RustDeskService {
match result {
Ok((stream, peer_addr)) => {
info!("Accepted direct connection from {}", peer_addr);
let conn_mgr = conn_mgr.clone();
tokio::spawn(async move {
if let Err(e) = conn_mgr.accept_direct_connection(stream, peer_addr).await {
error!("Failed to handle direct connection from {}: {}", peer_addr, e);
}
});
if let Err(e) = conn_mgr.accept_listener_connection(stream, peer_addr).await {
warn!("Rejected direct connection from {}: {}", peer_addr, e);
}
}
Err(e) => {
error!("TCP accept error: {}", e);
tokio::time::sleep(Duration::from_millis(100)).await;
}
}
}
@@ -348,7 +409,8 @@ impl RustDeskService {
handles.push(handle);
}
Ok((handles, listen_port))
*self.tcp_listener_handle.write() = Some(handles);
Ok(listen_port)
}
fn bind_direct_listeners(&self, port: u16) -> anyhow::Result<(Vec<TcpListener>, u16)> {
@@ -382,8 +444,8 @@ impl RustDeskService {
info!("Stopping RustDesk service");
let _ = self.shutdown_tx.send(());
self.connection_manager.close_all();
let _listener_guard = self.listener_start_lock.lock().await;
*self.status.write() = ServiceStatus::Stopped;
if let Some(mediator) = self.rendezvous.read().as_ref() {
mediator.stop();
@@ -401,13 +463,15 @@ impl RustDeskService {
}
}
// No listener can admit a new session after this point.
self.connection_manager.close_all().await;
*self.rendezvous.write() = None;
*self.status.write() = ServiceStatus::Stopped;
Ok(())
}
pub async fn restart(&self, config: RustDeskConfig) -> anyhow::Result<()> {
pub async fn restart(self: &Arc<Self>, config: RustDeskConfig) -> anyhow::Result<()> {
self.stop().await?;
self.update_config(config);
self.start().await

View File

@@ -128,7 +128,7 @@ pub type RelayCallback = Arc<dyn Fn(String, String, String, Vec<u8>, String) + S
pub type PunchCallback =
Arc<dyn Fn(Option<SocketAddr>, String, String, String, Vec<u8>, String) + Send + Sync>;
pub type IntranetCallback = Arc<dyn Fn(String, Vec<u8>, SocketAddr, String, String) + Send + Sync>;
pub type IntranetCallback = Arc<dyn Fn(String, Vec<u8>, IpAddr, String, String) + Send + Sync>;
pub struct RendezvousMediator {
config: Arc<RwLock<RustDeskConfig>>,
@@ -143,7 +143,6 @@ pub struct RendezvousMediator {
relay_callback: Arc<RwLock<Option<RelayCallback>>>,
punch_callback: Arc<RwLock<Option<PunchCallback>>>,
intranet_callback: Arc<RwLock<Option<IntranetCallback>>>,
listen_port: Arc<RwLock<u16>>,
shutdown_tx: broadcast::Sender<()>,
}
@@ -166,23 +165,10 @@ impl RendezvousMediator {
relay_callback: Arc::new(RwLock::new(None)),
punch_callback: Arc::new(RwLock::new(None)),
intranet_callback: Arc::new(RwLock::new(None)),
listen_port: Arc::new(RwLock::new(21118)),
shutdown_tx,
}
}
pub fn set_listen_port(&self, port: u16) {
let old_port = *self.listen_port.read();
if old_port != port {
*self.listen_port.write() = port;
self.increment_serial();
}
}
pub fn listen_port(&self) -> u16 {
*self.listen_port.read()
}
pub fn increment_serial(&self) {
let mut serial = self.serial.write();
*serial = serial.wrapping_add(1);
@@ -430,7 +416,9 @@ impl RendezvousMediator {
) -> anyhow::Result<()> {
let id = self.device_id();
let local_addrs = get_local_addresses();
let local_addrs = tokio::task::spawn_blocking(get_local_addresses)
.await
.map_err(|error| anyhow::anyhow!("Failed to inspect local addresses: {error}"))?;
if local_addrs.is_empty() {
debug!("No local addresses available for LocalAddr response");
return Ok(());
@@ -439,21 +427,18 @@ impl RendezvousMediator {
let config = self.config.read().clone();
let rendezvous_addr = config.rendezvous_addr();
let listen_port = self.listen_port();
let local_ip = local_addrs[0];
let local_sock_addr = SocketAddr::new(local_ip, listen_port);
info!(
"FetchLocalAddr: calling intranet callback with local_addr={}, rendezvous={}",
local_sock_addr, rendezvous_addr
"FetchLocalAddr: requesting an on-demand listener for {}, rendezvous={}",
local_ip, rendezvous_addr
);
if let Some(callback) = self.intranet_callback.read().as_ref() {
callback(
rendezvous_addr,
peer_socket_addr.to_vec(),
local_sock_addr,
local_ip,
relay_server.to_string(),
id,
);

View File

@@ -47,6 +47,9 @@ async fn current_status(
config: RustDeskConfigResponse::from(&config),
service_status: runtime.service_status,
rendezvous_status: runtime.rendezvous_status,
connection_count: runtime.connection_count,
listening: runtime.listening,
listen_port: runtime.listen_port,
}
}
@@ -86,6 +89,9 @@ pub struct RustDeskStatusResponse {
pub config: RustDeskConfigResponse,
pub service_status: String,
pub rendezvous_status: Option<String>,
pub connection_count: usize,
pub listening: bool,
pub listen_port: Option<u16>,
}
pub async fn get_rustdesk_config(