use rat_rdp_bulk::BulkCompressor;
use rat_rdp_core::{WriteBuf, decode};
use rat_rdp_dvc::{DrdynvcClient, DvcClientProcessor, DynamicChannelRef};
use rat_rdp_pdu::gcc::{ChannelName, Monitor};
use rat_rdp_pdu::mcs::{DisconnectProviderUltimatum, DisconnectReason, McsMessage, SendDataIndicationCtx};
use rat_rdp_pdu::rdp::autodetect::{AutoDetectReqPdu, AutoDetectRequest, AutoDetectResponse, AutoDetectRspPdu};
use rat_rdp_pdu::rdp::client_info::CompressionType;
use rat_rdp_pdu::rdp::headers::{CompressionFlags, ShareDataCtx, ShareDataPdu};
use rat_rdp_pdu::rdp::multitransport::MultitransportRequestPdu;
use rat_rdp_pdu::rdp::server_error_info::{ErrorInfo, ProtocolIndependentCode, ServerSetErrorInfoPdu};
use rat_rdp_pdu::rdp::session_info::{InfoData, SaveSessionInfoPdu, ServerAutoReconnect};
use rat_rdp_pdu::x224::X224;
use rat_rdp_svc::{
StaticChannelSet, SvcMessage, SvcProcessor, SvcProcessorMessages, client_encode_svc_messages_with_max_chunk_len,
};
use tracing::debug;
use crate::{SessionError, SessionErrorExt as _, SessionResult, reason_err};
#[derive(Debug, Clone)]
pub enum ProcessorOutput {
ResponseFrame(Vec<u8>),
Disconnect(DisconnectDescription),
DeactivateAll,
SaveSessionInfo { logon_complete: bool },
MultitransportRequest(MultitransportRequestPdu),
AutoReconnectCookie(ServerAutoReconnect),
AutoReconnectFailed,
AutoDetect(AutoDetectRequest),
GraphicsUpdate(Vec<u8>),
PointerUpdate(Vec<u8>),
MonitorLayout(Vec<Monitor>),
}
#[derive(Debug, Clone)]
pub enum DisconnectDescription {
McsDisconnect(DisconnectReason),
ErrorInfo(ErrorInfo),
}
pub struct Processor {
static_channels: StaticChannelSet,
user_channel_id: u16,
io_channel_id: u16,
message_channel_id: Option<u16>,
share_id: u32,
}
impl Processor {
pub fn new(
static_channels: StaticChannelSet,
user_channel_id: u16,
io_channel_id: u16,
message_channel_id: Option<u16>,
share_id: u32,
) -> Self {
Self {
static_channels,
user_channel_id,
io_channel_id,
message_channel_id,
share_id,
}
}
pub fn set_share_id(&mut self, share_id: u32) {
self.share_id = share_id;
}
pub fn set_static_channel_chunk_size(&mut self, maximum_chunk_size: usize) -> bool {
self.static_channels.set_maximum_chunk_size(maximum_chunk_size)
}
pub fn static_channel_chunk_size(&self) -> usize {
self.static_channels.maximum_chunk_size()
}
pub fn get_svc_processor<T: SvcProcessor + 'static>(&self) -> Option<&T> {
self.static_channels
.get_by_type::<T>()
.and_then(|svc| svc.channel_processor_downcast_ref())
}
pub fn get_svc_processor_mut<T: SvcProcessor + 'static>(&mut self) -> Option<&mut T> {
self.static_channels
.get_by_type_mut::<T>()
.and_then(|svc| svc.channel_processor_downcast_mut())
}
pub fn process_svc_processor_messages<C: SvcProcessor + 'static>(
&self,
messages: SvcProcessorMessages<C>,
) -> SessionResult<Vec<u8>> {
let channel_id = self
.static_channels
.get_channel_id_by_type::<C>()
.ok_or_else(|| reason_err!("SVC", "channel not found"))?;
process_svc_messages(
messages.into(),
channel_id,
self.user_channel_id,
self.static_channels.maximum_chunk_size(),
)
}
pub fn process_svc_messages_by_name(
&self,
channel_name: &ChannelName,
messages: Vec<SvcMessage>,
) -> SessionResult<Vec<u8>> {
let channel_id = self
.static_channels
.get_channel_id_by_channel_name(channel_name)
.ok_or_else(|| reason_err!("SVC", "channel not found"))?;
process_svc_messages(
messages,
channel_id,
self.user_channel_id,
self.static_channels.maximum_chunk_size(),
)
}
pub fn get_dvc<T: DvcClientProcessor + 'static>(&self) -> Option<DynamicChannelRef<'_, T>> {
self.get_svc_processor::<DrdynvcClient>()?.get_dvc::<T>()
}
pub fn get_dvc_by_channel_id<T: DvcClientProcessor + 'static>(
&self,
channel_id: u32,
) -> Option<DynamicChannelRef<'_, T>> {
self.get_svc_processor::<DrdynvcClient>()?
.get_dvc_by_channel_id(channel_id)
}
pub fn process(
&mut self,
frame: &[u8],
bulk_decompressor: &mut Option<BulkCompressor>,
) -> SessionResult<Vec<ProcessorOutput>> {
let data_ctx: SendDataIndicationCtx<'_> = match rat_rdp_pdu::mcs::decode_send_data_indication(frame) {
Ok(data_ctx) => data_ctx,
Err(error) => {
if let Ok(X224(McsMessage::DisconnectProviderUltimatum(ultimatum))) =
decode::<X224<McsMessage<'_>>>(frame)
{
debug!(reason = ?ultimatum.reason, "Received Disconnect Provider Ultimatum, session will be closed");
return Ok(vec![ProcessorOutput::Disconnect(DisconnectDescription::McsDisconnect(
ultimatum.reason,
))]);
}
return Err(SessionError::decode(error));
}
};
let channel_id = data_ctx.channel_id;
if channel_id == self.io_channel_id {
self.process_io_channel(data_ctx, bulk_decompressor)
} else if self.message_channel_id == Some(channel_id) {
self.process_message_channel(data_ctx)
} else {
let maximum_chunk_size = self.static_channels.maximum_chunk_size();
if let Some(svc) = self.static_channels.get_by_channel_id_mut(channel_id) {
let response_pdus = svc.process(data_ctx.user_data).map_err(SessionError::pdu)?;
process_svc_messages(response_pdus, channel_id, data_ctx.initiator_id, maximum_chunk_size)
.map(|data| vec![ProcessorOutput::ResponseFrame(data)])
} else {
Err(reason_err!("X224", "unexpected channel received: ID {channel_id}"))
}
}
}
fn process_io_channel(
&mut self,
data_ctx: SendDataIndicationCtx<'_>,
bulk_decompressor: &mut Option<BulkCompressor>,
) -> SessionResult<Vec<ProcessorOutput>> {
debug_assert_eq!(data_ctx.channel_id, self.io_channel_id);
let io_channel = rat_rdp_pdu::rdp::headers::decode_io_channel(data_ctx).map_err(SessionError::decode)?;
match io_channel {
rat_rdp_pdu::rdp::headers::IoChannelPdu::Data(ctx) => Self::process_share_data(ctx, bulk_decompressor),
rat_rdp_pdu::rdp::headers::IoChannelPdu::MultitransportRequest(pdu) => {
debug!(
"Received Initiate Multitransport Request: request_id={}",
pdu.request_id
);
Ok(vec![ProcessorOutput::MultitransportRequest(pdu)])
}
rat_rdp_pdu::rdp::headers::IoChannelPdu::DeactivateAll(_) => Ok(vec![ProcessorOutput::DeactivateAll]),
}
}
fn process_share_data(
ctx: ShareDataCtx,
bulk_decompressor: &mut Option<BulkCompressor>,
) -> SessionResult<Vec<ProcessorOutput>> {
let ShareDataCtx {
compression_flags,
compression_type,
pdu,
..
} = ctx;
let (pdu, compression_flags) = match pdu {
ShareDataPdu::Compressed { pdu_type, data } => {
let data = Self::decompress_share_data(data, compression_flags, compression_type, bulk_decompressor)?;
(
ShareDataPdu::decode_with_type(&data, pdu_type).map_err(SessionError::decode)?,
CompressionFlags::empty(),
)
}
pdu => (pdu, compression_flags),
};
match pdu {
ShareDataPdu::SaveSessionInfo(session_info) => {
debug!("Got Session Save Info PDU: {session_info:?}");
let mut outputs = vec![ProcessorOutput::SaveSessionInfo {
logon_complete: is_logon_complete(&session_info),
}];
if let InfoData::LogonExtended(extended) = &session_info.info_data {
if let Some(cookie) = &extended.auto_reconnect {
outputs.push(ProcessorOutput::AutoReconnectCookie(cookie.clone()));
}
}
Ok(outputs)
}
ShareDataPdu::ArcStatusPdu(status) => {
if status != [0; 4] {
return Err(reason_err!("IO channel", "invalid auto-reconnect status PDU"));
}
Ok(vec![ProcessorOutput::AutoReconnectFailed])
}
ShareDataPdu::SetKeyboardIndicators(data) => {
debug!("Got Keyboard Indicators PDU: {data:?}");
Ok(Vec::new())
}
ShareDataPdu::ServerSetErrorInfo(ServerSetErrorInfoPdu(ErrorInfo::ProtocolIndependentCode(
ProtocolIndependentCode::None,
))) => {
debug!("Received None server error");
Ok(Vec::new())
}
ShareDataPdu::ServerSetErrorInfo(ServerSetErrorInfoPdu(e)) => {
let desc = DisconnectDescription::ErrorInfo(e);
Ok(vec![ProcessorOutput::Disconnect(desc)])
}
ShareDataPdu::ShutdownDenied => {
debug!("ShutdownDenied received, session will be closed");
let ultimatum = McsMessage::DisconnectProviderUltimatum(DisconnectProviderUltimatum::from_reason(
DisconnectReason::UserRequested,
));
let encoded_pdu = rat_rdp_core::encode_vec(&X224(ultimatum)).map_err(SessionError::encode);
Ok(vec![
ProcessorOutput::ResponseFrame(encoded_pdu?),
ProcessorOutput::Disconnect(DisconnectDescription::McsDisconnect(DisconnectReason::UserRequested)),
])
}
ShareDataPdu::Update(data) => {
let data = Self::decompress_share_data(data, compression_flags, compression_type, bulk_decompressor)?;
debug!("Got slow-path graphics update ({} bytes)", data.len());
Ok(vec![ProcessorOutput::GraphicsUpdate(data)])
}
ShareDataPdu::Pointer(data) => {
let data = Self::decompress_share_data(data, compression_flags, compression_type, bulk_decompressor)?;
debug!("Got slow-path pointer update ({} bytes)", data.len());
Ok(vec![ProcessorOutput::PointerUpdate(data)])
}
ShareDataPdu::MonitorLayout(monitor_layout) => {
Ok(vec![ProcessorOutput::MonitorLayout(monitor_layout.monitors)])
}
pdu => Err(reason_err!("IO channel", "unhandled PDU: {:?}", pdu.as_short_name())),
}
}
fn decompress_share_data(
data: Vec<u8>,
compression_flags: CompressionFlags,
compression_type: CompressionType,
bulk_decompressor: &mut Option<BulkCompressor>,
) -> SessionResult<Vec<u8>> {
if compression_flags.is_empty() {
return Ok(data);
}
let decompressor = bulk_decompressor
.as_mut()
.ok_or_else(|| reason_err!("slow-path", "received compressed share data without a decompressor"))?;
let flags = u32::from(compression_flags.bits()) | u32::from(compression_type.as_u8());
let decompressed = decompressor
.decompress(&data, flags)
.map_err(|error| reason_err!("slow-path", "bulk decompression failed: {error}"))?
.to_vec();
debug!(
compressed_size = data.len(),
decompressed_size = decompressed.len(),
?compression_type,
"Decompressed slow-path share data"
);
Ok(decompressed)
}
fn process_message_channel(&self, data_ctx: SendDataIndicationCtx<'_>) -> SessionResult<Vec<ProcessorOutput>> {
let Some(message_channel_id) = self.message_channel_id else {
return Err(reason_err!("message channel", "no message channel negotiated"));
};
let req = decode::<AutoDetectReqPdu>(data_ctx.user_data).map_err(SessionError::decode)?;
match req.request {
AutoDetectRequest::RttRequest { sequence_number, .. } => {
let response = AutoDetectRspPdu::new(AutoDetectResponse::RttResponse { sequence_number });
let mut frame = WriteBuf::new();
rat_rdp_pdu::mcs::encode_send_data_request(
self.user_channel_id,
message_channel_id,
&response,
&mut frame,
)
.map_err(SessionError::encode)?;
debug!(sequence_number, "Responded to auto-detect RTT request");
Ok(vec![ProcessorOutput::ResponseFrame(frame.into_inner())])
}
req @ AutoDetectRequest::NetworkCharacteristicsResult { .. } => {
debug!(?req, "Received network characteristics from server");
Ok(vec![ProcessorOutput::AutoDetect(req)])
}
req => {
debug!(?req, "Auto-detect request not yet implemented");
Ok(Vec::new())
}
}
}
pub fn encode_static(&self, output: &mut WriteBuf, pdu: ShareDataPdu) -> SessionResult<usize> {
let written = rat_rdp_pdu::rdp::headers::encode_share_data(
self.user_channel_id,
self.io_channel_id,
self.share_id,
pdu,
output,
)
.map_err(SessionError::encode)?;
Ok(written)
}
}
fn is_logon_complete(session_info: &SaveSessionInfoPdu) -> bool {
matches!(
session_info.info_data,
InfoData::LogonInfoV1(_) | InfoData::LogonInfoV2(_) | InfoData::PlainNotify
)
}
fn process_svc_messages(
messages: Vec<SvcMessage>,
channel_id: u16,
initiator_id: u16,
maximum_chunk_size: usize,
) -> SessionResult<Vec<u8>> {
client_encode_svc_messages_with_max_chunk_len(messages, channel_id, initiator_id, maximum_chunk_size)
.map_err(SessionError::encode)
}
#[cfg(test)]
mod tests {
use rat_rdp_bulk::{CompressionType as BulkCompressionType, flags};
use rat_rdp_core::encode_vec;
use rat_rdp_pdu::gcc::MonitorFlags;
use rat_rdp_pdu::rdp::finalization_messages::MonitorLayoutPdu;
use rat_rdp_pdu::rdp::headers::ShareDataPduType;
use rat_rdp_pdu::rdp::session_info::{InfoType, LogonExFlags, LogonInfoExtended};
use super::*;
#[test]
fn processor_decompresses_slow_path_share_data() {
let source = vec![b'A'; 1024];
let mut compressor = BulkCompressor::new(BulkCompressionType::Rdp5);
let (compressed_size, flags) = compressor.compress(&source).expect("source should compress");
assert_ne!(flags & flags::PACKET_COMPRESSED, 0, "test data must be compressed");
let compressed = compressor.compressed_data(compressed_size).to_vec();
let mut bulk_decompressor = Some(BulkCompressor::new(BulkCompressionType::Rdp5));
let compression_flags = CompressionFlags::from_bits_retain(
u8::try_from(flags & !flags::COMPRESSION_TYPE_MASK).expect("bulk flags should fit in a byte"),
);
assert_eq!(
Processor::decompress_share_data(
compressed,
compression_flags,
CompressionType::K64,
&mut bulk_decompressor
)
.expect("compressed slow-path data should decompress"),
source
);
}
#[test]
fn processor_rejects_compressed_slow_path_data_without_a_decompressor() {
let mut bulk_decompressor = None;
assert!(
Processor::decompress_share_data(
vec![0],
CompressionFlags::COMPRESSED,
CompressionType::K64,
&mut bulk_decompressor
)
.is_err()
);
}
#[test]
fn processor_decompresses_compressed_save_session_info() {
let session_info = SaveSessionInfoPdu {
info_type: InfoType::PlainNotify,
info_data: InfoData::PlainNotify,
};
let source = encode_vec(&session_info).expect("encode save session info");
let mut compressor = BulkCompressor::new(BulkCompressionType::Rdp5);
let (compressed_size, flags) = compressor.compress(&source).expect("source should compress");
assert_ne!(flags & flags::PACKET_COMPRESSED, 0, "test data must be compressed");
let compressed = compressor.compressed_data(compressed_size).to_vec();
let mut bulk_decompressor = Some(BulkCompressor::new(BulkCompressionType::Rdp5));
let compression_flags = CompressionFlags::from_bits_retain(
u8::try_from(flags & !flags::COMPRESSION_TYPE_MASK).expect("bulk flags should fit in a byte"),
);
let outputs = Processor::process_share_data(
ShareDataCtx {
initiator_id: 0,
channel_id: 0,
share_id: 0,
pdu_source: 0,
compression_flags,
compression_type: CompressionType::K64,
pdu: ShareDataPdu::Compressed {
pdu_type: ShareDataPduType::SaveSessionInfo,
data: compressed,
},
},
&mut bulk_decompressor,
)
.expect("compressed save session info should be processed");
assert!(matches!(
outputs.as_slice(),
[ProcessorOutput::SaveSessionInfo { logon_complete: true }]
));
}
#[test]
fn extended_session_info_does_not_signal_login_completion() {
let session_info = SaveSessionInfoPdu {
info_type: InfoType::LogonExtended,
info_data: InfoData::LogonExtended(LogonInfoExtended {
present_fields_flags: LogonExFlags::AUTO_RECONNECT_COOKIE,
auto_reconnect: None,
errors_info: None,
}),
};
assert!(!is_logon_complete(&session_info));
}
#[test]
fn processor_gracefully_disconnects_on_provider_ultimatum() {
let frame = encode_vec(&X224(McsMessage::DisconnectProviderUltimatum(
DisconnectProviderUltimatum::from_reason(DisconnectReason::ProviderInitiated),
)))
.expect("encode disconnect provider ultimatum");
let mut processor = Processor::new(StaticChannelSet::new(), 1002, 1003, None, 0);
let outputs = processor
.process(&frame, &mut None)
.expect("disconnect provider ultimatum should not be a protocol error");
assert!(matches!(
outputs.as_slice(),
[ProcessorOutput::Disconnect(DisconnectDescription::McsDisconnect(
DisconnectReason::ProviderInitiated
))]
));
}
#[test]
fn plain_notify_signals_login_completion() {
let session_info = SaveSessionInfoPdu {
info_type: InfoType::PlainNotify,
info_data: InfoData::PlainNotify,
};
assert!(is_logon_complete(&session_info));
}
#[test]
fn processor_surfaces_valid_auto_reconnect_status() {
let mut bulk_decompressor = None;
let outputs = Processor::process_share_data(
ShareDataCtx {
initiator_id: 0,
channel_id: 0,
share_id: 0,
pdu_source: 0,
compression_flags: CompressionFlags::empty(),
compression_type: CompressionType::K64,
pdu: ShareDataPdu::ArcStatusPdu(vec![0; 4]),
},
&mut bulk_decompressor,
)
.expect("valid auto-reconnect status PDU should be processed");
assert!(matches!(outputs.as_slice(), [ProcessorOutput::AutoReconnectFailed]));
}
#[test]
fn processor_surfaces_monitor_layout() {
let monitors = vec![Monitor {
left: 0,
top: 0,
right: 799,
bottom: 599,
flags: MonitorFlags::PRIMARY,
}];
let mut bulk_decompressor = None;
let outputs = Processor::process_share_data(
ShareDataCtx {
initiator_id: 0,
channel_id: 0,
share_id: 0,
pdu_source: 0,
compression_flags: CompressionFlags::empty(),
compression_type: CompressionType::K64,
pdu: ShareDataPdu::MonitorLayout(MonitorLayoutPdu {
monitors: monitors.clone(),
}),
},
&mut bulk_decompressor,
)
.expect("monitor layout PDU should be processed");
let [ProcessorOutput::MonitorLayout(actual)] = outputs.as_slice() else {
panic!("expected a monitor layout output");
};
assert_eq!(actual, &monitors);
}
#[test]
fn processor_rejects_invalid_auto_reconnect_status() {
let mut bulk_decompressor = None;
assert!(
Processor::process_share_data(
ShareDataCtx {
initiator_id: 0,
channel_id: 0,
share_id: 0,
pdu_source: 0,
compression_flags: CompressionFlags::empty(),
compression_type: CompressionType::K64,
pdu: ShareDataPdu::ArcStatusPdu(vec![0, 0, 0]),
},
&mut bulk_decompressor,
)
.is_err()
);
}
}