Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 14 additions & 7 deletions crates/ironrdp-graphics/src/progressive.rs
Original file line number Diff line number Diff line change
Expand Up @@ -165,7 +165,8 @@ fn decode_first_pass_to_dwtq(
///
/// # Errors
///
/// Returns [`SrlError`] for a malformed or truncated SRL stream.
/// Returns [`SrlError`] when an SRL magnitude requires an invalid number of bits.
/// Missing trailing SRL entries read as zero bits, so truncation is not detected.
/// See MS-RDPEGFX section 3.3.8.2.1.2.
pub fn decode_upgrade_pass(
srl_data: &[u8],
Expand All @@ -187,7 +188,7 @@ pub fn decode_upgrade_pass(
.saturating_sub(curr_prog_quant.for_band(band_idx));
band_idx != NUM_BANDS - 1 && num_bits != 0 && zero_counts[band_idx] != 0
});
let mut srl_decoder = has_srl_values.then(|| srl::SrlDecoder::new(srl_data)).transpose()?;
let mut srl_decoder = has_srl_values.then(|| srl::SrlDecoder::new(srl_data));
let mut srl_values = Vec::with_capacity(NUM_BANDS);

for (band_idx, _) in bands.iter().enumerate() {
Expand Down Expand Up @@ -2075,14 +2076,16 @@ mod tests {
}

#[test]
fn upgrade_pass_rejects_truncated_srl() {
fn upgrade_pass_tolerates_truncated_srl() {
let mut coefficients = [0i16; COEFFICIENTS_PER_COMPONENT];
let mut sign = [SIGN_POSITIVE; COEFFICIENTS_PER_COMPONENT];
sign[0] = SIGN_ZERO;

let mut prev_prog_quant = ComponentCodecQuant::LOSSLESS;
prev_prog_quant.hl1 = 4;

// Bits past the end of the SRL stream read as zeros, as in the reference decoder,
// so the cut-off unary magnitude decodes as the maximum rather than failing.
assert_eq!(
decode_upgrade_pass(
&[0x80, 0x00],
Expand All @@ -2093,8 +2096,10 @@ mod tests {
&mut coefficients,
&mut sign,
),
Err(SrlError::Truncated)
Ok(())
);
assert_eq!(coefficients[0], 15);
assert_eq!(sign[0], SIGN_POSITIVE);
}

#[test]
Expand All @@ -2108,22 +2113,24 @@ mod tests {
tile.sign[0][0] = SIGN_ZERO;
tile.sign[1][0] = SIGN_ZERO;

// The second component asks for a 20-bit magnitude, which SRL cannot represent.
tile.prog_quant[1].hl1 = 20;
let prog_quant = tile.prog_quant;
let coefficients = tile.coefficients;
let sign = tile.sign;

assert_eq!(
tile.decode_upgrade(
[&[0x90, 0x00], &[0x80, 0x00], &[]],
[&[], &[], &[]],
[ComponentCodecQuant::LOSSLESS; 3],
75,
),
Err(SrlError::Truncated)
Err(SrlError::InvalidBitCount(20))
);

assert_eq!(tile.coefficients, coefficients);
assert_eq!(tile.sign, sign);
assert_eq!(tile.prog_quant, [prev_prog_quant; 3]);
assert_eq!(tile.prog_quant, prog_quant);
assert_eq!(tile.pass, 1);
assert_eq!(tile.quality, 50);
}
Expand Down
175 changes: 110 additions & 65 deletions crates/ironrdp-graphics/src/srl.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,29 +6,28 @@

const INITIAL_KP: u8 = 8;
const MAX_KP: u8 = 80;
// This conservative malformed-stream bound includes LL3 entries, although LL3 is raw-coded.
// One component holds 4096 coefficients (LL3 included, although LL3 is raw-coded).
// The encoder rejects longer runs; the decoder consumes individual run events as
// coefficients are requested, so an encoded run may exceed the remaining entries.
const MAX_ZERO_RUN: usize = 4096;

/// Errors encountered while decoding or encoding an SRL stream.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SrlError {
/// The required trailing zero byte is absent.
MissingTerminator,
/// The stream ended before a complete code word was read.
Truncated,
/// An SRL value requires between one and fifteen magnitude bits.
InvalidBitCount(u8),
/// A value cannot be represented by the magnitude width.
MagnitudeOutOfRange { magnitude: u16, max: u16 },
/// A zero run exceeds the number of coefficients in one component.
///
/// Only the encoder reports this; the decoder accepts overshooting runs because
/// Windows legitimately produces them.
ZeroRunTooLong,
}

impl core::fmt::Display for SrlError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::MissingTerminator => write!(f, "srl stream is missing its trailing zero byte"),
Self::Truncated => write!(f, "srl stream is truncated"),
Self::InvalidBitCount(bits) => write!(f, "invalid srl magnitude bit count {bits}"),
Self::MagnitudeOutOfRange { magnitude, max } => {
write!(f, "srl magnitude {magnitude} exceeds maximum {max}")
Expand All @@ -52,22 +51,21 @@ pub struct SrlDecoder<'a> {
}

impl<'a> SrlDecoder<'a> {
/// Create a decoder for an SRL stream, excluding its required trailing zero byte.
pub fn new(data: &'a [u8]) -> Result<Self, SrlError> {
let Some((&terminator, payload)) = data.split_last() else {
return Err(SrlError::MissingTerminator);
};

if terminator != 0 {
return Err(SrlError::MissingTerminator);
}

Ok(Self {
reader: BitReader::new(payload),
/// Create a decoder for an SRL stream.
///
/// All bytes are retained as data until the requested coefficients have been decoded.
/// A final zero byte may contain a value's sign or magnitude bits, so it cannot be
/// identified as a terminator in advance. Windows also omits trailing zero entries;
/// bits past the end read as zeros, matching the reference decoder.
///
/// This tolerance cannot distinguish omitted trailing entries from a truncated stream.
pub fn new(data: &'a [u8]) -> Self {
Self {
reader: BitReader::new(data),
kp: INITIAL_KP,
zero_run_remaining: 0,
nonzero_pending: false,
})
}
}

/// Decode `num_values` entries for one DWT band.
Expand All @@ -90,55 +88,40 @@ impl<'a> SrlDecoder<'a> {
continue;
}

self.zero_run_remaining = self.decode_zero_run()?;
self.nonzero_pending = true;
}

Ok(output)
}

fn decode_zero_run(&mut self) -> Result<usize, SrlError> {
let mut zeros = 0usize;

loop {
// Consume one zero-run event at a time. A zero bit encodes a chunk of
// zeros; a one bit encodes the final chunk followed by a nonzero value.
// Retaining that pending value even at EOF matters: its zero-filled sign
// and unary bits represent the positive maximum, not another zero run.
let k = self.kp / 8;

if self.reader.read_bit()? {
let tail = usize::try_from(self.reader.read_bits(k)?).map_err(|_| SrlError::ZeroRunTooLong)?;
if self.reader.read_bit() {
self.zero_run_remaining = usize::from(self.reader.read_bits(k));
self.kp = self.kp.saturating_sub(6);

let zeros = zeros.checked_add(tail).ok_or(SrlError::ZeroRunTooLong)?;
return (zeros <= MAX_ZERO_RUN).then_some(zeros).ok_or(SrlError::ZeroRunTooLong);
self.nonzero_pending = true;
} else {
self.zero_run_remaining = 1usize << k;
self.kp = self.kp.saturating_add(4).min(MAX_KP);
}

let chunk = 1usize << k;
zeros = zeros.checked_add(chunk).ok_or(SrlError::ZeroRunTooLong)?;
if zeros > MAX_ZERO_RUN {
return Err(SrlError::ZeroRunTooLong);
}

self.kp = self.kp.saturating_add(4).min(MAX_KP);
}

Ok(output)
}

fn decode_nonzero(&mut self, num_bits: u8) -> Result<i16, SrlError> {
let maximum = max_magnitude(num_bits)?;
let sign = self.reader.read_bit()?;
let sign = self.reader.read_bit();
let mut zero_count = 0u16;

while zero_count + 1 < maximum {
if self.reader.read_bit()? {
if self.reader.read_bit() {
break;
}

zero_count += 1;
}

let magnitude = zero_count + 1;
let magnitude = i16::try_from(magnitude).map_err(|_| SrlError::MagnitudeOutOfRange {
magnitude,
max: maximum,
})?;
// INVARIANT: maximum is at most 32767 because num_bits is in 1..=15.
let magnitude = i16::try_from(magnitude).expect("magnitude fits in i16");

Ok(if sign { -magnitude } else { magnitude })
}
Expand Down Expand Up @@ -238,7 +221,7 @@ impl Default for SrlEncoder {
/// magnitude width. Progressive tile decoding should use [`SrlDecoder`]
/// directly so its state continues between bands.
pub fn decode_srl(data: &[u8], num_values: usize, num_bits: u8) -> Result<Vec<i16>, SrlError> {
let mut decoder = SrlDecoder::new(data)?;
let mut decoder = SrlDecoder::new(data);
decoder.decode(num_values, num_bits)
}

Expand Down Expand Up @@ -276,9 +259,12 @@ impl<'a> BitReader<'a> {
}
}

fn read_bit(&mut self) -> Result<bool, SrlError> {
/// Reads one bit; bits past the end of the stream read as zero, like the reference
/// decoder's zero-filled bit accumulator, so a stream that omits its trailing zero
/// entries still decodes.
fn read_bit(&mut self) -> bool {
let Some(&byte) = self.data.get(self.byte_idx) else {
return Err(SrlError::Truncated);
return false;
};

let bit = (byte >> (7 - self.bit_idx)) & 1 != 0;
Expand All @@ -288,15 +274,15 @@ impl<'a> BitReader<'a> {
self.byte_idx += 1;
}

Ok(bit)
bit
}

fn read_bits(&mut self, count: u8) -> Result<u32, SrlError> {
let mut value = 0u32;
fn read_bits(&mut self, count: u8) -> u16 {
let mut value = 0u16;
for _ in 0..count {
value = (value << 1) | u32::from(self.read_bit()?);
value = (value << 1) | u16::from(self.read_bit());
}
Ok(value)
value
}
}

Expand Down Expand Up @@ -355,7 +341,7 @@ mod tests {
fn preserves_zero_run_and_kp_between_bands() {
// A two-zero run (010) spans the first and second calls.
// The following positive magnitude-one value uses K=0 after the run.
let mut decoder = SrlDecoder::new(&[0x48, 0x00]).unwrap();
let mut decoder = SrlDecoder::new(&[0x48, 0x00]);
assert_eq!(decoder.decode(1, 4), Ok(vec![0]));
assert_eq!(decoder.decode(2, 4), Ok(vec![0, 1]));
}
Expand All @@ -368,13 +354,55 @@ mod tests {
}

#[test]
fn rejects_truncated_stream() {
assert_eq!(decode_srl(&[0x80, 0x00], 1, 4), Err(SrlError::Truncated));
fn reads_past_the_end_as_zero_bits() {
// Zero run 0 (10), then the sign and unary magnitude bits run off the end of
// the stream and read as zeros: an unterminated maximum magnitude, as the
// reference decoder's zero-filled accumulator produces.
assert_eq!(decode_srl(&[0x80, 0x00], 1, 4), Ok(vec![15]));
}

#[test]
fn preserves_a_zero_data_byte_and_a_pending_value_at_eof() {
// 100001 is +3. The next 10 is a zero-length run at K=0 followed by
// the next value's positive sign. Its magnitude is all zeros: +15.
// With or without a final zero byte, the zero-filled reader must finish
// this pending value instead of replacing it with omitted zero entries.
for data in [&[0x86, 0x00][..], &[0x86][..], &[0x86, 0x00, 0x00][..]] {
assert_eq!(decode_srl(data, 3, 4), Ok(vec![3, 15, 0]));
}

// A run's final codeword ends exactly at the byte boundary here. The
// pending value crosses a band boundary and obtains all its bits at EOF.
let mut decoder = SrlDecoder::new(&[0x85]);
assert_eq!(decoder.decode(2, 4), Ok(vec![3, 0]));
assert_eq!(decoder.decode(2, 4), Ok(vec![15, 0]));
}

#[test]
fn rejects_missing_terminator() {
assert_eq!(decode_srl(&[0x84], 1, 4), Err(SrlError::MissingTerminator));
fn tolerates_a_missing_terminator() {
// Same bits as `decodes_initial_kp_and_unary_magnitude` without the zero byte.
assert_eq!(decode_srl(&[0x84], 1, 4), Ok(vec![3]));
}

#[test]
fn overshooting_zero_runs_are_accepted() {
// Forty `0` bits grow KP to 80 and add far more than 4096 zeros; every entry is zero.
assert_eq!(decode_srl(&[0x00; 5], 6, 4), Ok(vec![0; 6]));
// An explicit run of 4096 followed by a value: the run covers the entries asked for.
let mut decoder = SrlDecoder::new(&[0x00; 5]);
assert_eq!(decoder.decode(4096, 4), Ok(vec![0; 4096]));
}

#[test]
fn trailing_entries_omitted_by_the_encoder_decode_as_zeros() {
// Windows stops writing once every remaining entry is zero.
assert_eq!(decode_srl(&[], 5, 4), Ok(vec![0; 5]));
assert_eq!(decode_srl(&[0x84, 0x00], 4, 4), Ok(vec![3, 0, 0, 0]));

// The implicit zeros continue across bands, like an explicit run would.
let mut decoder = SrlDecoder::new(&[0x84]);
assert_eq!(decoder.decode(2, 4), Ok(vec![3, 0]));
assert_eq!(decoder.decode(3, 4), Ok(vec![0, 0, 0]));
}

#[test]
Expand All @@ -400,7 +428,7 @@ mod tests {
let original = [0, 0, 1, -1, 0, 3];
let encoded = encode_srl(&original, 4).unwrap();
assert_eq!(encoded, vec![0x4F, 0x44, 0x00]);
let mut decoder = SrlDecoder::new(&encoded).unwrap();
let mut decoder = SrlDecoder::new(&encoded);
assert_eq!(decoder.decode(2, 4), Ok(vec![0, 0]));
assert_eq!(decoder.decode(1, 4), Ok(vec![1]));
assert_eq!(decoder.decode(1, 4), Ok(vec![-1]));
Expand All @@ -409,6 +437,23 @@ mod tests {
assert_eq!(decode_srl(&encoded, original.len(), 4), Ok(original.to_vec()));
}

#[test]
fn round_trips_streams_with_and_without_terminators() {
for num_bits in 1..=5 {
let maximum = i16::try_from(max_magnitude(num_bits).unwrap()).unwrap();
for zeros in 0..16 {
for value in [-maximum, -1, 1, maximum] {
let mut values = vec![0; zeros];
values.extend([3.min(maximum), value, 0, -value]);
let encoded = encode_srl(&values, num_bits).unwrap();
for data in [encoded.as_slice(), &encoded[..encoded.len() - 1]] {
assert_eq!(decode_srl(data, values.len(), num_bits).unwrap(), values);
}
}
}
}
}

#[test]
fn preserves_zero_runs_between_bands_when_encoding() {
let mut encoder = SrlEncoder::new();
Expand Down
Loading