diff --git a/crates/ironrdp-graphics/src/progressive.rs b/crates/ironrdp-graphics/src/progressive.rs index 5ea16f38a6..ec6d7d72fd 100644 --- a/crates/ironrdp-graphics/src/progressive.rs +++ b/crates/ironrdp-graphics/src/progressive.rs @@ -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], @@ -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() { @@ -2075,7 +2076,7 @@ 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; @@ -2083,6 +2084,8 @@ mod tests { 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], @@ -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] @@ -2108,9 +2113,11 @@ 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], &[]], @@ -2118,12 +2125,12 @@ mod tests { [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); } diff --git a/crates/ironrdp-graphics/src/srl.rs b/crates/ironrdp-graphics/src/srl.rs index 9503dc7380..3f032f33ae 100644 --- a/crates/ironrdp-graphics/src/srl.rs +++ b/crates/ironrdp-graphics/src/srl.rs @@ -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}") @@ -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 { - 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. @@ -90,44 +88,31 @@ 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 { - 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 { 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; } @@ -135,10 +120,8 @@ impl<'a> SrlDecoder<'a> { } 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 }) } @@ -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, SrlError> { - let mut decoder = SrlDecoder::new(data)?; + let mut decoder = SrlDecoder::new(data); decoder.decode(num_values, num_bits) } @@ -276,9 +259,12 @@ impl<'a> BitReader<'a> { } } - fn read_bit(&mut self) -> Result { + /// 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; @@ -288,15 +274,15 @@ impl<'a> BitReader<'a> { self.byte_idx += 1; } - Ok(bit) + bit } - fn read_bits(&mut self, count: u8) -> Result { - 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 } } @@ -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])); } @@ -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] @@ -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])); @@ -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();